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=32、npr=256
BaseModel
通用能力:
_forward_once():按 YAML 图执行_profile_one_layer():逐层耗时/FLOPsfuse():Conv+BN 融合_apply():迁移 stride/grid/anchor_grid 到设备或精度
DetectionModel
职责:
- 读取模型 YAML
- 覆盖
nc或 anchors - 调用
parse_model - 用 256×256 假输入推断 P3/P4/P5 stride
- 检查 Anchor 顺序并归一化
- 初始化 Detect bias
它还支持三尺度 TTA:缩放为 1、0.83、0.67,并做水平翻转后合并结果。
ClassificationModel
可从检测模型截取前若干层作为 backbone,并用 Classify 替换尾部模块。
parse_model
这是读懂 YOLOv5 模型图的核心函数:
1 | YAML 字典 |
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 | .pt / .torchscript / .onnx / .engine / .mlmodel |
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 | 总损失 = box gain × CIoU loss |
它还支持:
- Label smoothing
- FocalLoss
- 各检测尺度 Objectness balance
- Anchor 宽高比匹配
- 相邻网格偏移匹配
3.7 utils/general.py
高频函数包括:
non_max_suppressionscale_boxesxyxy2xywh/xywh2xyxycheck_dataset/check_file/check_img_sizeincrement_pathstrip_optimizer- YAML 加载/保存
这是入口脚本最常依赖的通用工具集。
3.8 utils/metrics.py
负责:
box_iou- AP/mAP 计算
ap_per_class- 混淆矩阵
- fitness 计算
验证脚本在 IoU 阈值 0.50 到 0.95 上判断预测正确性,输出 Precision、Recall、mAP@0.5、mAP@0.5:0.95。
3.9 utils/torch_utils.py
关键工程能力:
select_devicesmart_optimizersmart_DDPModelEMA- AMP 检查
- EarlyStopping
- Conv+BN 融合
- 分布式同步辅助
3.10 回调和日志
utils/callbacks.py::Callbacks:维护事件与动作utils/loggers/__init__.py::Loggers:统一日志适配
主训练循环只触发事件,具体日志平台通过注册回调接入,实现控制流与外部服务解耦。
正在加载留言…