如何为已构建的神经网络模型添加自定义属性并调用?
解决神经网络模型中自定义编码方法跨阶段调用问题
问题分析
你已构建ImageClassifier模型,需要集成encode_image和encode_text两个自定义数据预处理方法,要求这两个方法既能在模型定义阶段调用,也能在训练完成(包括JIT编译后)正常使用。但直接将方法写入模型类中会导致JIT编译后调用时出现“函数未定义”错误——这是因为PyTorch JIT默认仅追踪forward方法中用到的代码逻辑,未被调用的自定义方法不会被编译到模型中。
解决方案
方案一:用torch.jit.script_method显式注册方法
通过给自定义编码方法添加@torch.jit.script_method装饰器,强制JIT编译这些方法,确保编译后仍可正常调用。
修改后的完整代码:
import torch import torch.nn as nn # 假设qlayer和ClassicalLayer已提前定义 qlayer = ... ClassicalLayer = ... class ImageClassifier(nn.Module): def __init__(self, n_qubits, n_layers): super().__init__() # 初始化模型依赖的组件(示例,需根据实际场景补充) self.visual = nn.Linear(224*224, 512) # 示例视觉编码器 self.token_embedding = nn.Embedding(1000, 512) self.positional_embedding = nn.Parameter(torch.randn(1, 100, 512)) self.transformer = nn.TransformerEncoder(nn.TransformerEncoderLayer(512, 8), 6) self.ln_final = nn.LayerNorm(512) self.text_projection = nn.Parameter(torch.randn(512, 512)) self.dtype = torch.float32 self.model = nn.Sequential( qlayer, ClassicalLayer(2) ) @torch.jit.script_method def encode_image(self, image): return self.visual(image.type(self.dtype)) @torch.jit.script_method def encode_text(self, text): x = self.token_embedding(text).type(self.dtype) # [batch_size, n_ctx, d_model] x = x + self.positional_embedding.type(self.dtype) x = x.permute(1, 0, 2) # NLD -> LND x = self.transformer(x) x = x.permute(1, 0, 2) # LND -> NLD x = self.ln_final(x).type(self.dtype) # 取eot token对应的特征(eot_token是每个序列中的最大值) x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection return x def forward(self, x): result = self.model(x) return result
方案二:将编码逻辑封装为独立子模块
把encode_image和encode_text的逻辑拆分为独立的nn.Module子类,作为ImageClassifier的属性。JIT会自动追踪所有子模块的代码,确保编译后可正常调用。
示例代码:
import torch import torch.nn as nn qlayer = ... ClassicalLayer = ... class ImageEncoder(nn.Module): def __init__(self, input_dim, output_dim, dtype=torch.float32): super().__init__() self.visual = nn.Linear(input_dim, output_dim) self.dtype = dtype def forward(self, image): return self.visual(image.type(self.dtype)) class TextEncoder(nn.Module): def __init__(self, vocab_size, d_model, n_ctx, n_layers, dtype=torch.float32): super().__init__() self.token_embedding = nn.Embedding(vocab_size, d_model) self.positional_embedding = nn.Parameter(torch.randn(1, n_ctx, d_model)) self.transformer = nn.TransformerEncoder(nn.TransformerEncoderLayer(d_model, 8), n_layers) self.ln_final = nn.LayerNorm(d_model) self.text_projection = nn.Parameter(torch.randn(d_model, d_model)) self.dtype = dtype def forward(self, text): x = self.token_embedding(text).type(self.dtype) x = x + self.positional_embedding.type(self.dtype) x = x.permute(1, 0, 2) x = self.transformer(x) x = x.permute(1, 0, 2) x = self.ln_final(x).type(self.dtype) x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection return x class ImageClassifier(nn.Module): def __init__(self, n_qubits, n_layers): super().__init__() # 初始化编码器子模块 self.image_encoder = ImageEncoder(224*224, 512) self.text_encoder = TextEncoder(1000, 512, 100, 6) self.model = nn.Sequential( qlayer, ClassicalLayer(2) ) # 封装调用接口,保持原有调用习惯 def encode_image(self, image): return self.image_encoder(image) def encode_text(self, text): return self.text_encoder(text) def forward(self, x): result = self.model(x) return result
验证方法
编译模型后测试编码方法是否可用:
# 初始化模型 model = ImageClassifier(n_qubits=2, n_layers=2) # JIT编译模型 jit_model = torch.jit.script(model) # 生成测试数据 test_image = torch.randn(2, 224*224) test_text = torch.randint(0, 1000, (2, 100)) # 原始模型调用编码方法 img_feat = model.encode_image(test_image) txt_feat = model.encode_text(test_text) # JIT编译后模型调用编码方法 jit_img_feat = jit_model.encode_image(test_image) jit_txt_feat = jit_model.encode_text(test_text) # 验证结果一致性 print(torch.allclose(img_feat, jit_img_feat)) print(torch.allclose(txt_feat, jit_txt_feat))
内容的提问来源于stack exchange,提问作者Hrridoy V2
相关产品推荐
相关产品推荐

