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

PyTorch训练循环中如何获取中间输出的真实梯度?

问题原因与解决方法

为什么a_s.grad和b_s.grad返回None

PyTorch默认仅为叶子节点(比如模型的可训练参数)保存梯度,对于计算过程中的中间张量,除非在其创建后立即调用retain_grad()标记要保留梯度,否则反向传播时不会存储梯度信息。你在模型外部对返回的a_s、b_s调用retain_grad()的时机太晚,此时计算图已构建完成,标记无法生效,因此梯度会是None。

解决方法

方法一:在模型内部标记保留梯度

修改模型的forward方法,在堆叠得到a_s、b_s后立即调用retain_grad():

class SmallModel(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.lstm_cell = torch.nn.LSTMCell(input_size=features.shape[2], hidden_size=16)
        self.fc = torch.nn.Linear(in_features=16, out_features=targets.shape[2])
        
    def forward(self, features):
        hx = torch.randn(64, 16)
        cx = torch.randn(64, 16)
        
        a_s = []
        b_s = []
        c_s = []
        
        for t in range(num_time_steps):
            features_t = features[:, t, :]
            hx, cx = self.lstm_cell(features_t, (hx, cx))
            out_t = torch.relu(self.fc(hx))
            
            a = out_t * 0.8 + 20
            b = a * 2
            c = b * 0.9
            
            a_s.append(a)
            b_s.append(b)
            c_s.append(c)
            
        a_s = torch.stack(a_s, dim=1)
        b_s = torch.stack(b_s, dim=1)
        c_s = torch.stack(c_s, dim=1)
        
        # 标记保留梯度
        a_s.retain_grad()
        b_s.retain_grad()
        c_s.retain_grad()
        
        return a_s, b_s, c_s

之后训练循环中无需再调用retain_grad(),反向传播后即可直接获取梯度:

for epoch in range(n_epoch):
    optimizer.zero_grad()
    a_s, b_s, c_s = model(features)
    loss = loss_fn(c_s, targets)
    loss.backward()
    
    # 现在可以打印实际梯度
    print(a_s.grad)
    print(b_s.grad)
    
    optimizer.step()

方法二:使用torch.autograd.grad直接计算梯度

如果不想修改模型,可以在训练循环中使用torch.autograd.grad直接计算损失对a_s、b_s的梯度,注意需要保留计算图直到所有梯度计算完成:

for epoch in range(n_epoch):
    optimizer.zero_grad()
    a_s, b_s, c_s = model(features)
    loss = loss_fn(c_s, targets)
    
    # 反向传播时保留计算图
    loss.backward(retain_graph=True)
    
    # 计算损失对a_s的梯度
    grad_a = torch.autograd.grad(loss, a_s, retain_graph=True)[0]
    # 计算损失对b_s的梯度
    grad_b = torch.autograd.grad(loss, b_s)[0]
    
    print(grad_a)
    print(grad_b)
    
    optimizer.step()

注意:retain_graph=True会占用更多内存,若仅需计算梯度而不需要更新参数,也可以不调用loss.backward(),直接用torch.autograd.grad计算。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 02:10:15