Ultralytics:解读 YOLO26 知识蒸馏

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.py L201-242))
      • [4.3 常规检测损失(`v8DetectionLoss`,`loss.py` L478-482)](#4.3 常规检测损失(v8DetectionLossloss.py L478-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.py L796-804))
      • [8.2 EMA 也剥离教师(`torch_utils.py` L723-727)](#8.2 EMA 也剥离教师(torch_utils.py L723-727))
      • [8.3 恢复训练(`trainer.py` L777-789)](#8.3 恢复训练(trainer.py L777-789))
      • [8.4 NaN 恢复(`trainer.py` L1006-1011)](#8.4 NaN 恢复(trainer.py L1006-1011))
    • [9. 关键设计要点与性能参考](#9. 关键设计要点与性能参考)
      • [9.1 设计要点](#9.1 设计要点)
      • [9.2 当前 YOLO26 蒸馏 mAP 参考(官方文档)](#9.2 当前 YOLO26 蒸馏 mAP 参考(官方文档))
  • 参考
  • 由于本人水平有限,难免出现错漏,敬请批评改正。
  • 更多精彩内容,可点击进入我的个人主页查看
    整合自两份源码精读:YOLO26 知识蒸馏(训练与推理全过程) + YOLO26 蒸馏损失逐行精讲

基于 ultralytics 源码:ultralytics/nn/distill_model.pyultralytics/engine/trainer.pyultralytics/utils/loss.pyultralytics/nn/modules/head.pyultralytics/nn/tasks.pyultralytics/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.yamlDetect 在层 52,其 from[16, 19, 22]):

复制代码
feats_idx = [16, 19, 22, 23]
  • feats_idx[:-1] = [16, 19, 22]三个 neck 特征层(P3/P4/P5),蒸馏的特征就在这里取。
  • feats_idx[-1] = 23Detect 头,用来生成教师分类置信度(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 常规检测损失(v8DetectionLossloss.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)具体算三部分:

  1. 正样本分配 :用 TaskAlignedAssigner(L362-369)把真值框匹配到 anchor 上,得到 fg_mask(前景掩码)、target_scores(每类 soft label)、target_bboxes
  2. 分类损失loss[1]):bce = BCEWithLogitsLoss(pred_scores, target_scores),除以 target_scores_sum
  3. 框损失loss[0])与 DFL 损失loss[2]):由 BboxLoss 计算(CIoU + Distribution Focal Loss)。
  4. 乘以增益系数: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 --- 乘教师置信度 wmse * 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_lossloss_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_itemslabel_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.yamlend2end: 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_modelprojector 单独 load 回去,保证 key 严格匹配。


9. 关键设计要点与性能参考

9.1 设计要点

  1. 特征来源 :蒸馏发生在三个 neck 层(P3/P4/P5),用 forward hook 自动采集,无需改网络结构。
  2. 教师完全冻结eval() + requires_grad=False + torch.no_grad(),不参与梯度。
  3. 投影器:学生通常比教师窄,用 3 个 1×1 卷积 MLP 把学生通道对齐到教师通道,投影器随学生一起训练。
  4. score-weighted L2:用教师分类置信度(sigmoid 后取各类最大)作为空间权重,让损失聚焦在有物体的位置。
  5. 逐点加权:教师置信度高的位置,特征对齐需求更强烈 → 放大误差;背景/不确定处 → 忽略。
  6. 多尺度:P3/P4/P5 三个尺度都蒸馏,覆盖大/中/小目标。
  7. dis 平衡 :控制蒸馏损失与常规损失的量级关系(默认 6.0),每个尺度蒸馏损失 × dis,再 ×batch_size。
  8. zero-cost 推理:checkpoint 只存学生,推理是标准学生模型,无教师、无投影器、无蒸馏开销。
  9. 任务支持 :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

参考

1 https://docs.ultralytics.com/

2 https://github.com/ultralytics/ultralytics.git

相关推荐
蓝田~1 小时前
AI Demo到上线有多远?→ 安全纵深防御+Token成本精确计算+可观测性,PrismAI三周工程化复盘
人工智能·安全
东坡肘子1 小时前
Apple Intelligence 已通过审核,即将在中国提供服务 -- 肘子的 Swift 周报 #148
人工智能·swiftui·swift
BUG研究员_1 小时前
工具绑定与调用
python·agent
苏灿烤鱼1 小时前
AI Agent 深拆 | 图能让 AI 决策可追责吗?Semantica 登顶拆解
python·github·agent
watersink1 小时前
机器学习FM、FFM
人工智能·机器学习
九硕智慧建筑一体化厂家1 小时前
从光伏板到护眼灯,直流照明如何重构绿色校园的能源基因
运维·人工智能·重构·智慧城市·能源
user-猴子2 小时前
2026国产AI智能体七强横评:WorkBuddy、AiPy、悟空、Kimi Work、TRAE Work、QoderWork、百度搭子怎么选?
人工智能·百度·dubbo
Sirius.z9 小时前
第R6周:LSTM实现糖尿病探索与预测
python
名字还没想好☜9 小时前
Python f-string 进阶:数字格式化、对齐填充、调试 = 号与嵌套表达式
开发语言·数据库·python·字符串格式化·f-string