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

TorchScript在tch-rs中运行维度越界RuntimeError的调试求助

问题解决:TorchScript模型在tch-rs中运行维度错误

错误原因

报错Dimension out of range (expected to be in range of [-1, 0], but got 1),核心问题有两点:

  1. 传入BERT的input_id张量维度不符合要求:BERT要求输入为二维张量([batch_size, sequence_length]),但Rust中传入的是一维张量,导致模型内部尝试获取第1维度(序列长度)时失败。
  2. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 07:15:32