关于PyTorch Seq2Seq教程解码器中两种输入分支张量形状的疑问
嘿,这个问题抓得特别准!其实你不用担心,两种分支输出的decoder_input形状是完全一致的,都是**(batch_size, 1)**的2D张量,刚好适配后续模型的输入要求。咱们一步步拆解来看:
第一种情况:有teacher forcing(target_tensor不为空)
decoder_input = target_tensor[:, i].unsqueeze(1)
这里的target_tensor本身是形状为(batch_size, MAX_LENGTH)的2D张量——每一行对应一个样本的完整目标序列。当我们取target_tensor[:, i]时,得到的是一个1D张量(batch_size,),代表当前step所有样本的目标输入词。
而unsqueeze(1)的作用就是在第1个维度(也就是列维度)上新增一个维度,把它从1D变成(batch_size, 1)的2D张量,刚好匹配模型后续forward_step对输入的形状要求。
第二种情况:无teacher forcing(target_tensor为空)
_, topi = decoder_output.topk(1) decoder_input = topi.squeeze(-1).detach()
先看decoder_output的形状:它是forward_step输出的结果,形状为(batch_size, 1, output_size)(因为我们每次只输入一个词,seq_len=1)。用topk(1)在最后一个维度(词表维度)取概率最大的1个词,得到的topi是形状为(batch_size, 1, 1)的3D张量——最后一个维度是取top1的索引位置。
接下来topi.squeeze(-1)是去掉最后一个多余的维度,把3D张量压缩成(batch_size, 1)的2D张量,和第一种情况的形状完全一致!detach()只是断开梯度传播,避免把预测值的梯度回传,不影响张量形状。
所以两种分支最终的decoder_input都是(batch_size, 1),完全适配后续embedding层和GRU的输入要求(因为GRU设置了batch_first=True,需要输入(batch_size, seq_len, hidden_size),而embedding层会把(batch_size,1)的输入转换成(batch_size,1,hidden_size))。
备注:内容来源于stack exchange,提问作者RoomTemperature

