使用JIT保存含自定义类的PyTorch模型时遇未知类型注解错误
我碰到过好几次类似的PyTorch JIT自定义类注解问题,结合你已经做的尝试,给你几个具体的排查方向和解决方案:
1. 检查类型注解的导入路径一致性
PyTorch JIT对模块导入路径的一致性非常挑剔——哪怕是同一个类,如果在Generator的代码里和你的脚本化主脚本里用了不同的导入方式,JIT就会把它们当成不同的类型。比如:
- 如果
Generator里用了相对导入from .blocks import SpadeBlock,但主脚本用的是绝对导入from gan.blocks import SpadeBlock,JIT就会识别失败。 - 解决方法:统一所有涉及
SpadeBlock的类型注解为绝对路径,比如在Generator类里把类型注解写成gan.blocks.SpadeBlock,而不是仅仅SpadeBlock。
2. 单独测试SpadeBlock的JIT兼容性
虽然你确认了它继承自nn.Module,但类内部的一些实现可能违反了JIT的规则:
- 先试试单独脚本化
SpadeBlock:
如果这一步也报错,说明问题出在from gan.blocks import SpadeBlock # 传入你实际使用的参数初始化一个实例 test_block = SpadeBlock(in_channels=64, out_channels=64) jit_block = torch.jit.script(test_block)SpadeBlock自身:- 检查
__init__方法是否用了JIT不支持的语法(比如*args/**kwargs、动态赋值未声明的属性); - 给所有类属性加上明确的类型注解,比如
self.conv: nn.Conv2d = nn.Conv2d(...),而不是只写self.conv = nn.Conv2d(...)。
- 检查
3. 确认包的导入上下文正确
你用了pip install -e安装开发包,但有时候运行环境可能导入了其他位置的旧版本gan包:
- 在主脚本开头加一段代码验证:
确保输出的路径是你的开发目录,而不是系统的site-packages或者其他地方。import gan print(gan.__file__) - 调整导入顺序:先导入
SpadeBlock再导入Generator,有时候JIT对导入顺序敏感:import torch from gan.blocks import SpadeBlock from gan import Generator
4. 显式注册自定义类给JIT
如果上面的方法都没用,可以强制让JIT提前识别SpadeBlock的类型:
import torch from gan.blocks import SpadeBlock from gan import Generator # 先显式脚本化SpadeBlock类,让JIT把它加入类型系统 torch.jit.script(SpadeBlock) # 再继续模型加载和脚本化流程 generator = Generator() generator.load_weights("path/to/weigts") jitted = torch.jit.script(generator) torch.jit.save(jitted, "model.pt")
另外还有个小技巧:如果Generator里的SpadeBlock属性用了类型注解,比如self.spade: SpadeBlock = ...,可以把注解改成字符串形式self.spade: "SpadeBlock" = ...,这在导入路径复杂或者有循环依赖的场景里经常能解决识别问题。
内容的提问来源于stack exchange,提问作者Arthur
相关产品推荐
相关产品推荐

