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

PyTorch Seq2Seq模型RNN报错:输入需3维但实际为2维问题排查

问题原因

你定义输入字段时错误设置了sequential = False,这是报错的根本原因。
torchtext的Field参数中,sequential=False代表该字段是非序列类型的离散/连续值,此时你设置的tokenize、init_token、eos_token、fix_length等序列相关的参数全部不会生效。该模式下每个样本的输入会被处理为单个整数,单个批次的输入张量形状为(batch_size, ),经过Embedding层后输出形状为(batch_size, emb_dim),是二维张量,不符合LSTM要求的三维输入格式,因此触发报错。

解决方案

你添加<sos>、<eos>标记的操作本身就说明输入是需要按序列处理的文本,只需要把inp_fields的sequential参数修改为True即可解决问题:

inp_fields = Field(sequential = True, use_vocab = True, tokenize = tokenize_text, lower = False, fix_length=100, init_token="<sos>", eos_token="<eos>")

修改后,批次输入的形状会变为(序列长度, 批量大小),经过Embedding层后输出形状为(序列长度, 批量大小, 嵌入维度),刚好匹配LSTM默认要求的三维输入格式。

缺失维度说明

你之前的输入缺少的就是序列长度维度,sequential=False的配置让框架把整个输入当作单个token处理,自动消去了序列长度的维度。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 14:15:02