Ultralytics YOLO 模型头(head.py)深度解析:Detect、Segment、Pose、OBB、RTDETRDecoder 等全部任务头的结构与实现
【免费下载链接】ultralyticsUltralytics YOLO26, YOLO11, YOLOv8 — object detection, instance segmentation, semantic segmentation, image classification, pose estimation, object tracking项目地址: https://gitcode.com/GitHub_Trending/ul/ultralytics
本文基于 docs/en/reference/nn/modules/head.md 的 API 参考,结合 ultralytics/nn/modules/head.py 的源码实现,系统讲解 Ultralytics YOLO 系列模型中 17 个任务头(Head)模块的结构、前向流程与关键参数。读完后,你将理解检测、实例分割、旋转框、关键点、深度估计、分类、语义分割以及开放词汇检测(YOLOE / YOLO-World)与 RT-DETR 各自在"模型最后一层"是如何输出预测结果的,并能看懂模型 YAML 配置中 head 部分与源码类的对应关系。
1. head.py 在整个模型中的地位
Ultralytics 的 YOLO 模型由 Backbone、Neck、Head 三段组成,而 ultralytics/nn/modules/head.py 是"最后一层"——它把 Neck 输出的多尺度特征图(通常是 P3/8、P4/16、P5/32)转换成本任务所需的最终预测(边框、类别、掩码系数、关键点、深度图、logits 等)。
该文件的模块导出列表(__all__)位于 head.py#L21-L33,包含 OBB、Classify、Depth、Detect、Pose、RTDETRDecoder、Segment、SemanticSegment、YOLOEDetect、YOLOESegment、v10Detect 等公开类。文档页面 docs/en/reference/nn/modules/head.md 则额外列出了 Segment26、OBB26、Pose26、WorldDetect、LRPCHead、YOLOESegment26 这些衍生类,全部 17 个类都在同一文件内实现。
理解这些类之前,先记住两条贯穿全文件的通用设计:
- one2many / one2one 双分支结构。以 Detect 为例,训练时使用"一对多"(one-to-many,NMS 或 TAL 分配)分支,推理时使用"一对一"(one-to-one,端到端免 NMS)分支。当
end2end=True时,构造函数会深拷贝出one2one_cv2/one2one_cv3(见 head.py#L139-L141)。 - fuse() 剪枝。导出模型前会调用各 head 的
fuse(),把不再使用的 one2many 分支置为None,例如 Detect.fuse 执行self.cv2 = self.cv3 = None,从而减小推理图体积。
2. Detect:YOLO 检测头
Detect 是所有检测类头的基类。构造函数签名为Detect(nc=80, reg_max=16, end2end=False, ch=())(head.py#L106):
nc:类别数;reg_max:DFL(Distribution Focal Loss)bins 数,输出通道no = nc + reg_max * 4(head.py#L119);end2end:是否启用端到端免 NMS 检测;ch:来自 Backbone 的特征图通道元组,其长度即检测层数nl。
2.1 网络结构
每个尺度上有两条独立的头(head.py#L121-L137):
- 边框回归分支 cv2:
Conv(x, c2, 3) → Conv(c2, c2, 3) → Conv2d(c2, 4*reg_max, 1),其中c2 = max(16, ch[0]//4, reg_max*4); - 分类分支 cv3:新版结构使用深度可分离卷积降低算力,即
DWConv(x, x, 3) → Conv(x, c3, 1)两层堆叠后接 1×1 卷积输出nc个分数;legacy=True时回退为 v3/v5/v8/v9 的普通 Conv 结构以兼容旧权重; - DFL 层:
reg_max > 1时为DFL(self.reg_max),否则退化为nn.Identity()。
注意 YOLO26 的配置里reg_max: 1(见 yolo26.yaml#L10),即边框输出直接是 4 个距离值,不做 DFL 分布回归——这也是reg_max > 1 if DFL else Identity设计的实际意义。
2.2 前向与输出格式
forward()的核心逻辑(head.py#L174-L188):
- 先用
self.one2many调用forward_head得到{"boxes", "scores", "feats"}; - 若开启 end2end,则对特征 detach 后(训练中防止 one2one 分支梯度回传进 Backbone)再跑一次
self.one2one,返回{"one2many": ..., "one2one": ...}; - 训练态返回预测字典供损失函数使用;推理态调用
_inference解码边框并拼接sigmoid(scores),若 end2end 还会经postprocess做 top-k 筛选,最终输出(y, preds)或纯y(export 模式)。
_get_decode_boxes(head.py#L203-L211)展示了 anchor 的惰性重建:当输入 shape 变化或dynamic=True时,通过make_anchors(来自 ultralytics/utils/tal.py)生成 anchors 与 strides,再经dist2bbox完成l,t,r,b距离到 xywh 的解码。
2.3 端到端后处理与 top-k
postprocess(head.py#L236-L251)接收(B, num_anchors, 4+nc+extra)的原始预测,通过get_topk_index选出每图前max_det(默认 300,类属性)个检测,返回[x1, y1, x2, y2, max_prob, class_idx, extra]。其中_grouped_topk(head.py#L89-L100)把 anchor 轴分成 8 组分别 topk 再合并,避免对上万 anchor 直接全排序,在 TensorRT engine 导出(self.format == "engine"且非 dynamic)时会启用分组加速。
2.4 偏置初始化
bias_init(head.py#L213-L225)在训练开始、stride 计算完成后调用:边框分支偏置置 2.0,分类偏置置log(5/nc/(640/stride)²),即假设"约 1% 的 anchor 是前景、640 分辨率",帮助早期收敛。该注释同时说明此函数"requires stride availability"。
3. Segment 与 Segment26:实例分割头
Segment 继承 Detect,新增参数nm=32(掩码系数数)与npr=256(原型通道数),核心部件是(head.py#L301-L320):
Proto(ch[0], npr, nm):作用在最高分辨率特征图 x[0] 上的原型(prototypes)生成模块;cv4:每尺度一个Conv(x, c4, 3) → Conv(c4, c4, 3) → Conv2d(c4, nm, 1)的掩码系数分支,c4 = max(ch[0]//4, nm)。
前向时在 Detect 输出之外再算proto = self.proto(x[0]),并把proto塞进训练预测字典(end2end 时 one2one 侧用 detach 版本);推理态返回((y, proto), preds)(head.py#L332-L345)。解码侧_inference把mask_coefficient拼在解码结果后面,即torch.cat([preds, x["mask_coefficient"]], dim=1)。
Segment26 是 YOLO26 版本:唯一区别是把原型模块换为Proto26(ch, npr, nm, nc),且forward中调用Detect.forward并把整个特征列表(而非仅 x[0])传给 proto(head.py#L399-L417)。对应的模型配置见 yolo26-seg.yaml。
4. OBB 与 OBB26:旋转框检测头
OBB 在 Detect 基础上增加ne=1个角度参数分支cv4。两个关键实现细节:
- 角度映射:
forward_head中角度 logits 经(angle.sigmoid() - 0.25) * math.pi映射到[-π/4, 3π/4](head.py#L489-L493),与旋转框的周期性匹配; - 旋转解码:
decode_bboxes改用dist2rbox(bboxes, self.angle, anchors, dim=1)而非普通dist2bbox(head.py#L496-L498)。
OBB26 的差异更微妙:它直接调用Detect.forward_head并输出原始 angle logits,不做 sigmoid 变换(head.py#L525-L536),即角度处理被推迟到更后面的阶段。配置可参考 yolo26-obb.yaml 与 yolo11-obb.yaml。
5. Pose 与 Pose26:关键点头
Pose 的关键参数是kpt_shape,格式为(K, D):K 个关键点、D 维(2 为 x,y;3 为 x,y,可见性)。总标量数nk = K*D,由cv4分支逐尺度输出。
kpts_decode(head.py#L608-L624)把相对偏移还原为绝对坐标:xy = (offset * 2 + (anchor - 0.5)) * stride,可见性维做 sigmoid;export 路径走 reshape 版本以保证图编译友好。
Pose26 引入了更精细的参数化(head.py#L648-L671):
- 引入
RealNVP()归一化流模块flow_model; - 中间特征分支
cv4输出宽度扩为K*(D+2),拆成cv4_kpts(nk 维关键点)与cv4_sigma(nk_sigma = K*2,即每个关键点独立的 σx、σy 方差)两个 1×1 卷积头; - 训练时额外输出
kpts_sigma供不确定性建模,fuse()会把cv4_sigma、flow_model等推理无用组件置空。
6. Depth:单目深度估计头
Depth 不走 Detect 体系,而是稠密预测解码器:proj把 P3/P4/P5 三个尺度统一投影到c_mid=256通道,然后自粗到细"上采样 2 倍 → 相加 → refine 两卷 3×3 卷积"逐层融合,最后head网络经两次 Conv+ConvTranspose 上采样到输入 1/4 分辨率输出单通道(head.py#L760-L787)。
几个值得注意的实现点:
- 输出偏置初始化为 0.182,使早期
exp()输出约 1.2 m,保持数值良态(head.py#L782-L783); - 深度取
exp(out.clamp(-4.0, 5.0)),限制在e^-4 ~ e^5区间,避免溢出; - 提供
cal_a/cal_b两个 buffer 做对数仿射标定d' = d^a * e^b,默认为单位映射,可用于推理期尺度校正(head.py#L785-L820); - export 模式下会再上采样 4 倍到输入分辨率输出
(B, 1, H, W)。
训练态返回{"depth": ...}字典,eval 态返回应用标定后的张量。对应配置为 yolo26-depth.yaml。
7. Classify:分类头
Classify 结构直白:Conv(c1, 1280, k, s) → AdaptiveAvgPool2d(1) → Dropout(0) → Linear(1280, c2),其中 1280 取自 EfficientNet-b0 的隐层宽度(head.py#L859)。输入若是特征图列表会先沿通道拼接。训练态返回 logits,推理态返回(softmax, logits)或 export 时的纯概率。模型配置参考 yolo11-cls.yaml / yolo26-cls.yaml。
8. 开放词汇检测头:WorldDetect 与 YOLOE 系列
8.1 WorldDetect(YOLO-World)
WorldDetect 让检测头接受文本嵌入text(B, C, embed)作为输入。其cv3分支输出 embed 维(默认 512)视觉特征而非类别分数,类别得分由对比头算出:x[i] = cat(cv2i, cv4i, text), 1)(head.py#L922-L935),即逐位置余弦相似度充当逐类置信度。由于推理时文本可变,self.nc会在 forward 内被文本类别数动态改写(head.py#L927)。cv4可选BNContrastiveHead(带 BN 与可学习 logit_scale)或ContrastiveHead。
8.2 YOLOEDetect:文本提示可融合的开放词汇头
YOLOEDetect 在 WorldDetect 思想之上增加了两样东西:
- 文本/视觉提示嵌入:
reprta(Residual + SwiGLUFFN)处理文本提示嵌入并 L2 归一化(get_tpe),savpe(SAVPE)处理带空间位置的视觉提示(get_vpe,支持 (B,N,H,W) 输入)(head.py#L1142-L1153); - 权重级融合:
fuse(txt_feats)把文本嵌入"烧进"卷积权重——_fuse_tp中用t @ w、t @ b把 (embed→nc) 的 1×1 卷积折叠为 (1→K) 的固定卷积,随后reprta被替换为 Identity(head.py#L1087-L1140)。融合后模型即变成"prompt-free"形态,走forward_lrpc路径:仅使用 one2one 分支,逐尺度由 LRPCHead 完成轻量区域筛选与分类,可用conf阈值(静态导出时 conf=0,全量保留 anchor)过滤低分 anchor 以减少推理开销。
YOLOESegment 与 YOLOESegment26 在该基础上加上Proto/Proto26原型与cv5掩码系数分支,forward_lrpc中同时筛选掩码系数mc[..., index]。对应配置:yoloe-26.yaml、yoloe-26-seg.yaml、yoloe-11.yaml。
9. RTDETRDecoder:Transformer 检测头
RTDETRDecoder 与基于 anchor grid 的 YOLO 头完全不同,它是"查询式"检测头。关键参数(head.py#L1492-L1510):
| 参数 | 默认 | 含义 |
|---|---|---|
nc | 80 | 类别数 |
ch | (512, 1024, 2048) | 骨干特征通道 |
hd | 256 | Transformer 隐层维度 |
nq | 300 | 查询(query)数量 |
nh | 8 | 多头注意力头数 |
ndl | 6 | 解码器层数 |
d_ffn | 1024 | FFN 维度 |
nd | 100 | 去噪(denoising)查询数 |
label_noise_ratio/box_noise_scale | 0.5 / 1.0 | 训练期标签/框噪声 |
learnt_init_query | False | 是否学习初始查询嵌入 |
前向流程(head.py#L1571-L1621):
input_proj将各尺度特征投影到hd维并 flatten 为 (B, H*W, C)(_get_encoder_input);get_cdn_group(来自 ultralytics/models/utils/ops.py)生成去噪训练用的 dn 嵌入与注意力掩码;_get_decoder_input做查询选择:在编码器输出上用enc_score_head打分类分,取 top-nq anchor 特征作为初始查询、top-nq 网格 anchor 作为参考框(head.py#L1737-L1765);DeformableTransformerDecoder(定义于 ultralytics/nn/modules/transformer.py)逐层迭代出dec_bboxes、dec_scores;- 训练态返回
(dec_bboxes, dec_scores, enc_bboxes, enc_scores, dn_meta)五元组供匈牙利匹配损失;推理态经postprocess取 top-k 输出(bs, nq, 6)的[cx, cy, w, h, prob, cls]。
_reset_parameters展示了典型的 DETR 式初始化:分类偏置用bias_init_with_prob(0.01)、回归 MLP 末层置零(保证初始预测等于 anchor),enc_output用 xavier 初始化(head.py#L1767-L1789)。
10. v10Detect 与 SemanticSegment
v10Detect 是 YOLOv10 的检测头:类属性直接end2end = True(head.py#L1817),并把分类分支替换为"深度可分卷积 + 1×1 卷积"的轻量结构(Conv(x, x, 3, g=x)即逐通道 3×3 卷积),配合 dual-assignment 训练实现 NMS-free 推理;fuse()同样剪掉 one2many 分支。
SemanticSegment 是语义分割头:输入 P3、P4 两路特征,classifier在 P3 上输出[B, nc, H/8, W/8]像素级 logits,aux_head在 P4 上做辅助监督(仅训练态返回)(head.py#L1882-L1899)。导出时有个很实用的优化:ONNX/MNN/OpenVINO 以及 TensorRT≥10 / 多类 Hailo 会把 argmax 烘焙进图内,直接输出紧凑的[B, H, W]类别图(uint8/int32),减少约 80 倍的显存拷贝(head.py#L1900-L1908)。对应配置 yolo26-sem.yaml。
11. 从 YAML 配置到源码类的映射
模型 YAML 的 head 段直接指明使用哪个头。对比两份配置即可看出差异:
- yolo11.yaml 第 50 行
- [[16, 19, 22], 1, Detect, [nc]]:P3/P4/P5 三个 Neck 输出送入 Detect,无 end2end 参数(默认 False,推理走 one2many + 外部 NMS 或后处理); - yolo26.yaml 第 9 行
end2end: True与第 10 行reg_max: 1:整个模型启用端到端免 NMS 分支且不做 DFL,与 head.py#L137 的DFL(reg_max) if reg_max > 1 else Identity分支精确对应;第 52 行- [[16, 19, 22], 1, Detect, [nc]]的[16, 19, 22]即 Backbone/Neck 中三个尺度特征层的索引。
各任务模型配置可按ultralytics/cfg/models/{26,11,12,10,...}/目录查找,如 yolo26-pose.yaml、yolo26-seg.yaml、yolo26-obb.yaml。
12. 独立验证一个头:可运行的最小示例
各类 docstring 均给出了可直接运行的构造示例,例如检测头(对应 head.py#L71-L76):
import torch from ultralytics.nn.modules.head import Detect, Segment, OBB, Pose, Depth, Classify # 构造三级特征图(P3/8、P4/16、P5/32,160 输入对应 20 网格) x = [ torch.randn(1, 256, 80, 80), torch.randn(1, 512, 40, 40), torch.randn(1, 1024, 20, 20), ] detect = Detect(nc=80, ch=(256, 512, 1024)) out = detect(x) # 推理态返回 (y, preds),y 形状为 (1, 84+4*reg_max, A) # 实例分割头:多返回 proto seg = Segment(nc=80, nm=32, npr=256, ch=(256, 512, 1024)) # 旋转框 / 关键点 / 深度 obb = OBB(nc=80, ne=1, ch=(256, 512, 1024)) pose = Pose(nc=80, kpt_shape=(17, 3), ch=(256, 512, 1024)) depth = Depth(ch=(256, 512, 1024)) cls_head = Classify(c1=1024, c2=1000)需要留意的是:解码类头(Detect 及其子类)在_get_decode_boxes中依赖make_anchors生成的 anchors 与self.stride(build 阶段由模型填充),单独构造头对象做推理时 anchors 会自动按输入 shape 重建,但bias_init依赖self.stride,因此"requires stride availability"。仓库的导出测试 tests/test_exports.py 会把这些头经真实模型导出到各格式,可作为行为正确性的回归依据。
13. 小结
- ultralytics/nn/modules/head.py 以
Detect为基类,通过子类化(Segment/OBB/Pose 及各自 26 版本、YOLOE 系列、v10Detect)覆盖了 Ultralytics 全部任务头的实现,另有 Depth、Classify、SemanticSegment、RTDETRDecoder 四个独立实现; - 统一的设计语言是:one2many/one2one 双分支、
forward_head抽象、_inference解码、fuse()导出剪枝,以及训练态返回结构化预测字典、推理态返回解码张量的双模式前向; - 阅读模型 YAML 时,head 段最后一行的类名(
Detect、Segment、OBB、Pose等)即可直接对应到本文件的类;end2end、reg_max、kpt_shape、nm/npr等全局/行内参数分别控制端到端推理、DFL 粒度、关键点形状与掩码原型; - 官方 API 参考页 docs/en/reference/nn/modules/head.md 按类分节列出了全部 17 个头的文档锚点,可与本文各节一一对照,深入每个方法的签名与示例。
【免费下载链接】ultralyticsUltralytics YOLO26, YOLO11, YOLOv8 — object detection, instance segmentation, semantic segmentation, image classification, pose estimation, object tracking项目地址: https://gitcode.com/GitHub_Trending/ul/ultralytics
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考