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

如何用脚本转换PyTorch Transformer并解决其C++运行编译报错?

解决PyTorch Transformer无法用TorchScript编译的问题

这个坑我之前踩过!你遇到的问题本质是TorchScript对可选Tensor参数的类型推断有严格限制——原TransformerEncoder的forward方法里,src_mask等参数默认值为None,TorchScript无法自动推断出这些参数对应的Tensor类型(比如数据类型、设备信息),因此抛出了RuntimeError。下面给你几个可靠的解决方案:

方案1:封装Wrapper类显式处理默认参数

最稳妥的方式是自己写一个包装类,把原TransformerEncoder包起来,在Wrapper的forward方法里显式声明参数类型,或者为可选参数生成合法的默认Tensor,让TorchScript能正确识别类型。

示例代码:

import torch
from torch.nn import TransformerEncoder, TransformerEncoderLayer

class TransformerEncoderWrapper(torch.nn.Module):
    def __init__(self, d_model, nhead, dim_feedforward, dropout, num_layers):
        super().__init__()
        encoder_layers = TransformerEncoderLayer(d_model, nhead, dim_feedforward, dropout)
        self.transformer = TransformerEncoder(encoder_layers, num_layers)
    
    def forward(
        self, 
        src: torch.Tensor, 
        src_mask: torch.Tensor = None, 
        src_key_padding_mask: torch.Tensor = None
    ):
        # 为可选参数生成默认的空mask(根据实际需求调整)
        if src_mask is None:
            # 生成与输入序列长度匹配的全False mask(不遮挡任何位置)
            src_mask = torch.zeros((src.size(0), src.size(0)), device=src.device, dtype=torch.bool)
        if src_key_padding_mask is None:
            src_key_padding_mask = torch.zeros((src.size(1), src.size(0)), device=src.device, dtype=torch.bool)
        
        return self.transformer(src, src_mask=src_mask, src_key_padding_mask=src_key_padding_mask)

# 实例化Wrapper并编译
wrapper = TransformerEncoderWrapper(1000, 8, 512, 0.1, 6)
sm = torch.jit.script(wrapper)

# 测试编译后的模型
test_src = torch.randn(10, 32, 1000)  # 形状:(序列长度, batch大小, 特征维度)
output = sm(test_src)
print(output.shape)  # 应该输出 torch.Size([10, 32, 1000])

# 保存模型供C++加载
sm.save("transformer_encoder.pt")

方案2:使用torch.jit.trace替代script

如果你的使用场景中,输入的形状和参数模式比较固定,可以用torch.jit.trace来编译模型。trace会根据你提供的示例输入,记录模型的计算路径,避免类型推断问题。

示例代码:

import torch
from torch.nn import TransformerEncoder, TransformerEncoderLayer

encoder_layers = TransformerEncoderLayer(1000, 8, 512, 0.1)
transf = TransformerEncoder(encoder_layers, 6)

# 准备一个符合实际使用场景的示例输入
example_src = torch.randn(10, 32, 1000)
# 用trace编译,传入示例输入
sm = torch.jit.trace(transf, (example_src,))

# 保存模型
sm.save("transformer_encoder_trace.pt")

⚠️ 注意:trace的局限性在于只能捕获示例输入走过的计算路径,如果后续使用时需要传入src_mask等示例中没用到的参数,可能会出现错误。

方案3:升级PyTorch版本

某些旧版本的PyTorch(比如1.8及更早)对TorchScript处理可选Tensor参数的支持不完善,升级到1.9+的稳定版后,这个问题可能已经被官方修复。你可以尝试更新PyTorch:

pip install --upgrade torch

C++加载模型示例

编译好模型后,你可以用PyTorch C++ API加载并运行:

#include <torch/script.h>
#include <iostream>

int main() {
    torch::jit::script::Module module;
    try {
        // 加载编译好的.pt模型文件
        module = torch::jit::load("transformer_encoder.pt");
    } catch (const c10::Error& e) {
        std::cerr << "加载模型失败:" << e.what() << std::endl;
        return -1;
    }

    // 准备输入Tensor(形状要和Python中测试的一致)
    torch::Tensor src = torch::randn({10, 32, 1000});
    std::vector<torch::jit::IValue> inputs = {src};
    
    // 运行模型
    torch::Tensor output = module.forward(inputs).toTensor();
    std::cout << "输出形状:" << output.sizes() << std::endl;

    return 0;
}

内容的提问来源于stack exchange,提问作者Andrey

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.07 21:02:39