深度学习训练中自定义Loss、Metric与Callback的PyTorch实践与坑点
先说我一次真实翻车经历。当时训练一个二分类模型训练集的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()间接影响当前batchMetric训练/验证循环内通常累积不需要可导只读不参与更新整个epoch累积Callbackepoch/step级钩子不参与可以EMA、恢复权重、改LR跨epoch/stepLoss的使命是给优化器一个标量它的梯度决定权重怎么更新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, param1.0): super().__init__() self.param param def forward(self, pred, target): # pred: [B, C] 或 [B, ...] # target: [B] 或 [B, ...] loss ... return loss.mean() # 注意必须是一个标量几个关键点forward必须返回标量。常用reductionmean因为它不随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_gradTrue) 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甚至更小时权重接近1loss保留完整。class FocalLoss(nn.Module): def __init__(self, gamma2.0, alphaNone): super().__init__() self.gamma gamma self.alpha alpha def forward(self, logits, targets): ce F.cross_entropy(logits, targets, reductionnone) 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越大易分样本被压得越狠但也会让难分样本的梯度偏高训练后期容易出现震荡。我习惯先从gamma2.0开始调如果发现训练初期loss掉得太猛后续验证集反而不涨就把gamma降到1.0或1.5。2.3 案例二Asymmetric Loss多标签分类的负样本压制Asymmetric LossASL是针对多标签分类设计的思路是对正负样本分别使用不同的指数γ公式如下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_pos0.0, gamma_neg4.0, clip0.05, eps1e-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_lossNone, aux_weight0.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_lossaux_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、resultPyTorch Lightning里也沿用了这套设计。它不是拍脑袋定的而是这种模式最贴合“跨batch累积”的需求。3.2 从零写一个F1 Score重点在累积器而不是公式F1是分类任务最常见的自定义指标。它的难点不在公式——公式谁都背得出来——而在累积状态的设计。我推荐用tp/fp/fn三个累积器而不是每batch算一个F1再平均后者只有在每个batch类别分布完全一致时才近似正确。class F1Score: def __init__(self, num_classes, averagemacro): 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(dim1).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, dtypetorch.long) def update(self, preds, targets): preds preds.argmax(dim1).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, minlengthself.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(dim0) self.cm.sum(dim1) - inter valid union 0 iou inter[valid].float() / union[valid].clamp(min1).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(gamma2.0) train_metric F1Score(num_classes10) val_metric F1Score(num_classes10) cbs CallbackList([EarlyStopping(...), ModelCheckpoint(...), EMA(model, decay0.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, monitorval_f1, patience5, min_delta1e-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, monitorval_f1, modemax): 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权重滑动平均CallbackEMA指数移动平均是我个人非常喜欢的一种“免费午餐”每次参数更新后用shadow decay * shadow (1 - decay) * weights维护一份影子权重推理时把影子权重临时载入通常比直接用最后一步权重泛化好。class EMA(CallbackBase): def __init__(self, model, decay0.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, alpha1.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或float32loss.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 1F1应该严格等于0而不是NaNCallback里的checkpoint确实在每个epoch结束时写入并且可以正常torch.load回来。我项目里一直留着一个test_custom_components.py每次改完训练代码先跑一遍它确认没问题才放全量训练。它已经救了我好几次最典型的一次是改了一个Metric的聚合方式结果宏平均和微平均都出现了NaN冒烟测试直接拦住了。这个脚本不复杂但值得长期维护因为训练代码越改越复杂自定义组件出错的可能性只会越来越高。