共计 3000 个字符,预计需要花费 8 分钟才能阅读完成。
背景与痛点
垃圾分类是城市管理中的重要环节,而计算机视觉技术可以大大提高分类效率。目标检测(Object Detection)技术能够自动识别图像中的垃圾并判断其类别,为智能垃圾桶、垃圾分拣机器人等应用提供技术支持。

然而在实际应用中,我们常常会遇到以下问题:
- 类别不平衡(Class Imbalance):某些类别的样本数量远多于其他类别
- 边界框偏移(Bounding Box Shift):标注框未能准确覆盖目标物体
- 标注不一致(Inconsistent Annotation):不同标注者对同一物体的标注存在差异
- 遮挡物体处理不当(Occlusion Handling):未能正确处理部分遮挡的物体
数据准备
标注规范示例
我们定义了 6 类生活垃圾:
- 可回收物(Recyclable)
- 厨余垃圾(Kitchen Waste)
- 有害垃圾(Hazardous Waste)
- 其他垃圾(Other Waste)
- 塑料瓶(Plastic Bottle)
- 纸张(Paper)
YOLO 标注格式说明
YOLO 格式的标注文件是.txt 文件,每行代表一个物体,格式为:
<class_id> <x_center> <y_center> <width> <height>
示例:
0 0.45 0.32 0.12 0.15
2 0.67 0.81 0.08 0.10
数据集划分建议
推荐的数据集划分比例:
- 训练集:70%
- 验证集:15%
- 测试集:15%
代码实战
YOLO 格式验证脚本
import os
def validate_yolo_annotation(img_path, txt_path):
"""
验证 YOLO 标注文件的有效性
:param img_path: 图像文件路径
:param txt_path: 标注文件路径
"""
try:
# 检查图像文件是否存在
if not os.path.exists(img_path):
raise FileNotFoundError(f"图像文件 {img_path} 不存在")
# 检查标注文件是否存在
if not os.path.exists(txt_path):
raise FileNotFoundError(f"标注文件 {txt_path} 不存在")
# 读取标注文件
with open(txt_path, 'r') as f:
lines = f.readlines()
for line in lines:
parts = line.strip().split()
# 检查字段数量
if len(parts) != 5:
raise ValueError(f"标注行'{line}'格式错误,应有 5 个字段")
# 检查数值范围
class_id = int(parts[0])
x_center = float(parts[1])
y_center = float(parts[2])
width = float(parts[3])
height = float(parts[4])
if not (0 <= x_center <= 1):
raise ValueError(f"x_center {x_center} 超出范围 [0,1]")
if not (0 <= y_center <= 1):
raise ValueError(f"y_center {y_center} 超出范围 [0,1]")
# 更多验证逻辑...
except Exception as e:
print(f"验证失败: {str(e)}")
return False
return True
PyTorch 数据加载器
from torch.utils.data import Dataset
import cv2
import torch
class YOLODataset(Dataset):
def __init__(self, img_dir, label_dir, transform=None):
self.img_dir = img_dir
self.label_dir = label_dir
self.transform = transform
self.img_files = [f for f in os.listdir(img_dir) if f.endswith('.jpg')]
def __len__(self):
return len(self.img_files)
def __getitem__(self, idx):
img_name = self.img_files[idx]
img_path = os.path.join(self.img_dir, img_name)
label_path = os.path.join(self.label_dir, img_name.replace('.jpg', '.txt'))
# 读取图像
img = cv2.imread(img_path)
img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
# 读取标注
labels = []
if os.path.exists(label_path):
with open(label_path, 'r') as f:
for line in f.readlines():
class_id, x_center, y_center, width, height = map(float, line.strip().split())
labels.append([class_id, x_center, y_center, width, height])
# 应用数据增强
if self.transform:
img, labels = self.transform(img, labels)
return img, torch.tensor(labels)
损失函数配置
YOLO 模型常用的损失函数包含三个部分:
- 边界框坐标损失(Bounding Box Loss)
- 物体置信度损失(Objectness Loss)
- 类别损失(Class Loss)
关键参数:
- 边界框损失权重:通常设为 5.0
- 无物体置信度损失权重:通常设为 0.5
- 类别损失权重:通常设为 1.0
避坑指南
标注时的 3 个典型错误
- 误标遮挡物体:不要标注被严重遮挡且无法辨认类别的物体
- 边界框过大 / 过小:边界框应紧贴物体边缘
- 类别混淆:确保每个物体的类别标注准确
学习率调整策略
推荐使用余弦退火(Cosine Annealing)学习率调度器:
from torch.optim.lr_scheduler import CosineAnnealingLR
optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
scheduler = CosineAnnealingLR(optimizer, T_max=100)
模型部署尺寸压缩
- 使用模型量化(Quantization):将浮点权重转换为 8 位整数
- 应用剪枝(Pruning):移除不重要的神经元连接
- 使用更小的骨干网络(如 MobileNet 代替 ResNet)
延伸思考
改进方向
- 加入半监督学习(Semi-supervised Learning):利用未标注数据提升模型性能
- 多任务学习(Multi-task Learning):同时预测物体类别和材质
推荐开源工具
- Albumentations:强大的数据增强库
- LabelImg:简单易用的标注工具
- Roboflow:在线数据增强和标注平台
总结
通过本文,我们系统地介绍了从数据准备到模型训练的完整流程。在实际应用中,建议先从小的数据集开始实验,验证流程正确性后再扩大数据规模。垃圾分类是一个非常有价值的应用场景,希望本文能帮助开发者快速入门。
正文完
发表至: 未分类
近两天内
