将PyTorch模块转为ScriptModule遇Vocab转换错误,如何使用自定义Python对象?
解决TorchScript转换自定义Python对象的报错问题
这个问题我之前也碰到过,核心原因是TorchScript无法自动解析普通的Python自定义类——它只支持特定的可序列化类型(比如张量、ScriptModule、ScriptClass,还有Python内置的基本类型)。你的Vocab是普通Python类,TorchScript不知道怎么把它转换成自己能识别的类型,所以抛出了这个错误。
下面给你两种针对性的解决方案,根据你的需求选择:
方案一:把自定义类改为TorchScript可识别的ScriptClass
如果你需要在TorchScript的方法(比如forward)中使用Vocab的属性或方法,那可以把Vocab改成torch.jit.ScriptClass,这是TorchScript专门用来定义自定义可序列化类型的方式:
import torch # 把普通Python类改成ScriptClass @torch.jit.script_class class Vocab(object): def __init__(self, name: str): self.name: str = name # 必须标注类型,TorchScript需要明确类型信息 def show(self): print("dict:" + self.name) class Model(torch.nn.Module): def __init__(self, ): super(Model, self).__init__() self.layers = torch.nn.Linear(2, 3) self.encoder = 4 self.vocab = Vocab("vocab") # 现在这个Vocab是TorchScript可识别的类型 def forward(self, x): name = self.vocab.name print("forward show encoder:" + str(self.encoder)) print("vocab:" + name) enc_hidden = [] step = len(x) // 2 for i in range(step): enc_hidden.append((x[2*i] + x[2*i + 1])/2) enc_hidden = torch.stack(enc_hidden, 0) enc_hidden = self.__show(enc_hidden) return self.layers(enc_hidden) @torch.jit.export def __show(self, x): return x + 1 model = Model() data = torch.randn(10, 2) script_model = torch.jit.script(model) print(script_model) r1 = model(data) print(r1)
修改要点:
- 给
Vocab加上@torch.jit.script_class装饰器 - 给类的属性(比如
self.name)标注明确的类型(str),TorchScript需要类型信息来完成转换 - 这样
Vocab就能被TorchScript正确序列化,Model里的self.vocab也能正常被转换了
方案二:标记属性为TorchScript忽略项
如果你的Vocab只在Python侧使用,TorchScript的方法(比如forward)其实不需要用到它,那可以用torch.jit.ignore装饰器告诉TorchScript跳过这个属性的转换:
import torch class Vocab(object): def __init__(self, name): self.name = name def show(self): print("dict:" + self.name) class Model(torch.nn.Module): def __init__(self, ): super(Model, self).__init__() self.layers = torch.nn.Linear(2, 3) self.encoder = 4 # 用torch.jit.ignore标记这个属性,TorchScript会忽略它 self.vocab = torch.jit.ignore()(Vocab("vocab")) def forward(self, x): # 注意:如果forward里需要用到self.vocab,这个方案就不适用了! name = self.vocab.name print("forward show encoder:" + str(self.encoder)) print("vocab:" + name) enc_hidden = [] step = len(x) // 2 for i in range(step): enc_hidden.append((x[2*i] + x[2*i + 1])/2) enc_hidden = torch.stack(enc_hidden, 0) enc_hidden = self.__show(enc_hidden) return self.layers(enc_hidden) @torch.jit.export def __show(self, x): return x + 1 model = Model() data = torch.randn(10, 2) script_model = torch.jit.script(model) print(script_model) r1 = model(data) print(r1)
注意事项:
- 这个方案只适合
Vocab不在TorchScript执行的代码路径里被使用的场景,如果forward里必须访问self.vocab,那还是得用方案一 torch.jit.ignore()会把这个属性标记为"Python-only",TorchScript不会尝试转换它,但Python侧依然可以正常使用
内容的提问来源于stack exchange,提问作者Chuanhua Yang
相关产品推荐
相关产品推荐

