Ultralytics:解读 YOLO26 知识蒸馏
-
- [1. 核心文件与启用方式](#1. 核心文件与启用方式)
- [2. 总体架构(mermaid 图解)](#2. 总体架构(mermaid 图解))
- [3. 初始化阶段:DistillationModel 构建](#3. 初始化阶段:DistillationModel 构建)
-
- [3.1 加载并冻结教师](#3.1 加载并冻结教师)
- [3.2 自动检测蒸馏特征层](#3.2 自动检测蒸馏特征层)
- [3.3 注册特征采集 hook](#3.3 注册特征采集 hook)
- [3.4 用 dummy 前向确定通道数,构建投影器](#3.4 用 dummy 前向确定通道数,构建投影器)
- [4. 训练阶段:单个 batch 的一次迭代](#4. 训练阶段:单个 batch 的一次迭代)
-
- [4.1 准备 batch](#4.1 准备 batch)
- [4.2 `loss()` 主流程骨架(`distill_model.py` L201-242)](#4.2
loss()主流程骨架(distill_model.pyL201-242)) - [4.3 常规检测损失(`v8DetectionLoss`,`loss.py` L478-482)](#4.3 常规检测损失(
v8DetectionLoss,loss.pyL478-482))
- [5. 蒸馏损失详解(核心)](#5. 蒸馏损失详解(核心))
-
- [5.1 蒸馏损失在整个 loss 中的位置](#5.1 蒸馏损失在整个 loss 中的位置)
- [5.2 蒸馏损失计算的完整代码(`loss()` L215-242)](#5.2 蒸馏损失计算的完整代码(
loss()L215-242)) - [5.3 特征与 score 的尺寸追踪(2 张 640×640)](#5.3 特征与 score 的尺寸追踪(2 张 640×640))
- [5.4 `loss_sl2` 逐行拆解(L244-264)](#5.4
loss_sl2逐行拆解(L244-264)) - [5.5 数值实例(完整演算)](#5.5 数值实例(完整演算))
-
- [5.5.1 微型示例(P3 层,标注通道=2,位置=2)](#5.5.1 微型示例(P3 层,标注通道=2,位置=2))
- [5.5.2 反向示例(误差在低置信度位置)](#5.5.2 反向示例(误差在低置信度位置))
- [5.5.3 完整 batch 尺度(2 张图,3 个 neck 层)](#5.5.3 完整 batch 尺度(2 张图,3 个 neck 层))
- [5.6 梯度流向](#5.6 梯度流向)
- [6. 训练日志中的 dis_loss](#6. 训练日志中的 dis_loss)
- [7. 推理阶段](#7. 推理阶段)
- [8. Checkpoint 保存 / 恢复 / EMA 处理](#8. Checkpoint 保存 / 恢复 / EMA 处理)
-
- [8.1 保存时剥离教师(`torch_utils.py` L796-804)](#8.1 保存时剥离教师(
torch_utils.pyL796-804)) - [8.2 EMA 也剥离教师(`torch_utils.py` L723-727)](#8.2 EMA 也剥离教师(
torch_utils.pyL723-727)) - [8.3 恢复训练(`trainer.py` L777-789)](#8.3 恢复训练(
trainer.pyL777-789)) - [8.4 NaN 恢复(`trainer.py` L1006-1011)](#8.4 NaN 恢复(
trainer.pyL1006-1011))
- [8.1 保存时剥离教师(`torch_utils.py` L796-804)](#8.1 保存时剥离教师(
- [9. 关键设计要点与性能参考](#9. 关键设计要点与性能参考)
-
- [9.1 设计要点](#9.1 设计要点)
- [9.2 当前 YOLO26 蒸馏 mAP 参考(官方文档)](#9.2 当前 YOLO26 蒸馏 mAP 参考(官方文档))
- 参考

- 由于本人水平有限,难免出现错漏,敬请批评改正。
- 更多精彩内容,可点击进入我的个人主页查看
整合自两份源码精读:YOLO26 知识蒸馏(训练与推理全过程) + YOLO26 蒸馏损失逐行精讲 。基于 ultralytics 源码:
ultralytics/nn/distill_model.py、ultralytics/engine/trainer.py、ultralytics/utils/loss.py、ultralytics/nn/modules/head.py、ultralytics/nn/tasks.py、ultralytics/utils/torch_utils.py。全程以一个具体设定为例:2 张 640×640 RGB 图像 组成一个 batch,学生模型为
yolo26n.pt,教师模型为yolo26s.pt,COCO 80 类,dis=6.0。
1. 核心文件与启用方式
| 文件 | 作用 |
|---|---|
ultralytics/nn/distill_model.py |
蒸馏核心:DistillationModel 包装 教师+学生,特征提取、投影器、蒸馏损失 |
ultralytics/engine/trainer.py |
训练器:组装 DistillationModel、损失加权、EMA、checkpoint |
ultralytics/utils/loss.py |
v8DetectionLoss:常规检测损失(box / cls / dfl) |
ultralytics/nn/modules/head.py |
Detect 头:one2many / one2one 双分支输出 |
ultralytics/utils/torch_utils.py |
保存时剥离教师、EMA 剥离教师 |
启用方式 (distill_model 参数 + 默认权重 dis=6.0):
bash
yolo detect train model=yolo26n.pt data=coco8.yaml epochs=100 distill_model=yolo26s.pt
python
from ultralytics import YOLO
student = YOLO("yolo26n.pt")
student.train(data="coco8.yaml", epochs=100, distill_model="yolo26s.pt", dis=6.0)
配置项定义在 ultralytics/cfg/default.yaml:
yaml
distill_model: # (str, optional) path to teacher model for knowledge distillation
dis: 6.0 # (float) distillation loss weight
2. 总体架构(mermaid 图解)
#mermaid-svg-oqVM7Du18r0dKBIF{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-oqVM7Du18r0dKBIF .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-oqVM7Du18r0dKBIF .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-oqVM7Du18r0dKBIF .error-icon{fill:#552222;}#mermaid-svg-oqVM7Du18r0dKBIF .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-oqVM7Du18r0dKBIF .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-oqVM7Du18r0dKBIF .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-oqVM7Du18r0dKBIF .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-oqVM7Du18r0dKBIF .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-oqVM7Du18r0dKBIF .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-oqVM7Du18r0dKBIF .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-oqVM7Du18r0dKBIF .marker{fill:#333333;stroke:#333333;}#mermaid-svg-oqVM7Du18r0dKBIF .marker.cross{stroke:#333333;}#mermaid-svg-oqVM7Du18r0dKBIF svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-oqVM7Du18r0dKBIF p{margin:0;}#mermaid-svg-oqVM7Du18r0dKBIF .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-oqVM7Du18r0dKBIF .cluster-label text{fill:#333;}#mermaid-svg-oqVM7Du18r0dKBIF .cluster-label span{color:#333;}#mermaid-svg-oqVM7Du18r0dKBIF .cluster-label span p{background-color:transparent;}#mermaid-svg-oqVM7Du18r0dKBIF .label text,#mermaid-svg-oqVM7Du18r0dKBIF span{fill:#333;color:#333;}#mermaid-svg-oqVM7Du18r0dKBIF .node rect,#mermaid-svg-oqVM7Du18r0dKBIF .node circle,#mermaid-svg-oqVM7Du18r0dKBIF .node ellipse,#mermaid-svg-oqVM7Du18r0dKBIF .node polygon,#mermaid-svg-oqVM7Du18r0dKBIF .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-oqVM7Du18r0dKBIF .rough-node .label text,#mermaid-svg-oqVM7Du18r0dKBIF .node .label text,#mermaid-svg-oqVM7Du18r0dKBIF .image-shape .label,#mermaid-svg-oqVM7Du18r0dKBIF .icon-shape .label{text-anchor:middle;}#mermaid-svg-oqVM7Du18r0dKBIF .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-oqVM7Du18r0dKBIF .rough-node .label,#mermaid-svg-oqVM7Du18r0dKBIF .node .label,#mermaid-svg-oqVM7Du18r0dKBIF .image-shape .label,#mermaid-svg-oqVM7Du18r0dKBIF .icon-shape .label{text-align:center;}#mermaid-svg-oqVM7Du18r0dKBIF .node.clickable{cursor:pointer;}#mermaid-svg-oqVM7Du18r0dKBIF .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-oqVM7Du18r0dKBIF .arrowheadPath{fill:#333333;}#mermaid-svg-oqVM7Du18r0dKBIF .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-oqVM7Du18r0dKBIF .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-oqVM7Du18r0dKBIF .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-oqVM7Du18r0dKBIF .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-oqVM7Du18r0dKBIF .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-oqVM7Du18r0dKBIF .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-oqVM7Du18r0dKBIF .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-oqVM7Du18r0dKBIF .cluster text{fill:#333;}#mermaid-svg-oqVM7Du18r0dKBIF .cluster span{color:#333;}#mermaid-svg-oqVM7Du18r0dKBIF div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-oqVM7Du18r0dKBIF .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-oqVM7Du18r0dKBIF rect.text{fill:none;stroke-width:0;}#mermaid-svg-oqVM7Du18r0dKBIF .icon-shape,#mermaid-svg-oqVM7Du18r0dKBIF .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-oqVM7Du18r0dKBIF .icon-shape p,#mermaid-svg-oqVM7Du18r0dKBIF .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-oqVM7Du18r0dKBIF .icon-shape .label rect,#mermaid-svg-oqVM7Du18r0dKBIF .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-oqVM7Du18r0dKBIF .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-oqVM7Du18r0dKBIF .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-oqVM7Du18r0dKBIF :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} dataloader (batch = 2×640×640×3)
forward hooks 采集特征
forward hooks 采集特征
projectori 1×1 卷积对齐通道
teacher_scores 作为权重
× dis=6.0
backward, 只更新学生+投影器
batch = {img, cls, bboxes, batch_idx}
teacher_model yolo26s (frozen, eval, no_grad)
student_model yolo26n (trainable, train mode)
教师 neck 特征 128,256,512 + 头 scores
学生 neck 特征 64,128,256
对齐后学生特征 128,256,512
score-weighted L2 蒸馏损失 loss_sl2
蒸馏损失 distill
Detect 头 one2many + one2one
常规损失 box + cls + dfl (v8DetectionLoss)
总损失 = (box+cls+dfl) + distill
optimizer (学生参数 + projector 参数)
一句话总结 :推同一个 batch 进冻结的教师 和可训练的学生 ,用 hook 抓取两者在三个 neck 特征层(P3/P4/P5)的输出;投影器把学生通道数对齐到教师;用教师分类置信度加权的 L2 距离作为蒸馏损失,与常规检测损失相加后反向传播------只更新学生和投影器,教师完全不更新。
3. 初始化阶段:DistillationModel 构建
训练器在 trainer.py 的 _setup_train 中把普通模型替换成蒸馏包装(L352-353):
python
if self.args.distill_model is not None and not isinstance(unwrap_model(self.model), DistillationModel):
self.model = DistillationModel(student_model=self.model, teacher_model=self.args.distill_model)
DistillationModel.__init__(distill_model.py L62-114)做如下几件事。
3.1 加载并冻结教师
python
ch = student_model.yaml.get("channels", 3) # 学生输入通道 = 3
if isinstance(teacher_model, (str, Path)):
teacher_model = load_checkpoint(teacher_model)[0] # 从 yolo26s.pt 加载
if teacher_model.yaml.get("channels", 3) != ch: # 若教师输入通道不同则重建
...
device = next(student_model.parameters()).device
self.teacher_model = teacher_model.to(device)
self._freeze_teacher() # 教师 eval() + 所有参数的 requires_grad=False
_freeze_teacher(L180-187):
python
self.teacher_model.eval()
for v in self.teacher_model.parameters():
v.requires_grad = False
教师被永久冻结在 eval 模式,不参与反向传播,也不更新。
3.2 自动检测蒸馏特征层
get_distill_layers(L168-178)扫描模型,找到 Detect 头,返回「喂给头的输入层索引 + 头本身的索引」:
python
for m in model.model:
if isinstance(m, Detect):
return [*list(m.f), m.i] # m.f = 头的输入来源层, m.i = 头所在层
对 YOLO26(yolo26.yaml 中 Detect 在层 52,其 from 是 [16, 19, 22]):
feats_idx = [16, 19, 22, 23]
feats_idx[:-1] = [16, 19, 22]→ 三个 neck 特征层(P3/P4/P5),蒸馏的特征就在这里取。feats_idx[-1] = 23→ Detect 头,用来生成教师分类置信度(teacher_scores)。
3.3 注册特征采集 hook
_register_feature_hooks(L154-166)在教师和学生相同索引 的层上挂 FeatureHook,把前向输出存进共享 dict:
python
self._student_hooks.append(
self.student_model.model[idx].register_forward_hook(FeatureHook(self._student_feats, idx))
)
self._teacher_hooks.append(
self.teacher_model.model[idx].register_forward_hook(FeatureHook(self._teacher_feats, idx))
)
FeatureHook(L17-30)非常简单:
python
def __call__(self, module, inputs, output):
self.feat_dict[self.idx] = output # 前向一次,把该层输出存进 dict
3.4 用 dummy 前向确定通道数,构建投影器
因为要逐层做蒸馏,需要知道每层特征通道数。用一个 (2, 3, 640, 640) 的全零张量分别前向教师和学生,hook 会自动抓取特征:
python
imgsz = student_model.args.imgsz # 640
student_model.eval()
with torch.no_grad():
im = torch.zeros(2, ch, imgsz, imgsz, device=device) # (2,3,640,640)
teacher_model(im)
student_model(im)
student_model.train()
得到教师/学生的 neck 特征后,为前三个特征层分别建一个 1×1 卷积的 MLP 投影器,把学生通道对齐到教师通道:
python
for student_out, teacher_out in zip(student_output[:-1], teacher_output[:-1]):
student_dim = self.decouple_outputs(student_out).shape[1] # 学生通道
teacher_dim = self.decouple_outputs(teacher_out).shape[1] # 教师通道
projectors.append(nn.Sequential(
nn.Conv2d(student_dim, teacher_dim, 1, padding=0), # 1×1 升/降维到教师通道
nn.ReLU(inplace=True),
nn.Conv2d(teacher_dim, teacher_dim, 1, padding=0), # 再过一个 1×1
))
self.projector = nn.ModuleList(projectors).to(device) # 3 个投影器
本例通道数(YOLO26n 学生 vs YOLO26s 教师,假设 640×640):
| 特征层 | 空间尺寸 | 学生通道(0.25×) | 教师通道(0.50×) | 投影器映射 |
|---|---|---|---|---|
| P3 / 8(层16) | 80×80 | 64 | 128 | 64→128 |
| P4 / 16(层19) | 40×40 | 128 | 256 | 128→256 |
| P5 / 32(层22) | 20×20 | 256 | 512 | 256→512 |
投影器本身也是可训练参数,会和学生一起更新,但在推理阶段被丢弃。
4. 训练阶段:单个 batch 的一次迭代
训练主循环在 trainer.py _do_train(L458-476)。关键点:self.loss, self.loss_items = self.model(batch)。
因为 DistillationModel.forward(L195-199)收到的是 dict(训练 batch),会走 loss():
python
def forward(self, x, *args, **kwargs):
if isinstance(x, dict): # 训练/验证中
return self.loss(x, *args, **kwargs)
return self.student_model.predict(x, *args, **kwargs) # 纯推理
4.1 准备 batch
一个 batch 含 2 张 640×640 图像:
python
batch = {
"img": (2, 3, 640, 640), # 归一化到 [0,1] 的 RGB 图像
"bboxes": (N, 4), # N 个真值框(xyxy/pixel)
"cls": (N,), # 每个框的类别索引
"batch_idx": (N,), # 每个框属于第几张图(0/1)
}
4.2 loss() 主流程骨架(distill_model.py L201-242)
python
loss_distill = torch.zeros(1, device=batch["img"].device)
if not self.training: # 训练中做 val 时:只用学生算常规损失
preds = self.student_model(batch["img"])
regular_loss, regular_loss_detach = self.student_model.loss(batch, preds)
return torch.cat([regular_loss, loss_distill]), torch.cat([regular_loss_detach, loss_distill])
# 训练模式:清空特征缓存
self._teacher_feats.clear()
self._student_feats.clear()
# ① 教师前向(no_grad,hook 抓教师特征)
with torch.no_grad():
self.teacher_model(batch["img"])
# ② 学生前向(可训练,hook 抓学生特征)
preds = self.student_model(batch["img"])
# ③ 常规检测损失
regular_loss, regular_loss_detach = self.student_model.loss(batch, preds)
4.3 常规检测损失(v8DetectionLoss,loss.py L478-482)
python
def loss(self, preds, batch):
batch_size = preds["boxes"].shape[0] # 2
loss, loss_detach = self.get_assigned_targets_and_loss(preds, batch)[1:]
return loss * batch_size, loss_detach # 每个分量 ×batch_size
get_assigned_targets_and_loss(L400-462)具体算三部分:
- 正样本分配 :用
TaskAlignedAssigner(L362-369)把真值框匹配到 anchor 上,得到fg_mask(前景掩码)、target_scores(每类 soft label)、target_bboxes。 - 分类损失 (
loss[1]):bce = BCEWithLogitsLoss(pred_scores, target_scores),除以target_scores_sum。 - 框损失 (
loss[0])与 DFL 损失 (loss[2]):由BboxLoss计算(CIoU + Distribution Focal Loss)。 - 乘以增益系数:
loss[0] *= hyp.box; loss[1] *= hyp.cls; loss[2] *= hyp.dfl。
返回 regular_loss = [box·2, cls·2, dfl·2](每个分量已经 ×batch_size)。
蒸馏损失部分在第 5 节单独深入展开。
5. 蒸馏损失详解(核心)
5.1 蒸馏损失在整个 loss 中的位置
#mermaid-svg-F8El2JTKsuVocMdY{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-F8El2JTKsuVocMdY .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-F8El2JTKsuVocMdY .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-F8El2JTKsuVocMdY .error-icon{fill:#552222;}#mermaid-svg-F8El2JTKsuVocMdY .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-F8El2JTKsuVocMdY .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-F8El2JTKsuVocMdY .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-F8El2JTKsuVocMdY .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-F8El2JTKsuVocMdY .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-F8El2JTKsuVocMdY .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-F8El2JTKsuVocMdY .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-F8El2JTKsuVocMdY .marker{fill:#333333;stroke:#333333;}#mermaid-svg-F8El2JTKsuVocMdY .marker.cross{stroke:#333333;}#mermaid-svg-F8El2JTKsuVocMdY svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-F8El2JTKsuVocMdY p{margin:0;}#mermaid-svg-F8El2JTKsuVocMdY .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-F8El2JTKsuVocMdY .cluster-label text{fill:#333;}#mermaid-svg-F8El2JTKsuVocMdY .cluster-label span{color:#333;}#mermaid-svg-F8El2JTKsuVocMdY .cluster-label span p{background-color:transparent;}#mermaid-svg-F8El2JTKsuVocMdY .label text,#mermaid-svg-F8El2JTKsuVocMdY span{fill:#333;color:#333;}#mermaid-svg-F8El2JTKsuVocMdY .node rect,#mermaid-svg-F8El2JTKsuVocMdY .node circle,#mermaid-svg-F8El2JTKsuVocMdY .node ellipse,#mermaid-svg-F8El2JTKsuVocMdY .node polygon,#mermaid-svg-F8El2JTKsuVocMdY .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-F8El2JTKsuVocMdY .rough-node .label text,#mermaid-svg-F8El2JTKsuVocMdY .node .label text,#mermaid-svg-F8El2JTKsuVocMdY .image-shape .label,#mermaid-svg-F8El2JTKsuVocMdY .icon-shape .label{text-anchor:middle;}#mermaid-svg-F8El2JTKsuVocMdY .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-F8El2JTKsuVocMdY .rough-node .label,#mermaid-svg-F8El2JTKsuVocMdY .node .label,#mermaid-svg-F8El2JTKsuVocMdY .image-shape .label,#mermaid-svg-F8El2JTKsuVocMdY .icon-shape .label{text-align:center;}#mermaid-svg-F8El2JTKsuVocMdY .node.clickable{cursor:pointer;}#mermaid-svg-F8El2JTKsuVocMdY .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-F8El2JTKsuVocMdY .arrowheadPath{fill:#333333;}#mermaid-svg-F8El2JTKsuVocMdY .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-F8El2JTKsuVocMdY .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-F8El2JTKsuVocMdY .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-F8El2JTKsuVocMdY .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-F8El2JTKsuVocMdY .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-F8El2JTKsuVocMdY .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-F8El2JTKsuVocMdY .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-F8El2JTKsuVocMdY .cluster text{fill:#333;}#mermaid-svg-F8El2JTKsuVocMdY .cluster span{color:#333;}#mermaid-svg-F8El2JTKsuVocMdY div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-F8El2JTKsuVocMdY .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-F8El2JTKsuVocMdY rect.text{fill:none;stroke-width:0;}#mermaid-svg-F8El2JTKsuVocMdY .icon-shape,#mermaid-svg-F8El2JTKsuVocMdY .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-F8El2JTKsuVocMdY .icon-shape p,#mermaid-svg-F8El2JTKsuVocMdY .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-F8El2JTKsuVocMdY .icon-shape .label rect,#mermaid-svg-F8El2JTKsuVocMdY .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-F8El2JTKsuVocMdY .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-F8El2JTKsuVocMdY .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-F8El2JTKsuVocMdY :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} DistillationModel.loss(batch)
① 清空特征缓存
② 教师前向 no_grad → hook 抓教师特征
③ 学生前向 → hook 抓学生特征
④ 常规损失 box+cls+dfl (v8DetectionLoss)
⑤ 教师头 scores → one2many/one2one 平均
⑥ sigmoid + 取各类 max → 空间置信度
⑦ 逐层 loss_sl2 × dis 累加
⑧ 蒸馏损失 × batch_size
总损失 = box, cls, dfl, distill
backward() 只更新学生+投影器
核心结论 :蒸馏损失是一个标量 ,与常规检测损失(box/cls/dfl)并列作为损失向量的第 4 个分量,一起 sum() 后反向传播。
5.2 蒸馏损失计算的完整代码(loss() L215-242)
python
# ① 清空特征缓存(每次都清理,避免上次残留)
self._teacher_feats.clear()
self._student_feats.clear()
# ② 教师前向(no_grad,完全冻结),hook 自动把 [16,19,22,23] 层输出存进 _teacher_feats
with torch.no_grad():
self.teacher_model(batch["img"])
# ③ 学生前向(可训练),hook 自动把 [16,19,22,23] 层输出存进 _student_feats
preds = self.student_model(batch["img"])
# ④ 常规检测损失(box, cls, dfl),每个分量已 ×batch_size
regular_loss, regular_loss_detach = self.student_model.loss(batch, preds)
# ⑤ 用「头」的输出(feats_idx[-1]=23)生成教师置信度
teacher_head_feat = self._teacher_feats[self.feats_idx[-1]] # Detect 头输出(dict)
teacher_scores = (
self.decouple_outputs(teacher_head_feat, branch="one2many")["scores"]
+ self.decouple_outputs(teacher_head_feat, branch="one2one")["scores"]
) / 2
# teacher_scores 形状: (2, 80, 8400) # logit,未 sigmoid
# ⑥ 按三个 neck 特征的空间尺寸切分 score,sigmoid 后取各类最大
neck_feats = [self._teacher_feats[idx] for idx in self.feats_idx[:-1]] # [16,19,22]
parts = torch.split(teacher_scores, [f.shape[-2] * f.shape[-1] for f in neck_feats], dim=-1)
# parts: [(2,80,6400), (2,80,1600), (2,80,400)]
teacher_scores = tuple(p.sigmoid().max(dim=1, keepdim=True).values for p in parts)
# 每个: (2,1,6400) / (2,1,1600) / (2,1,400) → 每个空间位置的「最大类概率」
# ⑦ 逐 neck 层累加蒸馏损失
for i, feat_idx in enumerate(self.feats_idx[:-1]):
teacher_feat = self.decouple_outputs(self._teacher_feats[feat_idx]) # (2, C_t, H, W)
student_feat = self.projector[i](self.decouple_outputs(self._student_feats[feat_idx])) # 对齐到 C_t
loss_distill += (self.loss_sl2(student_feat, teacher_feat,
feat_idx=i, teacher_scores=teacher_scores) * self.dis)
# ⑧ 蒸馏损失也 ×batch_size,与常规损失一起返回
distill_loss_detach = loss_distill.detach()
loss_distill = loss_distill * batch["img"].shape[0] # ×2
return torch.cat([regular_loss, loss_distill]), torch.cat([regular_loss_detach, distill_loss_detach])
5.3 特征与 score 的尺寸追踪(2 张 640×640)
三个 neck 特征层的尺寸:
| 层 | 名称 | 空间尺寸 | 学生通道 | 教师通道 | 投影后学生通道 |
|---|---|---|---|---|---|
| 16 | P3 / 8 | 80×80 | 64 | 128 | 128 |
| 19 | P4 / 16 | 40×40 | 128 | 256 | 256 |
| 22 | P5 / 32 | 20×20 | 256 | 512 | 512 |
投影器把学生特征投影到教师通道数,使两者可逐元素相减。
教师头 score 的尺寸:
text
总 anchor 数 = 80×80 + 40×40 + 20×20 = 6400 + 1600 + 400 = 8400
one2many["scores"] 形状 = (2, 80, 8400) # 2图 × 80类 × 8400个anchor
one2one["scores"] 形状 = (2, 80, 8400)
teacher_scores = 两者平均 = (2, 80, 8400) # logit
按尺度切分并转置信度:
text
parts = split((2,80,8400) → [(2,80,6400), (2,80,1600), (2,80,400)])
第 i 部分 → sigmoid() → max(dim=1, keepdim)
→ (2, 1, 6400) / (2, 1, 1600) / (2, 1, 400)
max(dim=1) 取 80 个类中概率最大的那个,得到「该空间位置最可能是什么类 + 多确信」:
#mermaid-svg-V8xQyTqZzvoCyLvq{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-V8xQyTqZzvoCyLvq .error-icon{fill:#552222;}#mermaid-svg-V8xQyTqZzvoCyLvq .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-V8xQyTqZzvoCyLvq .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-V8xQyTqZzvoCyLvq .marker{fill:#333333;stroke:#333333;}#mermaid-svg-V8xQyTqZzvoCyLvq .marker.cross{stroke:#333333;}#mermaid-svg-V8xQyTqZzvoCyLvq svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-V8xQyTqZzvoCyLvq p{margin:0;}#mermaid-svg-V8xQyTqZzvoCyLvq .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-V8xQyTqZzvoCyLvq .cluster-label text{fill:#333;}#mermaid-svg-V8xQyTqZzvoCyLvq .cluster-label span{color:#333;}#mermaid-svg-V8xQyTqZzvoCyLvq .cluster-label span p{background-color:transparent;}#mermaid-svg-V8xQyTqZzvoCyLvq .label text,#mermaid-svg-V8xQyTqZzvoCyLvq span{fill:#333;color:#333;}#mermaid-svg-V8xQyTqZzvoCyLvq .node rect,#mermaid-svg-V8xQyTqZzvoCyLvq .node circle,#mermaid-svg-V8xQyTqZzvoCyLvq .node ellipse,#mermaid-svg-V8xQyTqZzvoCyLvq .node polygon,#mermaid-svg-V8xQyTqZzvoCyLvq .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-V8xQyTqZzvoCyLvq .rough-node .label text,#mermaid-svg-V8xQyTqZzvoCyLvq .node .label text,#mermaid-svg-V8xQyTqZzvoCyLvq .image-shape .label,#mermaid-svg-V8xQyTqZzvoCyLvq .icon-shape .label{text-anchor:middle;}#mermaid-svg-V8xQyTqZzvoCyLvq .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-V8xQyTqZzvoCyLvq .rough-node .label,#mermaid-svg-V8xQyTqZzvoCyLvq .node .label,#mermaid-svg-V8xQyTqZzvoCyLvq .image-shape .label,#mermaid-svg-V8xQyTqZzvoCyLvq .icon-shape .label{text-align:center;}#mermaid-svg-V8xQyTqZzvoCyLvq .node.clickable{cursor:pointer;}#mermaid-svg-V8xQyTqZzvoCyLvq .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-V8xQyTqZzvoCyLvq .arrowheadPath{fill:#333333;}#mermaid-svg-V8xQyTqZzvoCyLvq .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-V8xQyTqZzvoCyLvq .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-V8xQyTqZzvoCyLvq .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-V8xQyTqZzvoCyLvq .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-V8xQyTqZzvoCyLvq .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-V8xQyTqZzvoCyLvq .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-V8xQyTqZzvoCyLvq .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-V8xQyTqZzvoCyLvq .cluster text{fill:#333;}#mermaid-svg-V8xQyTqZzvoCyLvq .cluster span{color:#333;}#mermaid-svg-V8xQyTqZzvoCyLvq div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-V8xQyTqZzvoCyLvq .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-V8xQyTqZzvoCyLvq rect.text{fill:none;stroke-width:0;}#mermaid-svg-V8xQyTqZzvoCyLvq .icon-shape,#mermaid-svg-V8xQyTqZzvoCyLvq .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-V8xQyTqZzvoCyLvq .icon-shape p,#mermaid-svg-V8xQyTqZzvoCyLvq .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-V8xQyTqZzvoCyLvq .icon-shape .label rect,#mermaid-svg-V8xQyTqZzvoCyLvq .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-V8xQyTqZzvoCyLvq .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-V8xQyTqZzvoCyLvq .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-V8xQyTqZzvoCyLvq :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 教师头 scores logit (2,80,6400)
sigmoid → (0,1) 概率
max over 80 classes, keepdim → (2,1,6400)
teacher_score: 每个网格点的最大类置信度
这一步把「80 维的类别概率」压缩成「1 维的空间置信度」,作为特征对齐的权重。
5.4 loss_sl2 逐行拆解(L244-264)
python
def loss_sl2(self, student_feat, teacher_feat, feat_idx, teacher_scores):
teacher_score = teacher_scores[feat_idx] # (2, 1, H*W) 该层空间置信度
n, c = student_feat.shape[:2] # n=2, c=教师通道数C_t
student_feat = student_feat.view(n, c, -1) # (2, C_t, H*W) 展平空间维
teacher_feat = teacher_feat.view(n, c, -1) # (2, C_t, H*W)
mse = F.mse_loss(student_feat, teacher_feat, reduction="none") # (2, C_t, H*W)
weighted_mse = (mse * teacher_score).sum() / (teacher_score.sum() * c + 1e-9)
return weighted_mse
公式(对每个特征层):
L s l 2 ( i ) = ∑ n , c , h , w ( s n , c , h , w ( i ) − t n , c , h , w ( i ) ) 2 ⋅ w n , h , w ( i ) ( ∑ n , h , w w n , h , w ( i ) ) ⋅ C t + ϵ L^{(i)}{sl2} = \frac{\sum{n,c,h,w} \Big( s^{(i)}{n,c,h,w} - t^{(i)}{n,c,h,w} \Big)^2 \cdot w^{(i)}{n,h,w}}{\big(\sum{n,h,w} w^{(i)}_{n,h,w}\big) \cdot C_t + \epsilon} Lsl2(i)=(∑n,h,wwn,h,w(i))⋅Ct+ϵ∑n,c,h,w(sn,c,h,w(i)−tn,c,h,w(i))2⋅wn,h,w(i)
其中:
- s s s = 投影后学生特征, t t t = 教师特征;
- w w w = 教师空间置信度(该位置最大类概率);
- C t C_t Ct = 教师通道数; ϵ = 10 − 9 \epsilon=10^{-9} ϵ=10−9 防除零。
5.5 数值实例(完整演算)
5.5.1 微型示例(P3 层,标注通道=2,位置=2)
取一小块数据演示算术过程(真实为 2×128×6400):
text
教师特征 teacher_feat (1, 2, 2) → channel0 = [0.5, 0.5] (位置0=0.5, 位置1=0.5)
channel1 = [1.0, 1.0] (位置0=1.0, 位置1=1.0)
投影后学生特征 student_feat (1, 2, 2) → channel0 = [0.4, 0.5] (位置0=0.4 → 误差在位置0)
channel1 = [1.0, 1.0]
教师空间置信度 teacher_score (1, 1, 2) → [0.9, 0.1]
# 位置0 教师很确信(0.9),位置1 教师不确定(0.1)
Step 1 --- 展平 :view(n, c, -1) 后形状不变 (1, 2, 2)。
Step 2 --- 逐元素平方差 (reduction='none'):
text
mse[0] =
channel0: [(0.5-0.4)², (0.5-0.5)²] = [0.01, 0.00]
channel1: [(1.0-1.0)², (1.0-1.0)²] = [0.00, 0.00]
Step 3 --- 乘教师置信度 w (mse * teacher_score,广播到每个通道):
text
位置0 (w=0.9): 0.01×0.9 + 0.00×0.9 = 0.009
位置1 (w=0.1): 0.00×0.1 + 0.00×0.1 = 0.000
分子 = 0.009
Step 4 --- 归一化:
text
denominator = teacher_score.sum() * c = (0.9 + 0.1) × 2 = 2.0
loss_sl2 = 0.009 / 2.0 = 0.0045
对比不加权(普通 MSE 平均):
text
plain_mse = sum(mse) / (c × H×W) = 0.01 / 4 = 0.0025
加权后 0.0045 > 不加权 0.0025:因为误差恰好出现在教师高置信度 的位置0,被放大;而低置信度的位置1误差为0,不受影响。这就是 score-weighting 的意义------把容量集中在教师认为有物体的地方。
5.5.2 反向示例(误差在低置信度位置)
若误差只在位置1(教师不确信处):
text
student channel0 = [0.5, 0.6] → mse channel0 = [0.00, 0.01]
teacher_score = [0.9, 0.1]
分子 = 0.00×0.9 + 0.01×0.1 = 0.001
loss_sl2 = 0.001 / 2.0 = 0.0005
误差在低置信度位置时,蒸馏损失被大幅压缩(0.0005),学生不必强行模仿教师没把握的特征。
5.5.3 完整 batch 尺度(2 张图,3 个 neck 层)
对每一层 i(P3/P4/P5)都算一个 loss_sl2,然后累加并乘 dis:
text
loss_sl2_P3 = 0.0045 (示意)
loss_sl2_P4 = 0.0032 (示意)
loss_sl2_P5 = 0.0028 (示意)
loss_distill = (0.0045 + 0.0032 + 0.0028) × dis(6.0)
= 0.0105 × 6.0
= 0.063
最终蒸馏分量 = loss_distill × batch_size(2) = 0.126
返回的损失向量:
text
loss_items = [box×2, cls×2, dfl×2, distill×2]
= [ ... , 0.126]
训练器: self.loss = loss.sum() → 反向传播
5.6 梯度流向
#mermaid-svg-GGNfsaRiepwefxf3{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-GGNfsaRiepwefxf3 .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-GGNfsaRiepwefxf3 .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-GGNfsaRiepwefxf3 .error-icon{fill:#552222;}#mermaid-svg-GGNfsaRiepwefxf3 .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-GGNfsaRiepwefxf3 .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-GGNfsaRiepwefxf3 .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-GGNfsaRiepwefxf3 .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-GGNfsaRiepwefxf3 .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-GGNfsaRiepwefxf3 .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-GGNfsaRiepwefxf3 .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-GGNfsaRiepwefxf3 .marker{fill:#333333;stroke:#333333;}#mermaid-svg-GGNfsaRiepwefxf3 .marker.cross{stroke:#333333;}#mermaid-svg-GGNfsaRiepwefxf3 svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-GGNfsaRiepwefxf3 p{margin:0;}#mermaid-svg-GGNfsaRiepwefxf3 .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-GGNfsaRiepwefxf3 .cluster-label text{fill:#333;}#mermaid-svg-GGNfsaRiepwefxf3 .cluster-label span{color:#333;}#mermaid-svg-GGNfsaRiepwefxf3 .cluster-label span p{background-color:transparent;}#mermaid-svg-GGNfsaRiepwefxf3 .label text,#mermaid-svg-GGNfsaRiepwefxf3 span{fill:#333;color:#333;}#mermaid-svg-GGNfsaRiepwefxf3 .node rect,#mermaid-svg-GGNfsaRiepwefxf3 .node circle,#mermaid-svg-GGNfsaRiepwefxf3 .node ellipse,#mermaid-svg-GGNfsaRiepwefxf3 .node polygon,#mermaid-svg-GGNfsaRiepwefxf3 .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-GGNfsaRiepwefxf3 .rough-node .label text,#mermaid-svg-GGNfsaRiepwefxf3 .node .label text,#mermaid-svg-GGNfsaRiepwefxf3 .image-shape .label,#mermaid-svg-GGNfsaRiepwefxf3 .icon-shape .label{text-anchor:middle;}#mermaid-svg-GGNfsaRiepwefxf3 .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-GGNfsaRiepwefxf3 .rough-node .label,#mermaid-svg-GGNfsaRiepwefxf3 .node .label,#mermaid-svg-GGNfsaRiepwefxf3 .image-shape .label,#mermaid-svg-GGNfsaRiepwefxf3 .icon-shape .label{text-align:center;}#mermaid-svg-GGNfsaRiepwefxf3 .node.clickable{cursor:pointer;}#mermaid-svg-GGNfsaRiepwefxf3 .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-GGNfsaRiepwefxf3 .arrowheadPath{fill:#333333;}#mermaid-svg-GGNfsaRiepwefxf3 .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-GGNfsaRiepwefxf3 .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-GGNfsaRiepwefxf3 .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-GGNfsaRiepwefxf3 .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-GGNfsaRiepwefxf3 .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-GGNfsaRiepwefxf3 .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-GGNfsaRiepwefxf3 .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-GGNfsaRiepwefxf3 .cluster text{fill:#333;}#mermaid-svg-GGNfsaRiepwefxf3 .cluster span{color:#333;}#mermaid-svg-GGNfsaRiepwefxf3 div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-GGNfsaRiepwefxf3 .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-GGNfsaRiepwefxf3 rect.text{fill:none;stroke-width:0;}#mermaid-svg-GGNfsaRiepwefxf3 .icon-shape,#mermaid-svg-GGNfsaRiepwefxf3 .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-GGNfsaRiepwefxf3 .icon-shape p,#mermaid-svg-GGNfsaRiepwefxf3 .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-GGNfsaRiepwefxf3 .icon-shape .label rect,#mermaid-svg-GGNfsaRiepwefxf3 .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-GGNfsaRiepwefxf3 .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-GGNfsaRiepwefxf3 .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-GGNfsaRiepwefxf3 :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 投影后学生特征
教师特征 (冻结, requires_grad=False)
学生特征 (可训练)
投影器 projector (可训练)
loss_sl2
教师置信度 w (no_grad)
× dis → distill 损失
backward()
更新 学生参数
更新 投影器参数
- 教师特征与置信度都来自
no_grad()前向,不产生梯度; - 梯度只经过投影器 流向学生特征,从而更新学生与投影器;
- 教师参数
requires_grad=False,永不更新。
6. 训练日志中的 dis_loss
trainer.py L376-377:有 distill_model 时追加 dis_loss 到 loss_names:
python
if self.args.distill_model is not None and "dis_loss" not in self.loss_names:
self.loss_names += ("dis_loss",)
训练器(L466-476)把 4 分量 loss_items 求和后反向传播:
python
loss, self.loss_items = self.model(batch) # 4 个分量 tensor
self.loss = loss.sum() # box + cls + dfl + distill
self.scaler.scale(self.loss).backward() # 反向传播
loss_items 由 label_loss_items(detect/train.py L213-228)映射成带名字的 dict:
text
Epoch GPU_mem box_loss cls_loss dfl_loss dis_loss Instances Size
1/100 4.52G 0.9123 0.5612 0.8001 0.3350 4 640
这里的 dis_loss=0.3350 就是 loss_distill(已 ×dis ×batch_size)的均值。
梯度只更新 :学生模型参数 + projector 投影器参数。教师参数 requires_grad=False,不产生梯度。
7. 推理阶段
蒸馏只在训练期生效,推理是纯学生模型,零额外开销。
- 训练时
DistillationModel.forward收到普通张量(非 dict)时走self.student_model.predict(x)。 - 保存的
best.pt已经把DistillationModel外壳(含教师、投影器)剥掉,只留学生权重(见第 8 节)。 - 因此推理时加载的是一个普通
DetectionModel,走标准 YOLO 推理:
#mermaid-svg-DwdTRrQ0Ghh5yUmR{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;fill:#333;}@keyframes edge-animation-frame{from{stroke-dashoffset:0;}}@keyframes dash{to{stroke-dashoffset:0;}}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-animation-slow{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 50s linear infinite;stroke-linecap:round;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-animation-fast{stroke-dasharray:9,5!important;stroke-dashoffset:900;animation:dash 20s linear infinite;stroke-linecap:round;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .error-icon{fill:#552222;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .error-text{fill:#552222;stroke:#552222;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-thickness-normal{stroke-width:1px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-thickness-thick{stroke-width:3.5px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-pattern-solid{stroke-dasharray:0;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-thickness-invisible{stroke-width:0;fill:none;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-pattern-dashed{stroke-dasharray:3;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edge-pattern-dotted{stroke-dasharray:2;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .marker{fill:#333333;stroke:#333333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .marker.cross{stroke:#333333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR svg{font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:16px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR p{margin:0;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .label{font-family:"trebuchet ms",verdana,arial,sans-serif;color:#333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .cluster-label text{fill:#333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .cluster-label span{color:#333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .cluster-label span p{background-color:transparent;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .label text,#mermaid-svg-DwdTRrQ0Ghh5yUmR span{fill:#333;color:#333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .node rect,#mermaid-svg-DwdTRrQ0Ghh5yUmR .node circle,#mermaid-svg-DwdTRrQ0Ghh5yUmR .node ellipse,#mermaid-svg-DwdTRrQ0Ghh5yUmR .node polygon,#mermaid-svg-DwdTRrQ0Ghh5yUmR .node path{fill:#ECECFF;stroke:#9370DB;stroke-width:1px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .rough-node .label text,#mermaid-svg-DwdTRrQ0Ghh5yUmR .node .label text,#mermaid-svg-DwdTRrQ0Ghh5yUmR .image-shape .label,#mermaid-svg-DwdTRrQ0Ghh5yUmR .icon-shape .label{text-anchor:middle;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .node .katex path{fill:#000;stroke:#000;stroke-width:1px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .rough-node .label,#mermaid-svg-DwdTRrQ0Ghh5yUmR .node .label,#mermaid-svg-DwdTRrQ0Ghh5yUmR .image-shape .label,#mermaid-svg-DwdTRrQ0Ghh5yUmR .icon-shape .label{text-align:center;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .node.clickable{cursor:pointer;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .root .anchor path{fill:#333333!important;stroke-width:0;stroke:#333333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .arrowheadPath{fill:#333333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edgePath .path{stroke:#333333;stroke-width:2.0px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .flowchart-link{stroke:#333333;fill:none;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edgeLabel{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edgeLabel p{background-color:rgba(232,232,232, 0.8);}#mermaid-svg-DwdTRrQ0Ghh5yUmR .edgeLabel rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-DwdTRrQ0Ghh5yUmR .labelBkg{background-color:rgba(232, 232, 232, 0.5);}#mermaid-svg-DwdTRrQ0Ghh5yUmR .cluster rect{fill:#ffffde;stroke:#aaaa33;stroke-width:1px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .cluster text{fill:#333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .cluster span{color:#333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR div.mermaidTooltip{position:absolute;text-align:center;max-width:200px;padding:2px;font-family:"trebuchet ms",verdana,arial,sans-serif;font-size:12px;background:hsl(80, 100%, 96.2745098039%);border:1px solid #aaaa33;border-radius:2px;pointer-events:none;z-index:100;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .flowchartTitleText{text-anchor:middle;font-size:18px;fill:#333;}#mermaid-svg-DwdTRrQ0Ghh5yUmR rect.text{fill:none;stroke-width:0;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .icon-shape,#mermaid-svg-DwdTRrQ0Ghh5yUmR .image-shape{background-color:rgba(232,232,232, 0.8);text-align:center;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .icon-shape p,#mermaid-svg-DwdTRrQ0Ghh5yUmR .image-shape p{background-color:rgba(232,232,232, 0.8);padding:2px;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .icon-shape .label rect,#mermaid-svg-DwdTRrQ0Ghh5yUmR .image-shape .label rect{opacity:0.5;background-color:rgba(232,232,232, 0.8);fill:rgba(232,232,232, 0.8);}#mermaid-svg-DwdTRrQ0Ghh5yUmR .label-icon{display:inline-block;height:1em;overflow:visible;vertical-align:-0.125em;}#mermaid-svg-DwdTRrQ0Ghh5yUmR .node .label-icon path{fill:currentColor;stroke:revert;stroke-width:revert;}#mermaid-svg-DwdTRrQ0Ghh5yUmR :root{--mermaid-font-family:"trebuchet ms",verdana,arial,sans-serif;} 2×640×640 图像
backbone
neck (P3/P4/P5)
Detect 头
one2one 分支解码 + sigmoid
NMS / 端到端 top-k
检测结果 boxes + scores + classes
注意:
yolo26.yaml中end2end: True,因此 YOLO26 推理默认走 one2one(NMS-free) 分支;关闭端到端则用传统 NMS。
8. Checkpoint 保存 / 恢复 / EMA 处理
8.1 保存时剥离教师(torch_utils.py L796-804)
python
if x.get("ema"):
x["model"] = x["ema"]
if isinstance(x["model"], DistillationModel):
x["model"]._remove_feature_hooks()
x["model"] = x["model"].student_model # 只存学生,去掉教师+投影器外壳
教师和投影器不写入 checkpoint,节省磁盘/内存。推理、加载都只拿到学生。
8.2 EMA 也剥离教师(torch_utils.py L723-727)
python
self.ema = deepcopy(unwrap_model(model)).eval()
if hasattr(self.ema, "teacher_model"):
self.ema.teacher_model = None # EMA 不复制一份完整教师
否则 EMA 会多存一份完整教师,浪费一倍内存。
8.3 恢复训练(trainer.py L777-789)
python
if isinstance(weights, DistillationModel):
student_model = self.get_model(cfg=cfg, weights=weights.student_model, ...)
teacher_model = weights.teacher_model if weights.teacher_model is not None else self.args.distill_model
model = DistillationModel(student_model=student_model, teacher_model=teacher_model)
if getattr(weights, "projector", None) is not None:
model.projector.load_state_dict(weights.projector.state_dict()) # 恢复投影器
self.model = model
因教师被剥离,恢复时按
distill_model路径重新加载教师,并把训练好的投影器权重恢复回去。
8.4 NaN 恢复(trainer.py L1006-1011)
若训练出现 NaN,从 last.pt 恢复 EMA 时,因为 EMA 里没有教师,只把 student_model 和 projector 单独 load 回去,保证 key 严格匹配。
9. 关键设计要点与性能参考
9.1 设计要点
- 特征来源 :蒸馏发生在三个 neck 层(P3/P4/P5),用
forward hook自动采集,无需改网络结构。 - 教师完全冻结 :
eval()+requires_grad=False+torch.no_grad(),不参与梯度。 - 投影器:学生通常比教师窄,用 3 个 1×1 卷积 MLP 把学生通道对齐到教师通道,投影器随学生一起训练。
- score-weighted L2:用教师分类置信度(sigmoid 后取各类最大)作为空间权重,让损失聚焦在有物体的位置。
- 逐点加权:教师置信度高的位置,特征对齐需求更强烈 → 放大误差;背景/不确定处 → 忽略。
- 多尺度:P3/P4/P5 三个尺度都蒸馏,覆盖大/中/小目标。
dis平衡 :控制蒸馏损失与常规损失的量级关系(默认 6.0),每个尺度蒸馏损失 ×dis,再 ×batch_size。- zero-cost 推理:checkpoint 只存学生,推理是标准学生模型,无教师、无投影器、无蒸馏开销。
- 任务支持 :detect / segment / pose / obb 均兼容(头继承
Detect),但官方仅验证了 detect 的精度提升。
9.2 当前 YOLO26 蒸馏 mAP 参考(官方文档)
| Model | size (pixels) | mAPval 50-95 baseline | mAPval 50-95 distilled | mAPval 50-95 (e2e) baseline | mAPval 50-95 (e2e) distilled |
|---|---|---|---|---|---|
| YOLO26n-distill | 640 | 40.9 | 41.5 | 40.1 | 40.9 |
| YOLO26s-distill | 640 | 48.6 | 49.2 | 47.8 | 48.6 |
| YOLO26m-distill | 640 | 53.1 | 53.9 | 52.5 | 53.3 |
| YOLO26l-distill | 640 | 55.0 | 56.0 | 54.4 | 55.5 |
| YOLO26x-distill | 640 | 57.5 | 57.9 | 56.9 | 57.4 |