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:
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 functionPero 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.