Business
Jobs
  • About Us
  • Solutions
    • Job Postings
      Post your job and receive qualified candidates in 48h.
    • Candidate Assessments
      500+ technical and psychological tests, plus anti-fraud.
    • Headhunting
      Tailor-made executive search from start to finish.
    • Payroll + EOR
      Payroll dispersal and EOR across 15+ LATAM countries.
  • Pricing
  • Jobs

0

280
Views
PyTorch no puede encurtir lambda

Tengo un modelo que usa un LambdaLayer personalizado de la siguiente manera:

 class LambdaLayer(LightningModule): def __init__(self, fun): super(LambdaLayer, self).__init__() self.fun = fun def forward(self, x): return self.fun(x) class TorchCatEmbedding(LightningModule): def __init__(self, start, end): super(TorchCatEmbedding, self).__init__() self.lb = LambdaLayer(lambda x: x[:, start:end]) self.embedding = torch.nn.Embedding(50, 5) def forward(self, inputs): o = self.lb(inputs).to(torch.int32) o = self.embedding(o) return o.squeeze()

El modelo funciona perfectamente bien en CPU o 1 GPU. Sin embargo, cuando se ejecuta con PyTorch Lightning en más de 2 GPU, ocurre este error:

 AttributeError: Can't pickle local object 'TorchCatEmbedding.__init__.<locals>.<lambda>'

El propósito de usar una función lambda aquí es que, dado un tensor de inputs , quiero pasar solo las inputs[:, start:end] a la capa de embedding .

Mis preguntas:

  • ¿Hay alguna alternativa al uso de una lambda en este caso?
  • si no, ¿qué se debe hacer para que la función lambda funcione en este contexto?
over 4 years ago · Santiago Trujillo
1 answers
Answer question

0

Entonces, el problema no es la función lambda per se, es que pickle no funciona con funciones que no son solo funciones de nivel de módulo (la forma en que pickle trata las funciones es solo como referencias a algún nombre de nivel de módulo). Entonces, desafortunadamente, si necesita capturar los argumentos de start y end , no podrá usar un cierre, normalmente solo querrá algo como:

 def function_maker(start, end): def function(x): return x[:, start:end] return function

Pero esto lo llevará de vuelta al punto de partida, en lo que respecta al problema del decapado.

Entonces, intenta algo como:

 class Slicer: def __init__(self, start, end): self.start = start self.end = end def __call__(self, x): return x[:, self.start:self.end])

Entonces puedes usar:

 LambdaLayer(Slicer(start, end))

No estoy familiarizado con PyTorch, aunque me sorprende que no ofrezca la posibilidad de usar un backend de serialización diferente. El proyecto pathos/ dill puede seleccionar funciones arbitrarias, por ejemplo, y a menudo es más fácil usar eso. Pero creo que lo anterior debería resolver el problema.

over 4 years ago · Santiago Trujillo Report
Answer question
Find remote jobs

Discover the new way to find a job!

Top jobs
Top job categories
Business
Post vacancy Pricing Sales
Legal
Terms and conditions Privacy policy
© 2026 PeakU Inc. All Rights Reserved.
Andres GPT
Show me some job opportunities
There's an error!