☰
深度学习训练中自定义Loss、Metric与Callback的PyTorch实践与坑点
2026/10/3 7:50:08 网站建设 项目流程

先说我一次真实翻车经历。当时训练一个二分类模型,训练集的loss掉得很漂亮,验证loss也在降,曲线怎么看我都很满意。可我把验证结果捞下来用 sklearn 手算 F1,发现连续几个epoch都在原地踏步。排查了一晚上,问题出在我自己写的自定义Metric上——对预测概率又做了一次softmax再取argmax,等于套了两遍argmax,评估口径整个废掉。

这种问题不会报错,不会闪红,训练速度不受影响,它只会安静地让你的验证指标失真。这也是我这次想写自定义Loss、Metric及Callback的原因。这三样东西在每个深度学习框架里都有“官方写法”,但官方文档通常只告诉你语法,很少告诉你“为什么这么写”,以及“写错了会出现什么鬼畜现象”。这篇文章默认你已经有跑通baseline的能力,想开始按自己的方式改造训练流程——无论是换损失函数、换评估指标,还是在每个epoch结束时做点额外动作。我会以PyTorch语法为主,思路在Keras和TensorFlow里同样适用。

1. 先分清楚三者的分工,再动手写,否则所有“魔改”都会互相打架

1.1 一张表理清Loss、Metric和Callback的真实边界

很多人习惯用“Metric就是不用反向传播的Loss”来理解评估指标,新手阶段这么记没问题,但真去写复杂项目时会吃亏。这三个组件在训练脚本里的职责完全不同,我把边界整理成了一张表:

组件调用时机是否参与反向传播能不能影响模型权重计算范围
Loss训练循环内,每个batch必须可导通过.backward()间接影响当前batch
Metric训练/验证循环内,通常累积不需要可导只读,不参与更新整个epoch累积
Callbackepoch/step级钩子不参与可以(EMA、恢复权重、改LR)跨epoch/step

Loss的使命是给优化器一个标量,它的梯度决定权重怎么更新;Metric的使命是给人类一个可解释的数字,它的计算方式甚至可以和损失函数完全不同;Callback则是训练流程里的“监工”,在batch结束、epoch结束这些时间点插入一段额外代码,读取日志、保存模型、调整学习率、直接改模型参数都可以。

1.2 为什么“反正都能算一个数”会把你带进沟里

我见过不少把三者混用的代码,典型的有三类问题。

第一类是把Accuracy或F1直接当Loss用。F1对输入的微小变化基本不可导,或者导数直接为0,模型根本学不动。这不是“换个损失函数”能救的,正确做法是给不可导指标找一个可导的代理目标,比如用Focal Loss或Dice Loss的平滑版本去逼近你要优化的方向。

第二类是把验证集的Loss当作业务评估指标。Loss通常包含正则项、标签平滑、自定义权重,它和真实业务指标(准确率、IOU、召回率)不是单调强相关。我见过一个项目在Loss里加了很大的权重衰减,训练loss一路降,验证集的recall反而在掉,最后发现是正则项把有效信息也压掉了。

第三类是把“每个epoch要做的杂事”全部写死在训练循环里,不抽象成Callback。第一次跑实验没问题,但当你第二份实验要换学习率策略、换保存逻辑、加一个EMA时,就得把训练循环从头翻一遍。抽象成Callback不是为了显得工程化,而是把变化点收敛在可控范围,减少改一处动全身的风险。

2. 自定义Loss的核心骨架:先保住梯度,再谈公式

2.1 一个标准的自定义Loss类长什么样

PyTorch里自定义Loss基本就是继承nn.Module,实现一个forward,返回标量。骨架长这样:

import torch import torch.nn as nn import torch.nn.functional as F class MyLoss(nn.Module): def __init__(self, param=1.0): super().__init__() self.param = param def forward(self, pred, target): # pred: [B, C] 或 [B, ...] # target: [B] 或 [B, ...] loss = ... return loss.mean() # 注意必须是一个标量

几个关键点:

  • forward必须返回标量。常用reduction='mean',因为它不随batch size变化;sum容易让loss绝对数值跟着batch size走,换batch size时学习率也得跟着调;none用于按样本加权,但最后必须自己做一次聚合,否则backward会直接报错。
  • __init__里可以放超参数,比如Focal Loss的gamma、Asymmetric Loss的正负样本指数,但不要放需要更新的状态,那属于Metric或Optimizer的职责。

真正难的是别在forward里写“断掉梯度”的操作。你可以在里面用clamp、where、max这些有数学意义的tensor操作,但随手把tensor转成numpy再算距离,那瞬间计算图就断了。表现很迷惑:loss还在下降,但模型参数几乎不动,或者loss卡在某个值附近震荡,没有任何报错。

提示:想快速验证自定义Loss有没有问题,构造一个固定batch,跑一遍loss.backward(),检查pred.grad是否非空且全是有限值。如果梯度为零或NaN,基本可以断定forward里混进了不可导操作。

pred = torch.randn(4, 10, requires_grad=True) target = torch.randint(0, 10, (4,)) loss = MyLoss()(pred, target) loss.backward() assert pred.grad is not None assert torch.isfinite(pred.grad).all()

这个方法我每次写完新Loss都会跑一遍,比肉眼检查代码可靠得多。

2.2 案例一:Focal Loss,类别不平衡时的首选

Focal Loss出自目标检测,核心想法是压低“易分样本”的loss贡献,把训练注意力引向“难分样本”。公式是:

FL = -(1 - p_t)^γ * log(p_t)

其中p_t是模型对真实类别的预测概率。当p_t接近1时,(1-p_t)^γ接近0,这个样本的loss被压得很低;当p_t接近0.5甚至更小时,权重接近1,loss保留完整。

class FocalLoss(nn.Module): def __init__(self, gamma=2.0, alpha=None): super().__init__() self.gamma = gamma self.alpha = alpha def forward(self, logits, targets): ce = F.cross_entropy(logits, targets, reduction='none') pt = torch.exp(-ce) # 交叉熵 -log(pt),反推 pt focal = (1.0 - pt) ** self.gamma * ce if self.alpha is not None: alpha_t = self.alpha[targets] # 每类的权重 focal = alpha_t * focal return focal.mean()

这里用torch.exp(-ce)反推pt,比自己算softmax再取索引更简洁,也能保证数值一致性。alpha可以是标量也可以是一维tensor,需要和类别数对齐;类别特别多时可以用1 / 类别频率初始化。

一个容易忽略的细节:gamma越大,易分样本被压得越狠,但也会让难分样本的梯度偏高,训练后期容易出现震荡。我习惯先从gamma=2.0开始调,如果发现训练初期loss掉得太猛,后续验证集反而不涨,就把gamma降到1.0或1.5。

2.3 案例二:Asymmetric Loss,多标签分类的负样本压制

Asymmetric Loss(ASL)是针对多标签分类设计的,思路是对正负样本分别使用不同的指数γ,公式如下:

L = -y · (1-p)^γ⁺ · log(p) - (1-y) · p^γ⁻ · log(1-p)

y是0/1标签,p是sigmoid概率。多标签场景下负样本通常远多于正样本,而且很多负样本“太容易学”,如果把它们的loss权重降下来,模型就能把容量留给更重要的正样本和难分负样本。

class AsymmetricLoss(nn.Module): def __init__(self, gamma_pos=0.0, gamma_neg=4.0, clip=0.05, eps=1e-8): super().__init__() self.gamma_pos = gamma_pos self.gamma_neg = gamma_neg self.clip = clip self.eps = eps def forward(self, logits, targets): prob = torch.sigmoid(logits) prob = torch.clamp(prob, self.eps, 1.0 - self.eps) # 防止 log(0) pos_loss = -targets * (1 - prob) ** self.gamma_pos * torch.log(prob) neg_loss = -(1 - targets) * (prob ** self.gamma_neg) * torch.log(1 - prob) return (pos_loss + neg_loss).mean()

γ⁺通常设为0或很小的值,γ⁻设成2到4,因为负样本“太好学”,需要压得狠一点。clip参数可以控制预测概率的裁剪范围,减少噪声标签对负样本的影响。注意如果标签是-1/1而不是0/1,需要先转换。

2.4 案例三:Intermediate Loss,把中间层也拉进优化目标

Intermediate Loss也叫中间层损失或辅助损失。最早出名是在GoogLeNet里,当时的想法是网络很深容易梯度消失,不如在中间某个特征图上加一个辅助分类器,让梯度能提前回传。后来语义分割里的Deep Supervision也走这个思路。

要写这种Loss,模型侧需要把中间层输出也返回出来:

class MultiHeadModel(nn.Module): def __init__(self, backbone, num_classes): super().__init__() self.backbone = backbone self.head = nn.Linear(backbone.out_features, num_classes) self.aux_head = nn.Linear(backbone.out_features, num_classes) def forward(self, x): feats = self.backbone(x) main = self.head(feats) aux = self.aux_head(feats) # 也可以是更浅的特征 return {'main': main, 'aux': aux}

Loss侧把主损失和辅助损失加权相加:

class IntermediateLoss(nn.Module): def __init__(self, primary_loss, aux_loss=None, aux_weight=0.4): super().__init__() self.primary = primary_loss self.aux = aux_loss if aux_loss is not None else F.cross_entropy self.aux_weight = aux_weight def forward(self, preds, target): main_loss = self.primary(preds['main'], target) aux_loss = self.aux(preds['aux'], target) return main_loss + self.aux_weight * aux_loss

aux_weight我一般取0.3到0.5。训练后期可以让辅助损失的权重逐渐衰减,让主分支的优先级慢慢提高,否则辅助头可能会把主干特征引导到“既能分类又能辅助”的妥协状态,反而影响主任务上限。

注意:如果模型在训练时返回dict,在验证推理时也要保持同样结构,否则preds['main']会直接KeyError。这个错误在训练循环里可能因为异常被埋在日志中,排查时容易忽略。

3. 自定义Metric的价值不在“算得准”,而在跨batch状态管理

3.1 Metric和Loss的本质区别

Loss是“每batch即时计算、立即消费”的数据,算完就扔;Metric则是“跨batch积累、epoch结束时统一结算”的数据。这带来两个实际问题:

  • 单个batch的统计量方差很大,尤其验证集被切分成很多小块时,最后几个batch的指标不能代表整个epoch。
  • 不同batch的类别分布可能不一样,微平均和宏平均的计算结果差异会被放大。如果只取最后一个batch的指标,你看到的往往是噪声最大的那个数字。

所以几乎每个框架的Metric生命周期都是reset -> update -> compute三段式。Keras里是reset_state、update_state、result,PyTorch Lightning里也沿用了这套设计。它不是拍脑袋定的,而是这种模式最贴合“跨batch累积”的需求。

3.2 从零写一个F1 Score,重点在累积器而不是公式

F1是分类任务最常见的自定义指标。它的难点不在公式——公式谁都背得出来——而在累积状态的设计。我推荐用tp/fp/fn三个累积器,而不是每batch算一个F1再平均,后者只有在每个batch类别分布完全一致时才近似正确。

class F1Score: def __init__(self, num_classes, average='macro'): self.num_classes = num_classes self.average = average self.reset() def reset(self): self.tp = torch.zeros(self.num_classes) self.fp = torch.zeros(self.num_classes) self.fn = torch.zeros(self.num_classes) def update(self, preds, targets): preds = preds.argmax(dim=1).view(-1) targets = targets.view(-1) for c in range(self.num_classes): p_mask = (preds == c) t_mask = (targets == c) self.tp[c] += (p_mask & t_mask).sum() self.fp[c] += (p_mask & ~t_mask).sum() self.fn[c] += (~p_mask & t_mask).sum() def compute(self): eps = 1e-12 precision = self.tp / (self.tp + self.fp + eps) recall = self.tp / (self.tp + self.fn + eps) f1 = 2 * precision * recall / (precision + recall + eps) if self.average == 'macro': return f1.mean().item() if self.average == 'micro': tp = self.tp.sum() fp = self.fp.sum() fn = self.fn.sum() precision = tp / (tp + fp + eps) recall = tp / (tp + fn + eps) return 2 * precision * recall / (precision + recall + eps) raise ValueError(self.average)

几个实操细节:

  • update里提前做argmax,这是和训练阶段的口径对齐。训练时模型输出的是logits,验证时如果直接拿logits算F1,必须先决定用argmax还是阈值。
  • view(-1)是为了兼容图像/序列任务的输出形状。分类任务直接( B, C ),分割任务可能是( B, C, H, W ),压平后再统计不会错。
  • eps加在每个分母上,防止某些类别在整个验证集里一个真值都没有时产生NaN。宏平均遇到这种情况,应该跳过该类别还是给它一个0分?我选择给0分,并在日志里标记这个类别样本不足。

3.3 IoU Metric:用混淆矩阵累积,比逐类求IoU更快更稳

分割任务里最常自定义的是IoU。朴素的写法是每batch逐类算inter / union再求平均,但batch小的时候数值波动大,而且处理“某个类在当前batch没出现”时很麻烦。更好的做法是累积一个混淆矩阵,最后统一计算:

class IoU: def __init__(self, num_classes): self.num_classes = num_classes self.reset() def reset(self): self.cm = torch.zeros(self.num_classes, self.num_classes, dtype=torch.long) def update(self, preds, targets): preds = preds.argmax(dim=1).view(-1) targets = targets.view(-1) keep = (targets >= 0) & (targets < self.num_classes) & (preds < self.num_classes) preds = preds[keep] targets = targets[keep] idx = targets * self.num_classes + preds counts = torch.bincount(idx, minlength=self.num_classes * self.num_classes) self.cm += counts.view(self.num_classes, self.num_classes) def compute(self): inter = torch.diag(self.cm) union = self.cm.sum(dim=0) + self.cm.sum(dim=1) - inter valid = union > 0 iou = inter[valid].float() / union[valid].clamp(min=1).float() return iou.mean().item()

混淆矩阵的好处是累积对象是整数矩阵,天然适合后面要讲的分布式all_reduce,而且计算复杂度不会随着类别数上升而变成灾难,因为用了torch.bincount一次搞定,而不是双重循环。

3.4 分布式训练下的Metric口径:本地统计再平均是错的

单卡训练时状态管理很简单,多卡DDP下就容易出现“指标对不上”的玄学问题。根本原因是每个GPU只看到自己shard的数据。如果每个rank本地算F1再取平均,会受数据切分影响。比如类别A的样本恰好集中在rank 0的shard里,rank 1在类别A上的统计几乎全是0,本地F1就会被严重拉低。

标准解法是把tp/fp/fn这些累积量先all_reduce,再做最终计算:

import torch.distributed as dist def sync_tensor(t): if dist.is_initialized(): dist.all_reduce(t)

在DDP训练时,每个Metric累积器update完,等epoch结束先同步再compute。单机多卡也好,多机多卡也好,这个套路都成立。就算你当前只在单卡训练,我也建议把Metric设计成“统计量累积”的形式,将来扩到多卡时不用推倒重来。

4. 自定义Callback的本质:在训练循环的指定时刻插入行为

4.1 Keras和PyTorch里的Callback流派差异

Keras很早就把Callback设计成一套完整钩子体系:on_train_begin、on_epoch_begin、on_batch_end、on_train_end等等。PyTorch原生没有这个统一抽象,最常见的做法是把训练循环写成普通Python函数,然后自己维护一个callbacks列表在关键位置调用。PyTorch Lightning则把Callback做成了正式接口,封装了更多细节。

我的建议很简单:如果项目已经从零起步、想长期维护,直接用Lightning会更省心,因为它的EarlyStopping和ModelCheckpoint已经写得很成熟;如果是在已有PyTorch脚本上做修改,自己实现一个Callback列表大概只需要二十行代码,改动最小,也更透明。

最小实现长这样:

class CallbackBase: def on_train_begin(self, model, optimizer, **kwargs): pass def on_batch_end(self, model, optimizer, batch_idx, logs): pass def on_epoch_end(self, model, optimizer, epoch, logs): pass class CallbackList: def __init__(self, callbacks): self.callbacks = callbacks def __getattr__(self, name): def call(*args, **kwargs): for cb in self.callbacks: getattr(cb, name)(*args, **kwargs) return call

训练循环里只需要写cbs.on_epoch_end(model, optimizer, epoch, logs),剩下的事由每个callback自己决定。

4.2 一个最小可用的训练循环生命周期骨架

把前面的自定义Loss和自定义Metric串进一个完整循环:

model = model.to(device) criterion = FocalLoss(gamma=2.0) train_metric = F1Score(num_classes=10) val_metric = F1Score(num_classes=10) cbs = CallbackList([EarlyStopping(...), ModelCheckpoint(...), EMA(model, decay=0.999)]) def train_one_epoch(): model.train() train_metric.reset() for x, y in train_loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() pred = model(x) loss = criterion(pred, y) loss.backward() optimizer.step() train_metric.update(pred, y) cbs.on_batch_end(model, optimizer, len(train_loader), {'loss': loss.item()}) return {'train_loss': loss.item(), 'train_f1': train_metric.compute()} def validate(): model.eval() val_metric.reset() with torch.no_grad(): for x, y in val_loader: x, y = x.to(device), y.to(device) pred = model(x) val_metric.update(pred, y) return {'val_f1': val_metric.compute()} for epoch in range(epochs): train_logs = train_one_epoch() val_logs = validate() logs = {**train_logs, **val_logs} cbs.on_epoch_end(model, optimizer, epoch, logs)

注意train_metric.reset()必须放在每个epoch开头,否则第3轮的F1会包含前两轮的数据,曲线看起来会异常平缓,失去评估意义。这个错误很隐蔽,因为指标不会报错,只是“钝化”。

4.3 EarlyStopping和ModelCheckpoint:两个必须分别实现的回调

EarlyStopping的核心是patience和best值跟踪,而不是“看着差不多了就停”。一个稳定的写法:

class EarlyStopping(CallbackBase): def __init__(self, monitor='val_f1', patience=5, min_delta=1e-4): self.monitor = monitor self.patience = patience self.min_delta = min_delta self.best = -float('inf') self.counter = 0 self.should_stop = False def on_epoch_end(self, model, optimizer, epoch, logs): score = logs.get(self.monitor) if score is None: raise ValueError(f'{self.monitor} not found in logs') if score > self.best + self.min_delta: self.best = score self.counter = 0 else: self.counter += 1 if self.counter >= self.patience: self.should_stop = True

不要用raise StopIteration去中断循环,除非你在最外层做了异常捕获。否则训练日志会缺尾巴,模型状态也可能处于半更新状态。设置should_stop标记,在epoch结束后统一检查,更安全。

ModelCheckpoint负责“把好状态留下来”:

class ModelCheckpoint(CallbackBase): def __init__(self, filepath, monitor='val_f1', mode='max'): self.filepath = filepath self.monitor = monitor self.mode = mode self.best = -float('inf') if mode == 'max' else float('inf') def on_epoch_end(self, model, optimizer, epoch, logs): score = logs[self.monitor] improved = (self.mode == 'max' and score > self.best) or \ (self.mode == 'min' and score < self.best) if improved: self.best = score state = { 'model': model.state_dict(), 'optimizer': optimizer.state_dict() if optimizer else None, 'epoch': epoch, 'score': score, } torch.save(state, self.filepath)

这里单独把EarlyStopping和ModelCheckpoint分成两个类,是因为它们关注点不同:一个决定“什么时候停”,一个决定“存哪一份”。合成一个类虽然也能跑,但后续你要调整保存频率、要改成每隔N个epoch保存一次,就得牵扯早停逻辑,没必要。

4.4 写一个EMA权重滑动平均Callback

EMA(指数移动平均)是我个人非常喜欢的一种“免费午餐”:每次参数更新后,用shadow = decay * shadow + (1 - decay) * weights维护一份影子权重,推理时把影子权重临时载入,通常比直接用最后一步权重泛化好。

class EMA(CallbackBase): def __init__(self, model, decay=0.999): self.decay = decay self.shadow = {k: v.detach().clone() for k, v in model.state_dict().items()} def on_batch_end(self, model, **kwargs): with torch.no_grad(): for k, v in model.state_dict().items(): s = self.shadow[k] s.mul_(self.decay).add_(v, alpha=1.0 - self.decay) def swap(self, model): model.load_state_dict(self.shadow)

这里有个细节:state_dict里除了可学习参数,还有BN的running_mean、running_var这些buffer。EMA会把它们也平滑掉,这在验证时可能会让BN的统计量滞后。如果你的模型BN比较重,可以只对“含权重”的key做EMA,把num_batches_tracked这类状态排除。

提示:用EMA做推理前,一定先把shadow覆盖回模型,再重新跑一遍完整验证集,确认指标没下降再提测。EMA权重在训练过程中的指标不能直接代表最终推理性能,因为它还没被完整地“安置”回模型里。

5. 三个自定义组件放进同一个脚本时,我踩过并修复的坑

5.1 train/eval模式与Metric记录错位的坑

最经典的问题是验证循环忘了切model.eval(),导致BN和Dropout一直处于训练模式。现象是验证Loss正常,但验证F1忽高忽低,尤其是batch size小、Dropout比例高的时候,曲线像在蹦极。另一个低阶问题是验证循环里依然调用了optimizer.zero_grad(),虽然没有大影响,但会让代码看起来在“训练”,很容易误导后来的人。

正确姿势是:验证循环前写model.eval(),并且整个验证过程包在torch.no_grad()里。Metric的update不需要梯度,所以放在no_grad下完全没问题。

5.2 设备、浮点精度与除零问题

这类问题不报错,只会悄悄污染你的指标。我给你一张排查表:

问题现象对策
model在GPU、Metric累积器在CPU偶尔报device mismatch,或隐式同步极慢统一.to(device),或者让Metric累积器留在CPU时先.cpu()再累加
把GPU tensor直接存进Python list显存不释放,训练越跑越卡在update阶段用.item()或转成CPU标量
某些类别0个真值F1/IoU出现NaN分母加eps,或者union为0时跳过该类
混合精度训练下累积器用了fp16指标精度漂移,tp/fp/fn逐渐变成0累积器用torch.long或float32

loss.item()和metric.compute()混用的时候尤其小心:loss计算图在backward之后会释放,但如果你在loss.backward()之后还留着loss tensor,内存不会立刻回收。习惯是print完就.item(),Metric里不要存logits本身,只存统计量。

5.3 给钩子代码加“幂等保护”和日志降噪

Callback里最容易翻车的是频繁文件写入。有人把checkpoint写在on_batch_end里,结果跑一天磁盘被写满,而且训练速度被IO拖慢一大截。我建议:所有文件操作只放在on_epoch_end,并且判断“是否更优”后再写。还有一点是要用“严格大于”还是“大于等于”的判断。我踩过一次用>=导致前两个epoch连续各存了一份,因为第二个epoch和第一个epoch分数恰好相同,结果磁盘瞬间爆了。

日志也是同理。batch循环内不要print太多。平均每100个batch打一次还能接受,但最优雅的做法是把所有结构化日志集中到on_epoch_end统一格式化。这样跑长训练时不会被刷屏,而且不同实验之间的日志格式也能保持一致,方便后面画图对比。

5.4 用最小回归集验证自定义组件

不管代码写得再小心,我强烈建议在做全量训练前先跑一个“冒烟测试”。做法很简单:拿一个小数据集,只跑2到3个epoch,重点确认三件事:

  • 自定义Loss的loss在下降,且pred.grad非零;
  • 自定义Metric在构造的极端标签分布下,和sklearn手算结果一致。比如让模型全员预测class 0、标签全是class 1,F1应该严格等于0,而不是NaN;
  • Callback里的checkpoint确实在每个epoch结束时写入,并且可以正常torch.load回来。

我项目里一直留着一个test_custom_components.py,每次改完训练代码先跑一遍它,确认没问题才放全量训练。它已经救了我好几次,最典型的一次是改了一个Metric的聚合方式,结果宏平均和微平均都出现了NaN,冒烟测试直接拦住了。这个脚本不复杂,但值得长期维护,因为训练代码越改越复杂,自定义组件出错的可能性只会越来越高。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询