You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为已构建的神经网络模型添加自定义属性并调用?

解决神经网络模型中自定义编码方法跨阶段调用问题

问题分析

你已构建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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.27 08:33:36