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

使用torch.where后模型参数梯度为零?求排查原因

问题描述

模型forward方法代码

def forward(self, x, output_type, *unused_args, **unused_kwargs):
    gru_output, gru_hn = self.gru(x)
    # Decoder (Graph Adjacency Reconstruction)
    for data_batch_idx in range(x.shape[0]):
        pred = self.decoder(gru_output[data_batch_idx, -1, :])  # gru_output[-1] => only take last time-step
        pred_graph_adj = pred.reshape(1, -1) if data_batch_idx == 0 else torch.cat((pred_graph_adj, pred.reshape(1, -1)), dim=0)
    if output_type == "discretize":
        bins = torch.tensor(self.model_cfg['output_bins']).reshape(-1, 1)
        num_bins = len(bins)-1
        bins = torch.concat((bins[:-1], bins[1:]), dim=1)
        discretize_values = np.linspace(2, 4, num_bins)
        for lower, upper, discretize_value in zip(bins[:, 0], bins[:, 1], discretize_values):
            pred_graph_adj = torch.where((pred_graph_adj <= upper) & (pred_graph_adj > lower), discretize_value, pred_graph_adj)
        pred_graph_adj = torch.where(pred_graph_adj < bins.min(), 2, pred_graph_adj)

    return pred_graph_adj

训练代码片段

pred = self.forward(x, output_type=self.model_cfg['output_type'])
batch_loss = self.loss_fn(pred, y)
self.optimizer.zero_grad()
batch_loss.backward()
self.optimizer.step()
self.scheduler.step()

观察结果

  • 当output_type不为"discretize"(不使用torch.where)时,sum([p.grad.sum() for p in self.decoder.parameters()])结果非零;但当output_type为"discretize"(使用torch.where)时,该结果为零。
  • 已确认batch_loss不为零。
  • 已检查模型所有权重的require_grad均为True。
  • 已确认计算图中pred和batch_loss与模型权重连通。

补充信息

  • pred_graph_adj初始值范围为-1 ~ 1,经过torch.where处理后最终值范围为2 ~ 4。
  • torch.where的参数对应关系:
    (lower, upper] -> discrete_values:
    (-1, -0.25] -> 2
    (-0.25, 0.25] -> 3
    (0.25, 1] -> 4
    

疑问

  1. 使用torch.where是否会导致模型参数梯度变为零?
  2. 若torch.where不是原因,还有哪些可能的因素?

解答

1. torch.where是否会导致梯度为零?

会,这就是梯度消失的直接原因。你的代码里discretize_values是用np.linspace生成的numpy数组,传入torch.where时会被转换成无梯度的常数张量。当torch.where把模型原始输出替换成这些常数后,对应位置的梯度传递会被完全截断——常数对模型参数的导数为0,最终导致解码器参数的梯度总和为0。

哪怕你用PyTorch张量存储离散值,只要没设置requires_grad=True,同样会阻断梯度传递。你当前的实现方式必然会导致梯度无法回传到解码器参数。

2. 其他可能的因素(排除torch.where的情况)

如果后续修改torch.where的使用方式后仍出现梯度为零的情况,可排查以下几点:

  • 损失函数特性:如果使用的损失函数在离散输出点上导数为0(比如某些自定义损失),会导致梯度无法传递。不过你非离散化时梯度正常,这个可能性较低。
  • 输出与标签的极端匹配:如果离散化后的输出和标签完全一致,损失的导数会为0,但你已经确认batch_loss不为零,可直接排除。
  • 计算图意外中断:比如在离散化前,pred_graph_adj被无意调用了.detach()、.numpy()等阻断梯度的操作,但你已确认计算图连通,这个可能性也很低。
  • 数值精度溢出:极端情况下,离散化操作导致数值溢出,梯度被置为0,但你的输出范围变化是合理的,这种情况几乎不可能出现。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 10:04:52