如何用脚本转换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
相关产品推荐
相关产品推荐

