PyTorch RNN代码片段为何返回0梯度?求技术解析
RNN输出切片后输入梯度为0的原因分析
我正在研究对RNN输出进行切片操作如何影响梯度,编写了如下测试代码:
# Test the gradient of a RNN import torch from torch import nn from torch.autograd import Variable rnn = nn.RNN(10, 20, 1) input = Variable(torch.randn(5, 3, 10), requires_grad=True) h0 = Variable(torch.randn(1, 3, 20), requires_grad=True) output, hn = rnn(input, h0) output = output[:, 1:, :] loss_fn = nn.CrossEntropyLoss() target = Variable(torch.empty((5,2), dtype=torch.long).random_(20), requires_grad=False) loss = loss_fn(output.reshape(-1, output.shape[2]), target.view(-1)) loss.backward() print(input.grad[:,0,:])
我原本认为梯度不应为0,因为即使基于切片后的输出计算损失,后续时间步的隐藏状态仍会受第一步输入的影响。为何此时梯度为0?
核心原因:切片操作针对的是batch维度,而非时间步维度
你的代码里的切片output[:, 1:, :]是对batch维度(第二个维度)进行切片,保留了batch中索引为1、2的样本,完全丢弃了索引为0的样本的所有输出。而PyTorch中RNN的batch维度是独立计算的——每个batch样本的前向传播、隐藏状态更新都是相互独立的,不同样本之间没有依赖关系。
这意味着:
- 损失计算只用到了batch中第1、2个样本的输出,和第0个样本的输出完全无关
- 第0个样本的输入没有参与任何损失的计算,因此其梯度自然为0
你误以为“后续时间步隐藏状态受第一步输入影响”是对同一个样本的时间步维度来说的,比如如果切片是针对时间步的output[1:, :, :](去掉第一个时间步的输出),那第一个时间步的输入确实会影响后续时间步的隐藏状态,梯度不会为0。但这里的切片是针对batch维度,直接丢掉了整个样本的输出,所以该样本的输入梯度必然为0。
内容的提问来源于stack exchange,提问作者Riccardo Ricci
相关产品推荐
相关产品推荐

