1. 项目概述:YOLO目标检测入门指南
作为一名计算机视觉方向的从业者,我经常被问到如何快速入门目标检测领域。YOLO(You Only Look Once)作为当前最流行的实时目标检测算法之一,以其出色的速度和精度平衡,成为许多初学者的首选。本文将从一个实践者的角度,带您从零开始搭建第一个YOLO目标检测项目。
对于完全没有深度学习基础的小白来说,YOLO可能是最友好的入门选择。相比传统的两阶段检测算法(如Faster R-CNN),YOLO采用单阶段检测架构,将目标检测任务转化为回归问题,这种端到端的设计理念大大简化了实现难度。最新版本的YOLOv8甚至提供了近乎"开箱即用"的体验,让初学者也能快速获得可观的检测效果。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境准备与工具选型
2.1 硬件配置建议
虽然YOLO对硬件要求相对友好,但合适的配置能显著提升开发体验。对于入门学习,我建议至少准备:
- GPU:NVIDIA GTX 1660及以上(4GB显存起步)
- 内存:16GB DDR4
- 存储:256GB SSD(用于系统和环境)+ 1TB HDD(存储数据集)
注意:如果没有独立GPU,可以使用Google Colab的免费GPU资源,但需要注意其运行时长限制。
2.2 软件环境搭建
推荐使用conda创建独立的Python环境,避免依赖冲突:
bash复制conda create -n yolo_env python=3.8
conda activate yolo_env
pip install torch torchvision torchaudio
pip install ultralytics # YOLOv8官方库
对于视觉任务,OpenCV是必不可少的工具:
bash复制pip install opencv-python
pip install opencv-contrib-python
2.3 开发工具选择
- IDE:VS Code + Python插件(轻量级)或PyCharm Professional(功能全面)
- 版本控制:Git + GitHub/GitLab
- 可视化工具:LabelImg(标注工具)、TensorBoard(训练监控)
3. YOLO核心原理浅析
3.1 YOLO算法思想精髓
YOLO的核心创新在于将目标检测视为单一的回归问题,直接从图像像素到边界框坐标和类别概率。这与传统的滑动窗口或区域提议方法形成鲜明对比。具体来说:
- 将输入图像划分为S×S的网格
- 每个网格预测B个边界框及其置信度
- 每个边界框包含5个预测值:(x,y,w,h,confidence)
- 每个网格还预测C个类别概率
这种设计使得YOLO可以一次性完成所有预测,实现了极高的推理速度。
3.2 YOLOv8架构改进
最新版的YOLOv8在之前版本基础上做了多项优化:
- 更高效的骨干网络(CSPDarknet53改进版)
- 更精确的锚框设计
- 改进的损失函数(CIoU Loss)
- 更灵活的多尺度预测
这些改进使得YOLOv8在保持实时性的同时,检测精度显著提升,特别适合新手快速获得不错的效果。
4. 实战:第一个YOLO检测项目
4.1 数据集准备与标注
对于初学者,建议从公开数据集开始:
- COCO:80类常见物体,通用性强
- Pascal VOC:20类物体,适合教学
- 自定义数据:使用LabelImg标注自己的数据集
标注格式示例(YOLO格式):
code复制<class_id> <x_center> <y_center> <width> <height>
4.2 模型训练基础流程
使用YOLOv8训练只需几行代码:
python复制from ultralytics import YOLO
# 加载预训练模型
model = YOLO('yolov8n.pt') # 纳米尺寸模型
# 训练配置
results = model.train(
data='coco128.yaml',
epochs=100,
imgsz=640,
batch=16,
device='0' # 使用GPU 0
)
关键参数说明:
imgsz:输入图像尺寸(影响精度和速度)batch:批大小(受显存限制)device:指定训练设备
4.3 模型评估与推理
训练完成后,可以方便地进行评估和推理:
python复制# 评估模型性能
metrics = model.val()
# 进行单张图像推理
results = model('path/to/image.jpg')
# 显示结果
results[0].show()
5. 常见问题与解决方案
5.1 训练过程中的典型问题
-
显存不足(OOM)
- 降低
batch_size - 减小
imgsz - 使用梯度累积
- 降低
-
损失值震荡
- 调整学习率(默认0.01可能过大)
- 增加数据增强
- 检查数据标注质量
-
过拟合
- 增加数据多样性
- 使用早停(Early Stopping)
- 添加正则化
5.2 推理阶段的常见问题
-
小目标检测效果差
- 使用更高分辨率的输入(增加
imgsz) - 尝试专门针对小目标优化的模型(如YOLOv8s)
- 使用更高分辨率的输入(增加
-
检测框偏离目标
- 检查标注是否准确
- 调整锚框参数
- 尝试CIoU或DIoU损失函数
-
多类别混淆
- 增加困难样本
- 调整类别权重
- 考虑使用Focal Loss
6. 性能优化技巧
6.1 模型压缩与加速
-
量化:将FP32模型转为INT8,显著减小模型体积并提升速度
python复制model.export(format='onnx', int8=True) -
剪枝:移除不重要的神经元连接
python复制
model.prune() -
知识蒸馏:用大模型指导小模型训练
6.2 部署优化
-
TensorRT加速:NVIDIA官方推理优化引擎
python复制model.export(format='engine') -
ONNX Runtime:跨平台高性能推理
python复制model.export(format='onnx') -
移动端部署:转换为CoreML或TFLite格式
python复制model.export(format='tflite')
7. 项目扩展与进阶
7.1 多路视频流处理
对于监控等需要处理多路视频的场景,可以使用多线程:
python复制import threading
def process_stream(rtsp_url, model):
cap = cv2.VideoCapture(rtsp_url)
while True:
ret, frame = cap.read()
results = model(frame)
# 显示或保存结果
# 创建多个处理线程
threads = []
urls = ['rtsp://cam1', 'rtsp://cam2']
for url in urls:
t = threading.Thread(target=process_stream, args=(url, model))
threads.append(t)
t.start()
7.2 自定义模型开发
当基础模型无法满足需求时,可以修改网络结构:
python复制from ultralytics.nn.tasks import DetectionModel
class CustomModel(DetectionModel):
def __init__(self, cfg='yolov8n.yaml'):
super().__init__(cfg)
# 添加自定义模块
def forward(self, x):
# 修改前向传播逻辑
return super().forward(x)
7.3 领域自适应技巧
-
迁移学习:冻结骨干网络,只训练检测头
python复制model.train(freeze=[0, 1, 2]) # 冻结前3层 -
数据增强策略:
- Mosaic增强
- MixUp
- 随机透视变换
-
测试时增强(TTA):
python复制results = model.predict(..., augment=True)
8. 学习资源与社区
8.1 推荐学习路径
-
基础理论:
- 《Deep Learning for Computer Vision》
- YOLO原始论文阅读
-
实战项目:
- Kaggle目标检测竞赛
- Roboflow提供的教程项目
-
前沿追踪:
- arXiv上的最新论文
- Ultralytics官方博客
8.2 优质开源项目
-
官方实现:
- Ultralytics YOLOv8
- YOLOv5 (旧版但稳定)
-
衍生项目:
- YOLOR (多任务学习)
- YOLOX (Anchor-free改进)
-
部署优化:
- TensorRT-YOLO
- ONNX-YOLO
9. 避坑指南与经验分享
9.1 新手常见误区
-
盲目追求最新模型
- 最新版不一定最适合你的场景
- 考虑YOLOv5等成熟版本可能更稳定
-
忽视数据质量
- 标注错误比模型问题更常见
- 建议至少检查10%的标注样本
-
超参数调整过度
- 初学者应先使用默认参数
- 一次只调整一个参数并记录变化
9.2 实用小技巧
-
学习率设置:
- 使用学习率预热(warmup)
- 余弦退火调度器效果不错
-
数据不平衡处理:
- 过采样少数类
- 使用类别加权损失
-
模型集成:
- 多个模型的预测结果加权融合
- 可以提升2-3%的mAP
10. 实际应用案例
10.1 智能安防系统
使用YOLO实现的人流统计系统架构:
- 视频流输入(RTSP/RTMP)
- YOLO实时检测(人/车/包裹)
- 基于检测结果的业务逻辑
- 人数统计
- 异常行为识别
- 遗留物检测
10.2 工业质检应用
PCB缺陷检测实现方案:
-
数据采集:
- 正常PCB图像
- 各类缺陷样本(短路、断路等)
-
模型训练:
- 使用YOLOv8s模型
- 高分辨率输入(1280x1280)
- 针对小缺陷优化
-
部署方案:
- 工厂边缘计算设备
- 实时检测速度≥30FPS
10.3 农业应用案例
果园果实检测与计数系统:
-
数据特点:
- 复杂自然环境(光照变化大)
- 果实遮挡严重
- 多尺度目标
-
解决方案:
- 多尺度训练(320-1280)
- 改进的NMS算法
- 针对遮挡优化的损失函数
-
业务价值:
- 产量预估准确率>90%
- 成熟度检测
- 自动化采摘引导
11. 模型解释与可视化
11.1 特征图可视化
理解模型"看到"的内容:
python复制import torch
from torchvision.utils import make_grid
# 获取中间层输出
activation = {}
def get_activation(name):
def hook(model, input, output):
activation[name] = output.detach()
return hook
model.model[10].register_forward_hook(get_activation('layer10'))
# 可视化
with torch.no_grad():
output = model(img_tensor)
act = activation['layer10']
grid = make_grid(act[0][:16].unsqueeze(1), nrow=4)
11.2 检测结果分析
使用Grad-CAM理解检测决策:
python复制from pytorch_grad_cam import GradCAM
target_layers = [model.model[-2]] # 最后卷积层
cam = GradCAM(model=model, target_layers=target_layers)
grayscale_cam = cam(input_tensor, targets=None)
# 叠加到原图
visualization = show_cam_on_image(img, grayscale_cam)
11.3 性能瓶颈分析
使用PyTorch Profiler找出耗时操作:
python复制with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU,
torch.profiler.ProfilerActivity.CUDA]
) as p:
model(input_tensor)
print(p.key_averages().table(sort_by="cuda_time_total"))
12. 持续学习建议
12.1 技能提升路线
-
基础夯实:
- Python编程
- PyTorch框架深入
- 计算机视觉基础
-
领域深入:
- 目标检测专题
- 模型优化技术
- 部署工程化
-
横向扩展:
- 多目标跟踪
- 实例分割
- 行为识别
12.2 实验记录方法
-
实验管理工具:
- Weights & Biases
- TensorBoard
- MLflow
-
记录要点:
- 超参数配置
- 数据版本
- 环境信息
- 性能指标
-
分析技巧:
- 消融实验设计
- 误差分析
- 可视化对比
13. 模型部署实战
13.1 本地服务化部署
使用FastAPI创建推理服务:
python复制from fastapi import FastAPI, File, UploadFile
import cv2
import numpy as np
app = FastAPI()
model = YOLO('best.pt')
@app.post("/predict")
async def predict(file: UploadFile = File(...)):
contents = await file.read()
nparr = np.frombuffer(contents, np.uint8)
img = cv2.imdecode(nparr, cv2.IMREAD_COLOR)
results = model(img)
return {"predictions": results[0].boxes.data.tolist()}
13.2 边缘设备部署
在Jetson系列设备上的优化:
-
转换为TensorRT引擎:
bash复制
trtexec --onnx=yolov8n.onnx --saveEngine=yolov8n.engine -
使用Triton推理服务器:
- 配置模型仓库
- 优化批处理策略
- 动态加载模型
13.3 移动端集成
Android端集成步骤:
-
转换为TFLite格式:
python复制model.export(format='tflite') -
Android Studio配置:
- 添加TFLite依赖
- 加载模型文件
- 实现预处理/后处理
-
性能优化:
- 使用GPU代理
- 量化推理
- 多线程处理
14. 行业应用展望
14.1 新兴应用场景
-
元宇宙与虚拟现实:
- 实时3D物体检测
- 虚拟物品交互
-
自动驾驶演进:
- 多传感器融合检测
- 时序上下文理解
-
医疗影像分析:
- 病灶自动标注
- 手术导航辅助
14.2 技术融合趋势
-
与Transformer结合:
- DETR系列算法
- 注意力机制改进
-
多模态学习:
- 视觉-语言联合理解
- 跨模态检索
-
自监督学习:
- 减少标注依赖
- 预训练范式创新
15. 社区贡献与开源文化
15.1 参与开源项目
-
贡献方式:
- 提交bug报告
- 完善文档
- 添加新特性
-
优秀项目推荐:
- Ultralytics YOLO
- MMDetection
- Detectron2
-
协作工具:
- GitHub工作流
- 代码审查规范
- 问题跟踪系统
15.2 知识分享途径
-
技术博客写作:
- 记录学习过程
- 分享解决方案
- 复盘项目经验
-
开源教程制作:
- Jupyter Notebook示例
- Colab实战项目
- 视频教程系列
-
社区互动:
- Stack Overflow解答
- GitHub Discussions
- 技术论坛参与
16. 伦理与责任
16.1 技术应用边界
-
隐私保护:
- 人脸检测的合规使用
- 数据脱敏处理
- 用户授权机制
-
偏见与公平:
- 数据集多样性检查
- 模型公平性评估
- 纠偏技术应用
-
可解释性:
- 检测决策透明化
- 错误案例分析
- 用户告知义务
16.2 可持续发展
-
能效优化:
- 模型轻量化
- 绿色计算实践
- 碳足迹评估
-
长期维护:
- 代码可读性
- 文档完整性
- 版本兼容性
-
人才培养:
- 新手友好设计
- 教学资源建设
- 社区互助文化
17. 故障排除手册
17.1 训练阶段问题
-
Loss不下降:
- 检查学习率是否合适
- 验证数据加载是否正确
- 尝试更简单的模型
-
显存泄漏:
- 检查张量是否及时释放
- 使用memory_profiler工具
- 减少数据加载线程
-
梯度爆炸:
- 添加梯度裁剪
- 使用更稳定的激活函数
- 检查输入数据范围
17.2 推理阶段问题
-
检测框抖动:
- 添加时序平滑滤波
- 提高置信度阈值
- 使用更稳定的模型
-
漏检率高:
- 调整NMS参数
- 降低置信度阈值
- 增加输入分辨率
-
推理速度慢:
- 启用半精度推理
- 使用TensorRT加速
- 优化预处理流水线
18. 性能基准测试
18.1 测试方法论
-
评估指标:
- mAP (mean Average Precision)
- FPS (Frames Per Second)
- 显存占用
-
测试环境:
- 固定硬件配置
- 控制温度条件
- 多次运行取平均
-
对比维度:
- 不同输入分辨率
- 不同批大小
- 不同精度模式
18.2 典型结果参考
YOLOv8各版本在COCO上的表现:
| 模型 | mAP@0.5 | FPS (T4) | 参数量(M) |
|---|---|---|---|
| YOLOv8n | 37.3 | 450 | 3.2 |
| YOLOv8s | 44.9 | 300 | 11.4 |
| YOLOv8m | 50.2 | 180 | 26.2 |
| YOLOv8l | 52.9 | 120 | 43.7 |
| YOLOv8x | 53.9 | 80 | 68.2 |
测试环境:NVIDIA T4 GPU, TensorRT 8.4, FP16精度
19. 模型版本管理
19.1 最佳实践
-
命名规范:
- 包含数据集和指标信息
- 示例:yolov8s_coco_640_v1.2
-
版本控制:
- Git LFS管理大文件
- 清晰的提交信息
- 定期打标签
-
模型注册表:
- MLflow Model Registry
- 自定义元数据存储
- 版本依赖关系
19.2 回滚策略
-
性能降级处理:
- 自动回归测试
- 快速回滚机制
- A/B测试部署
-
数据漂移应对:
- 监控输入分布
- 保留旧版本模型
- 渐进式更新
-
灾难恢复:
- 多地备份
- 模型签名验证
- 文档化恢复流程
20. 项目实战进阶
20.1 自定义任务扩展
-
关键点检测:
- 扩展检测头
- 修改损失函数
- 数据标注格式调整
-
实例分割:
- 添加掩码头
- 使用YOLOv8-seg模型
- 多边形标注数据
-
多任务学习:
- 共享骨干网络
- 任务特定头
- 平衡损失权重
20.2 工业级优化
-
数据流水线优化:
- 异步数据加载
- 智能缓存策略
- 分布式采样
-
训练加速:
- 混合精度训练
- 梯度累积
- 分布式数据并行
-
推理优化:
- 批处理策略
- 模型量化
- 算子融合
21. 跨平台开发
21.1 Web集成方案
-
前端展示:
- WebSocket实时传输
- Canvas绘制检测框
- 交互式结果筛选
-
后端服务:
- FastAPI异步框架
- Redis任务队列
- 结果缓存机制
-
全栈示例:
javascript复制// 前端调用 const response = await fetch('/predict', { method: 'POST', body: formData }); const results = await response.json();
21.2 桌面应用集成
-
PyQt方案:
python复制from PyQt5.QtCore import QThread, pyqtSignal class DetectionThread(QThread): result_ready = pyqtSignal(object) def run(self): results = model(frame) self.result_ready.emit(results) -
Electron方案:
- 使用node.js调用Python
- 进程间通信
- 原生UI集成
-
性能考量:
- 内存管理
- 线程安全
- 跨平台兼容性
22. 模型监控与维护
22.1 生产环境监控
-
关键指标:
- 推理延迟
- 吞吐量
- 成功率
-
异常检测:
- 输入分布偏移
- 输出置信度异常
- 资源使用突增
-
报警机制:
- Prometheus + Grafana
- 分级报警策略
- 自动扩容设置
22.2 模型迭代
-
持续训练:
- 新数据收集
- 增量学习
- 主动学习策略
-
A/B测试:
- 流量分配策略
- 指标对比
- 平稳切换
-
版本迁移:
- 兼容性测试
- 灰度发布
- 回滚预案
23. 团队协作规范
23.1 代码管理
-
Git规范:
- 功能分支工作流
- 有意义的提交信息
- 代码审查流程
-
项目结构:
code复制project/ ├── data/ ├── models/ ├── notebooks/ ├── src/ │ ├── train.py │ └── inference.py ├── tests/ └── README.md -
文档标准:
- API文档生成
- 示例代码
- 变更日志
23.2 知识共享
-
技术评审:
- 方案设计评审
- 代码走查
- 经验分享会
-
文档沉淀:
- 决策记录(ADR)
- 问题解决记录
- 最佳实践指南
-
新人培养:
- 导师制度
- 渐进式任务
- 定期反馈
24. 成本优化策略
24.1 训练成本控制
-
云资源使用:
- Spot实例利用
- 自动启停脚本
- 资源监控告警
-
算法优化:
- 早停策略
- 模型剪枝
- 数据筛选
-
分布式训练:
- 梯度压缩
- 异步更新
- 数据并行
24.2 推理成本优化
-
硬件选型:
- 性价比分析
- 能效比考量
- 长期成本预测
-
模型量化:
- FP16量化
- INT8量化
- 稀疏化
-
缓存策略:
- 结果缓存
- 模型预热
- 请求合并
25. 安全与隐私
25.1 模型安全
-
对抗攻击防护:
- 输入检测
- 对抗训练
- 鲁棒性评估
-
模型保护:
- 权重加密
- 模型水印
- 混淆技术
-
API安全:
- 速率限制
- 身份验证
- 请求验证
25.2 数据隐私
-
匿名化处理:
- 人脸模糊
- 元数据清除
- 差分隐私
-
联邦学习:
- 数据不出域
- 参数聚合
- 安全多方计算
-
合规审计:
- 数据流向追踪
- 访问日志
- 定期检查
26. 前沿技术追踪
26.1 YOLO系列演进
-
v1-v3经典架构:
- Darknet骨干
- 多尺度预测
- 基础设计理念
-
v4-v6优化期:
- CSP结构
- PANet改进
- 自注意力引入
-
v7-v8创新期:
- 可重参数化
- 动态标签分配
- 任务特定头
26.2 相关算法对比
主流目标检测算法比较:
| 算法 | 特点 | 适用场景 | 学习曲线 |
|---|---|---|---|
| YOLO | 速度快,精度平衡 | 实时检测 | 平缓 |
| Faster R-CNN | 精度高,速度慢 | 精度优先 | 陡峭 |
| SSD | 折中方案 | 移动端 | 中等 |
| RetinaNet | 处理类别不平衡 | 密集检测 | 较陡 |
| DETR | 端到端,无NMS | 研究创新 | 陡峭 |
27. 职业发展建议
27.1 技能矩阵构建
计算机视觉工程师核心能力:
-
基础能力:
- Python编程
- 深度学习框架
- 数学基础
-
核心能力:
- 目标检测专精
- 模型优化
- 部署工程化
-
扩展能力:
- 多模态学习
- 大模型应用
- 全栈开发
27.2 项目经验积累
-
个人项目:
- 复现经典论文
- 解决实际问题
- 参与开源贡献
-
竞赛经历:
- Kaggle
- 天池
- ECCV/ICCV比赛
-
实习经验:
- 工业级项目
- 团队协作流程
- 产品化思维
28. 学习路线图
28.1 新手阶段(0-3个月)
-
学习重点:
- Python基础
- PyTorch入门
- YOLO基础使用
-
推荐项目:
- 水果检测
- 车牌识别
- 安全帽检测
-
里程碑:
- 完整训练流程掌握
- 简单应用开发
- 基础性能调优
28.2 进阶阶段(3-6个月)
-
学习重点:
- 模型改进
- 部署优化
- 工程化实践
-
推荐项目:
- 多路视频分析
- 移动端部署
- 自定义模型开发
-
里程碑:
- 工业级应用开发
- 性能瓶颈分析
- 完整项目经验
28.3 专家阶段(6个月+)
-
学习重点:
- 算法创新
- 系统架构
- 团队管理
-
推荐方向:
- 学术研究
- 开源项目主导
- 技术方案设计
-
里程碑:
- 专利/论文产出
- 大型项目领导
- 技术决策能力
29. 工具链推荐
29.1 开发工具
-
IDE:
- VS Code + Python插件
- PyCharm Professional
- Jupyter Lab
-
调试工具:
- PyTorch Debugger
- CUDA-MEMCHECK
- Python Profiler
-
协作工具:
- Git + GitHub
- Docker
- CI/CD流水线
29.2 可视化工具
-
数据探索:
- Label Studio
- FiftyOne
- CVAT
-
训练监控:
- TensorBoard
- Weights & Biases
- MLflow
-
结果分析:
- Detectron2 Visualizer
- OpenCV绘图工具
- Plotly交互可视化
30. 实用代码片段
30.1 数据增强技巧
自定义Mosaic增强:
python复制def mosaic_augmentation(images, labels, size=640):
# 创建空白画布
mosaic_img = np.zeros((size*2, size*2, 3), dtype=np.uint8)
mosaic_labels = []
# 随机选择4个位置
positions = [(0,0), (size,0), (0,size), (size,size)]
random.shuffle(positions)
for i, (img, lbl) in enumerate(zip(images, labels)):
x, y = positions[i]
# 调整图像大小并放置到对应位置
img = cv2.resize(img, (size, size))
mosaic_img[y:y+size, x:x+size] = img
# 调整标签坐标
lbl[:, [1,3]] = (lbl[:, [1,3]] * size + x) / (size*2)
lbl[:, [2,4]] = (lbl[:, [2,4]] * size + y) / (size*2)
mosaic_labels.append(lbl)
return mosaic_img, np.concatenate(mosaic_labels)
30.2 模型导出优化
带后处理的ONNX导出:
python复制class YOLOWrapper(torch.nn.Module):
def __init__(self, model):
super().__init__()
self.model = model
def forward(self, x):
# 原始推理
preds = self.model(x)
# 添加NMS后处理
boxes, scores, labels = non_max_suppression(preds)
return boxes, scores, labels
# 导出
wrapped_model = YOLOWrapper(model)
torch.onnx.export(wrapped_model, dummy_input, "yolo_with_nms.onnx")
30.3 性能分析工具
GPU利用率监控:
python复制import pynvml
pynvml.nvmlInit()
handle = pynvml.nvmlDeviceGetHandleByIndex(0)
def get_gpu_utilization():
util = pynvml.nvmlDeviceGetUtilizationRates(handle)
return util.gpu, util.memory
while training:
gpu_util, mem_util = get_gpu_utilization()
print(f"GPU: {gpu_util}%, Memory: {mem_util}%")
time.sleep(1)
31. 模型解释性
31.1 注意力可视化
显示模型关注区域:
python复制def visualize_attention(img, model):
# 获取注意力图
features = model.backbone(img)
attention = model.neck(features)[-1].mean(dim=1)
# 可视化
plt.imshow(img[0].permute(1,2,0).cpu())
plt.imshow(attention[0].cpu(), alpha=0.5, cmap='jet')
plt.show()
31.2 错误案例分析
典型错误分类样本分析:
python复制def analyze_errors(model, val_loader):
errors = []
for imgs, targets in val_loader:
preds = model(imgs)
for i, (pred, target) in enumerate(zip(preds, targets)):
if not match(pred, target):
errors.append({
'image': imgs[i],
'pred': pred,
'target': target
})
return errors
31.3 特征相似度
计算特征空间距离:
python复制from sklearn.metrics.pairwise import cosine_similarity
def feature_similarity(model, img1, img2):
feat1 = model.backbone(img1).flatten()
feat2 = model.backbone(img2).flatten()
return cosine_similarity([feat1], [feat2])[0][0]
32. 多模型集成
32.1 加权框融合
合并多个模型的预测结果:
python复制def weighted_box_fusion(detections_list, weights=None):
if weights is None:
weights = [1/len(detections_list)] * len(detections_list)
all_boxes = []
for detections, weight in zip(detections_list, weights):
for box in detections:
box['weight'] = weight
all_boxes.append(box)
# 聚类相似框
clusters = []
while all_boxes:
box = all_boxes.pop()
cluster = [box]
i = 0
while i < len(all_boxes):
if iou(box, all_boxes[i]) > 0.5:
cluster.append(all_boxes.pop(i))
else:
i += 1
clusters.append(cluster)
# 计算加权平均
fused_boxes = []
for cluster in clusters:
avg_box = {
'x1': sum(b['x1']*b['weight'] for b in cluster)/sum(b['weight'] for b in cluster),
# 其他坐标同理...
'confidence': sum(b['confidence']*b['weight'] for b in cluster)/sum(b
