YOLOv5 7.0:核心模块详解

03 核心模块详解

3.1 models/yolo.py

Detect

YOLOv5 检测头。主要成员:

成员 含义
nc 类别数
no = nc + 5 每个锚点的输出维度
nl 检测层数量,P5 模型通常为 3
na 每层锚点数,通常为 3
anchors 以网格单位保存的 Anchor
stride 各检测层步长
m 每层一个 1×1 输出卷积
grid 网格中心坐标
anchor_grid Anchor 的像素尺度

训练时返回三层原始张量:

1
[B, na, H, W, nc+5]

推理时解码并拼接为:

1
[B, 所有候选数, nc+5]

Segment

继承 Detect

  • Proto 从高分辨率特征产生共享 mask 原型
  • 每个候选额外预测 nm 个 mask 系数
  • 默认 nm=32npr=256

BaseModel

通用能力:

  • _forward_once():按 YAML 图执行
  • _profile_one_layer():逐层耗时/FLOPs
  • fuse():Conv+BN 融合
  • _apply():迁移 stride/grid/anchor_grid 到设备或精度

DetectionModel

职责:

  1. 读取模型 YAML
  2. 覆盖 nc 或 anchors
  3. 调用 parse_model
  4. 用 256×256 假输入推断 P3/P4/P5 stride
  5. 检查 Anchor 顺序并归一化
  6. 初始化 Detect bias

它还支持三尺度 TTA:缩放为 1、0.83、0.67,并做水平翻转后合并结果。

ClassificationModel

可从检测模型截取前若干层作为 backbone,并用 Classify 替换尾部模块。

parse_model

这是读懂 YOLOv5 模型图的核心函数:

1
2
3
4
5
6
YAML 字典
→ 缩放深度/宽度
→ 计算每层输入输出通道
→ 创建模块
→ 记录 from 索引
→ 生成 nn.Sequential 和 save list

3.2 models/common.py

基础网络块

作用
Conv Conv2d + BatchNorm + SiLU
DWConv 深度/分组卷积
Bottleneck 残差瓶颈
C3 CSP 风格双分支与瓶颈堆叠
SPPF 串行 5×5 MaxPool 快速空间金字塔
Concat 通道拼接
Proto 分割原型掩码
Classify 分类头

YOLOv5s 的主体模块是 Conv + C3 + SPPF

DetectMultiBackend

它不是网络结构,而是部署适配层。构造函数通过权重后缀判断运行时,并加载:

1
2
3
4
5
.pt / .torchscript / .onnx / .engine / .mlmodel
OpenVINO 模型目录
SavedModel / .pb / .tflite / EdgeTPU
Paddle 模型目录
Triton URL

forward() 负责:

  • PyTorch Tensor ↔ NumPy
  • NCHW ↔ NHWC
  • FP16
  • TensorRT 动态 shape/binding
  • TFLite INT8 量化与反量化
  • 统一输出为当前设备上的 Tensor

AutoShape

PyTorch Hub 的易用包装器,接受:

  • 路径/URL
  • PIL
  • NumPy
  • OpenCV 图
  • Tensor
  • 图片列表

非 Tensor 输入会自动完成三通道转换、Letterbox、BCHW、归一化、前向、NMS、坐标缩放,并返回 Detections

注意:注释明确要求 OpenCV BGR 输入先转 RGB,例如 cv2.imread(...)[..., ::-1]

Detections

封装结果展示和导出:

  • print() / show() / save()
  • crop()
  • render()
  • pandas()
  • tolist()

3.3 models/experimental.py

主要用于:

  • attempt_load():加载一个或多个 .pt
  • Conv+BN 融合
  • 模型兼容处理
  • 多模型 Ensemble

DetectMultiBackend 加载 PyTorch 权重时调用此处。

3.4 utils/dataloaders.py

推理加载器

输入
LoadImages 图片、视频、目录、glob
LoadStreams 摄像头、RTSP/RTMP/HTTP、多路流
LoadScreenshots 屏幕截图

训练加载器

LoadImagesAndLabels 负责:

  • 搜索图像和映射标签路径
  • 验证图片/标签
  • 建立 .cache
  • 矩形训练
  • RAM/磁盘缓存
  • Mosaic、MixUp、随机透视、HSV、翻转
  • 标签坐标在归一化 xywh 与像素 xyxy 间转换
  • BGR→RGB、HWC→CHW、连续内存

create_dataloader() 再包成 InfiniteDataLoader 或普通 DataLoader,并在 DDP 下使用分布式采样器。

3.5 utils/augmentations.py

关键函数:

函数/类 作用
letterbox 保持比例缩放并填充到 stride 倍数
random_perspective 旋转、平移、缩放、剪切、透视
augment_hsv HSV LUT 颜色增强
mixup 两张图与标签混合
copy_paste 分割数据复制粘贴
Albumentations 可选第三方增强包装

3.6 utils/loss.py

ComputeLoss 的组成:

1
2
3
总损失 = box gain × CIoU loss
+ obj gain × BCE objectness
+ cls gain × BCE classification

它还支持:

  • Label smoothing
  • FocalLoss
  • 各检测尺度 Objectness balance
  • Anchor 宽高比匹配
  • 相邻网格偏移匹配

3.7 utils/general.py

高频函数包括:

  • non_max_suppression
  • scale_boxes
  • xyxy2xywh / xywh2xyxy
  • check_dataset / check_file / check_img_size
  • increment_path
  • strip_optimizer
  • YAML 加载/保存

这是入口脚本最常依赖的通用工具集。

3.8 utils/metrics.py

负责:

  • box_iou
  • AP/mAP 计算
  • ap_per_class
  • 混淆矩阵
  • fitness 计算

验证脚本在 IoU 阈值 0.50 到 0.95 上判断预测正确性,输出 Precision、Recall、mAP@0.5mAP@0.5:0.95。

3.9 utils/torch_utils.py

关键工程能力:

  • select_device
  • smart_optimizer
  • smart_DDP
  • ModelEMA
  • AMP 检查
  • EarlyStopping
  • Conv+BN 融合
  • 分布式同步辅助

3.10 回调和日志

  • utils/callbacks.py::Callbacks:维护事件与动作
  • utils/loggers/__init__.py::Loggers:统一日志适配

主训练循环只触发事件,具体日志平台通过注册回调接入,实现控制流与外部服务解耦。

文章互动

阅读 --

留言

0 条留言

正在加载留言…