如何为第三方库中PyTorch模块的自定义compute方法生成可用的TorchScript脚本?
如何为第三方库中PyTorch模块的自定义compute方法生成可用的TorchScript脚本?
我之前也碰到过类似的问题,TorchScript对PyTorch Module的默认处理逻辑确实只聚焦于forward方法,第三方库的自定义方法很容易被忽略。咱们先拆解下你遇到的错误原因,再给出可行的解决方案:
错误原因分析
第一个错误:
AttributeError: 'RecursiveScriptModule' object has no attribute 'compute'
TorchScript默认只会把Module的forward方法编译进脚本模块,其他方法如果没有被显式标记为需要导出,就不会被包含在生成的RecursiveScriptModule里,所以脚本化后找不到compute方法。第二个错误:
'Tensor (inferred)' object has no attribute or method 'linear'
直接脚本化实例的compute方法时,TorchScript无法正确捕获self的类型信息——它不知道self是LibraryModule的实例,自然无法解析self.linear这个属性。
解决方案:创建包装子类显式导出方法
因为咱们不能修改第三方库的代码,最稳妥的办法是继承原模块,在子类中显式标记要导出的compute方法:
import torch import torch.nn as nn # ---------------------- 第三方库代码(不可修改) ---------------------- class SomeClass: """A utility class in a library I cannot modify""" def __init__(self, x): self.x = x class LibraryModule(nn.Module): """A module provided in a library I cannot modify""" def __init__(self, in_features, out_features): super().__init__() self.linear = nn.Linear(in_features, out_features) def compute(self, x, some_class_object: SomeClass): """ Main function of my module; like forward, but takes a non-tensor argument """ return self.linear(x) * some_class_object.x # -------------------------------------------------------------------- # 我们的包装子类,用于导出compute方法 class WrappedLibraryModule(LibraryModule): @torch.jit.export def compute(self, x, some_class_object: SomeClass): # 直接复用父类的compute逻辑,不做任何修改 return super().compute(x, some_class_object) # 脚本化包装后的模块 scripted_module = torch.jit.script(WrappedLibraryModule(3, 2)) # 测试调用 test_tensor = torch.tensor([10, 20, 30]) test_obj = SomeClass(2) result = scripted_module.compute(test_tensor, test_obj) print(result)
关键说明
@torch.jit.export是核心:它告诉TorchScript要把这个方法纳入编译范围,这样脚本化后的模块就会保留compute方法- 无需额外处理
SomeClass:因为原方法里有明确的类型注解some_class_object: SomeClass,且这个类结构简单,TorchScript可以自动识别并处理它的类型
备注:内容来源于stack exchange,提问作者user118967
相关产品推荐
相关产品推荐

