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
相关产品推荐
相关产品推荐

