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

自定义nn.Linear反向函数与torch.autograd训练结果差异原因排查

自定义nn.Linear反向传播后损失偏高的问题

我想要自定义nn.Linear()的反向传播函数,实现代码如下:

class Linear(torch.autograd.Function):
    @staticmethod
    def forward(ctx, inputs, weight, bias):
        e = F.linear(inputs, weight, bias)
        ctx.save_for_backward(inputs, e, weight)
        return e

    @staticmethod
    def backward(ctx, grad):
        input, e, weight = ctx.saved_tensors
        grad_weight = grad.permute(0, 2, 1).matmul(input)    
        grad_bias = grad.sum(dim=(0))
        grad_input_t = grad.matmul(weight)   

        def check_numerical_stability(tensor, name):
            if torch.isnan(tensor).any() or torch.isinf(tensor).any():
                print(f"Warning: {name} contains NaN or Inf values!")
                tensor = torch.where(torch.isnan(tensor), torch.zeros_like(tensor), tensor)
                tensor = torch.where(torch.isinf(tensor), torch.zeros_like(tensor), tensor)
            return tensor

        grad_input_t = check_numerical_stability(grad_input_t, "grad_input_t")
        grad_weight = check_numerical_stability(grad_weight, "grad_weight")
        grad_bias = check_numerical_stability(grad_bias, "grad_bias")
        return grad_input_t, grad_weight, grad_bias

但使用该函数训练后的损失,比使用torch.autograd自动微分训练的损失更高。输入为3D张量[batch_size, len_seq, embed_size],我已检查微分公式、处理数值稳定性问题并添加权重初始化,仍未找到原因,恳请帮助。


问题分析与修复

你的自定义反向传播中梯度计算的维度处理存在两处关键错误,这是导致损失偏高的核心原因:

  • grad_weight维度错误:
    原代码计算出的grad_weight是[batch_size, out_features, in_features],但标准nn.Linear的权重梯度应该是[out_features, in_features],需要在batch和序列维度上求和。正确计算方式是先对grad转置后两维,与input做批量矩阵乘法,再求和:

    grad_weight = grad.transpose(1, 2).bmm(input).sum(dim=0)
    
  • grad_bias求和维度错误:
    输入是3D张量,grad的形状为[batch_size, len_seq, out_features],标准nn.Linear的偏置梯度需要在batch和len_seq两个维度上求和,原代码仅在dim=0求和,漏掉了序列维度,正确写法:

    grad_bias = grad.sum(dim=(0, 1))
    
  • grad_input_t计算正确:
    grad.matmul(weight)得到的[batch_size, len_seq, in_features]与官方反向传播的输入梯度一致,无需修改。

修复后的完整backward函数:

@staticmethod
def backward(ctx, grad):
    input, e, weight = ctx.saved_tensors
    # 修正权重梯度计算
    grad_weight = grad.transpose(1, 2).bmm(input).sum(dim=0)
    # 修正偏置梯度求和维度
    grad_bias = grad.sum(dim=(0, 1))
    grad_input_t = grad.matmul(weight)   

    def check_numerical_stability(tensor, name):
        if torch.isnan(tensor).any() or torch.isinf(tensor).any():
            print(f"Warning: {name} contains NaN or Inf values!")
            tensor = torch.where(torch.isnan(tensor), torch.zeros_like(tensor), tensor)
            tensor = torch.where(torch.isinf(tensor), torch.zeros_like(tensor), tensor)
        return tensor

    grad_input_t = check_numerical_stability(grad_input_t, "grad_input_t")
    grad_weight = check_numerical_stability(grad_weight, "grad_weight")
    grad_bias = check_numerical_stability(grad_bias, "grad_bias")
    return grad_input_t, grad_weight, grad_bias

另外,建议用torch.autograd.gradcheck工具验证自定义实现的正确性:对比自定义函数与官方nn.Linear的梯度,确保误差在可接受范围内,这能快速定位梯度计算的偏差。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 08:48:12