使用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
疑问
- 使用
torch.where是否会导致模型参数梯度变为零? - 若
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
相关产品推荐
相关产品推荐

