自定义nn.Linear反向函数与torch.autograd训练结果差异原因排查
我想要自定义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

