TorchScript在tch-rs中运行维度越界RuntimeError的调试求助
问题解决:TorchScript模型在tch-rs中运行维度错误
错误原因
报错Dimension out of range (expected to be in range of [-1, 0], but got 1),核心问题有两点:
- 传入BERT的
input_id张量维度不符合要求:BERT要求输入为二维张量([batch_size, sequence_length]),但Rust中传入的是一维张量,导致模型内部尝试获取第1维度(序列长度)时失败。 jit.trace的局限性:它仅基于给定输入记录计算路径,无法处理BERT内部的动态维度逻辑,若后续输入形状变化就会触发错误。
解决步骤
1. 修正Rust中输入张量的维度
在tch-rs中构造input_id和mask时,必须保证是二维张量。示例代码如下:
use tch::{Tensor, Kind}; // 构造符合要求的input_id:[批量大小, 序列长度],值范围对应BERT词表大小 let input_id = Tensor::randint(0, 28996, &[2, 128], (Kind::Int64, tch::Device::Cpu)); // 构造同形状的attention_mask:全1表示有效token let mask = Tensor::ones(&[2, 128], (Kind::Int64, tch::Device::Cpu));
2. 改用torch.jit.script导出模型
jit.script直接解析Python代码生成TorchScript,能更好兼容动态逻辑,替换jit.trace即可解决大部分兼容性问题:
import torch from your_model_file import BertClassifier # 加载训练完成的模型权重 model = BertClassifier() model.load_state_dict(torch.load("trained_weights.pt")) model.eval() # 用script导出模型 scripted_model = torch.jit.script(model) scripted_model.save("bert_classifier_scripted.pt")
3. 验证导出模型的有效性
在Python中用和Rust相同形状的输入测试导出后的模型,确保运行正常:
test_input_id = torch.randint(0, 28996, (2, 128)) test_mask = torch.ones((2, 128)) output = scripted_model(test_input_id, test_mask) print(output.shape) # 预期输出:torch.Size([2, 5])
4. 若坚持使用jit.trace
必须保证trace时的输入形状和Rust中使用的完全一致,后续不能动态调整批量大小或序列长度:
# 用与Rust一致的形状生成示例输入 example_input_id = torch.randint(0, 28996, (2, 128)) example_mask = torch.ones((2, 128)) # 基于固定形状trace模型 traced_model = torch.jit.trace(model, (example_input_id, example_mask)) traced_model.save("bert_classifier_traced.pt")
内容的提问来源于stack exchange,提问作者pekc
相关产品推荐
相关产品推荐

