必须继承 nn.Module,不能只写纯函数;因训练循环依赖 loss.backward(),仅 nn.Module 实例返回的 Tensor 才能保留计算图,普通函数易致梯度静默失败或 RuntimeError。

PyTorch自定义损失函数:该继承 nn.Module 还是写纯函数?
直接说结论:**必须继承 nn.Module(或其子类 nn.Loss)**,不能只写一个普通 Python 函数。PyTorch 的训练循环(如 optimizer.step())依赖 loss.backward() 触发自动求导,而只有 nn.Module 实例返回的 Tensor 才能正确保留计算图——普通函数若未显式确保 requires_grad=True 且参与图构建,反向传播会静默失败或报 RuntimeError: element 0 of tensors does not require grad。
继承 nn.Module 实现损失类的最小必要结构
核心是重写 forward 方法,并确保输入 pred 和 target 都参与计算图。不需要重写 __init__(除非要传超参),也不用调用 super().__init__() ——但建议加上,便于未来扩展。
常见错误:
- 在
forward中用了numpy或.item()/.detach().cpu().numpy(),导致计算图断裂 - 返回标量时用了
torch.mean(loss).item(),返回的是 Python float,不是Tensor - 对
target做了非可导操作(如torch.argmax后再算交叉熵),应改用原生支持的nn.CrossEntropyLoss或手动实现 soft 版本
正确示例(带权重的二分类 focal loss):
import torch
import torch.nn as nn
<p>class FocalLoss(nn.Module):
def <strong>init</strong>(self, alpha=1.0, gamma=2.0):
super().<strong>init</strong>()
self.alpha = alpha
self.gamma = gamma</p><pre class="brush:php;toolbar:false;">def forward(self, inputs, targets):
# inputs: [N, 2], targets: [N] with class indices (0 or 1)
logpt = nn.functional.log_softmax(inputs, dim=1)
pt = torch.exp(logpt)
logpt = logpt.gather(1, targets.unsqueeze(1))
pt = pt.gather(1, targets.unsqueeze(1))
focal_weight = ((1 - pt) ** self.gamma)
loss = -self.alpha * focal_weight * logpt
return loss.mean()为什么不用 nn.Loss 基类?
PyTorch 并没有名为 nn.Loss 的抽象基类。官方所有损失函数(如 nn.MSELoss、nn.CrossEntropyLoss)都直接继承自 nn.Module。所谓“继承 Loss 类”是误传——你看到的文档里写的 “Base class for all losses” 实际指向的是内部的 _Loss(带下划线,非公开),用户不应继承它。继承 nn.Module 是唯一受支持、稳定且可维护的方式。
使用场景差异:
- 需要保存超参(如
gamma、reduction)并在多个地方复用 → 必须用类 - 一次性实验、快速验证公式 → 可写闭包函数,但必须返回带梯度的
Tensor,且不能脱离Module上下文直接喂给optimizer - 想兼容
reduction='none'/'mean'/'sum'→ 类中需显式支持,否则默认只能做mean
性能与调试关键点:backward 无声失败的排查路径
自定义损失最常遇到的问题不是报错,而是 loss 不下降、梯度为 0 或 NaN。根本原因往往藏在张量形状和梯度流里。
- 检查
loss.grad_fn是否为None:如果是,说明计算图已断,往前查哪一步用了.detach()、.cpu()或 numpy - 确认
pred来自模型输出(即model(x)),且模型参数requires_grad=True(默认就是) - 避免在 loss 计算中引入非连续张量(如转置后未
.contiguous()),某些操作(如view)会报错,但有些不会,只悄悄破坏梯度 - 打印
loss.item()和pred[0].grad(反向后)对比验证是否真有梯度回传
复杂点永远在细节:哪怕公式完全正确,一个 .detach() 就能让整个训练失效,而且不报错。写完务必用小 batch 跑一两步,用 torch.autograd.gradcheck 对关键子函数做数值梯度校验(尤其涉及自定义运算时)。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











