Empresas
Empleos
  • Sobre nosotros
  • Soluciones
    • Publicación de vacantes
      Publica tu vacante y recibe candidatos calificados en 48h.
    • Evaluación de candidatos
      500+ pruebas técnicas y psicológicas, más anti-fraude.
    • Headhunting
      Búsqueda ejecutiva a la medida de principio a fin.
    • Nómina + EOR
      Dispersión de nómina y EOR en más de 15 países de LATAM.
  • Precios
  • Empleos

0

282
Vistas
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 Respuestas
Responde la pregunta

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 Denunciar
Responde la pregunta
Encuentra empleos remotos

¡Descubre la nueva forma de encontrar empleo!

Top de empleos
Top categorías de empleo
Empresas
Publicar vacante Precios Comercial
Legal
Términos y condiciones Política de privacidad
© 2026 PeakU Inc. All Rights Reserved.
Andres GPT
Recomiéndame algunas ofertas
Necesito ayuda