EditNTS无教师强制训练时torch.cat张量维度不匹配报错求助
问题原因
PyTorch中RNN模块默认输入输出维度顺序为[序列长度, 批次大小, 特征维度],你当前需要拼接的其他张量使用的是[批次大小, 序列长度, 特征维度]的批次优先格式,因此出现维度不匹配报错。
修正方案
方案1:直接调整张量维度(改动最小,优先选择)
在执行torch.cat操作时,对hidden_words[0]做维度交换即可,将原来的拼接代码修改为:
output_t = torch.cat((output_edits, attn_applied_org_t, c, hidden_words[0].transpose(0,1)), dim=2)
或者使用permute方法:
output_t = torch.cat((output_edits, attn_applied_org_t, c, hidden_words[0].permute(1,0,2)), dim=2)
两种方法都可以把[1, 32, 400]的张量调整为你需要的[32, 1, 400]格式,且不会影响梯度反向传播,训练可以正常运行。
方案2:调整RNN初始化参数
如果你的全流程都使用批次优先的维度格式,可以在self.rnn_words初始化时添加batch_first=True参数,示例如下:
self.rnn_words = nn.GRU(input_size=xxx, hidden_size=xxx, num_layers=xxx, batch_first=True)
修改后RNN的输入和输出都会自动适配[批次大小, 序列长度, 特征维度]的格式,不需要再单独调整hidden的维度。注意修改后要保证输入给rnn_words的embedded_words张量维度也符合批次优先的格式。
验证步骤
修改完成后可以先打印所有待拼接张量的前两个维度,确认除了拼接的dim=2之外,其他维度大小完全一致,就可以正常运行。
内容的提问来源于stack exchange,提问作者jaugustin12
相关产品推荐
相关产品推荐

