共计 2208 个字符,预计需要花费 6 分钟才能阅读完成。
背景介绍
CenterNet 是一种基于关键点检测的目标检测方法,其核心思想是将目标检测任务转化为关键点(通常是目标中心点)的预测问题。相比传统的两阶段检测器(如 Faster R-CNN)和单阶段检测器(如 YOLO、SSD),CenterNet 具有以下优势:

- 简单高效:模型结构简洁,不需要复杂的 anchor 设计和后处理
- 精度高:在 COCO 等基准数据集上达到 SOTA 性能
- 灵活性强:可以方便地扩展到其他任务如姿态估计、3D 检测等
痛点分析
在训练自定义数据集时,初学者常会遇到以下问题:
- 标注格式转换:公开数据集标注格式(如 COCO、VOC)与模型输入要求不匹配
- 小目标检测:默认配置对小目标(<32×32 像素)检测效果差
- 训练不稳定:损失值出现 NaN 或波动剧烈
- 低 mAP 问题:模型收敛但验证集指标不理想
- 推理速度慢:实际部署时无法达到预期帧率
技术实现
数据集准备
CenterNet 官方支持 COCO 格式数据集。如果你的数据是 VOC 格式,可以使用以下转换脚本:
import xml.etree.ElementTree as ET
import json
def voc_to_coco(voc_dir, output_path):
# 初始化 COCO 数据结构
coco = {"images": [],
"annotations": [],
"categories": [{"id": 1, "name": "your_class"}]
}
# 遍历 VOC 标注文件
for xml_file in Path(voc_dir).glob('*.xml'):
tree = ET.parse(xml_file)
# 解析逻辑...
# 完整代码见 GitHub 仓库
with open(output_path, 'w') as f:
json.dump(coco, f)
模型配置
关键参数说明(src/lib/opts.py):
head_conv: 输出头卷积通道数(默认 64)hm_weight: 中心点热图损失权重(默认 1)wh_weight: 宽高回归损失权重(默认 0.1)off_weight: 偏移量损失权重(默认 1)
推荐初始配置:
python main.py ctdet \
--exp_id your_experiment \
--arch dla_34 \
--batch_size 32 \
--master_batch 16 \
--lr 1.25e-4 \
--gpus 0,1
训练流程
完整训练代码示例:
from models.model import create_model
from datasets.dataset_factory import get_dataset
from trains.train_factory import train_factory
# 1. 初始化模型
model = create_model('dla_34', {'hm': num_classes, 'wh': 2, 'reg': 2})
# 2. 加载数据
train_dataset = get_dataset('coco', '/path/to/data', 'train')
# 3. 配置训练器
trainer = train_factory['ctdet'](
opt=opt, # 包含所有超参数
model=model,
optimizer=optimizer
)
# 4. 开始训练
for epoch in range(opt.num_epochs):
trainer.train(epoch)
trainer.save_model()
避坑指南
- NaN 损失问题
- 检查学习率是否过大(建议初始值 1e-4~5e-4)
-
添加梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), 0.1) -
低 mAP 解决方案
- 验证标注质量(推荐使用 LabelImg 检查)
- 调整损失权重(如小目标多的场景增加 wh_weight)
-
尝试更大 backbone(ResNet101→DLA-34)
-
显存不足处理
- 减小 batch_size(最低可到 8)
-
使用
--not_cuda_benchmark禁用 cudnn 自动优化 -
训练震荡
- 启用 warmup 学习率:
--lr_step 10,20 --lr_gamma 0.1 -
添加数据增强:
--flip 0.5 --scale 0.4 -
推理速度优化
- 使用 TensorRT 加速(可提速 2 - 3 倍)
- 降低输入分辨率(如从 512×512→384×384)
性能优化
-
多尺度测试技巧
python test.py ctdet \ --exp_id your_experiment \ --keep_res \ --flip_test \ --test_scales 0.5,0.75,1.0 -
模型量化
model = torch.quantization.quantize_dynamic(model, {torch.nn.Conv2d}, dtype=torch.qint8 )
延伸思考
- 如何改进 CenterNet 处理密集目标(如人群计数场景)?
- 当标注存在大量遮挡目标时,应该调整哪些参数?
- 如何设计更适合小目标检测的 head 结构?
训练曲线示例(TensorBoard):
# 启动监控
tensorboard --logdir=exp/your_experiment
经过一周的调参实战,我的自定义数据集 mAP@0.5 从 0.42 提升到了 0.68。最关键的是找到了合适的学习率衰减策略,并在第 3 个 epoch 后启用了更强的数据增强。建议新手一定要耐心做消融实验,CenterNet 对超参数其实相当敏感。
正文完
