CenterNet实战指南:从零开始训练自定义数据集

1次阅读
没有评论

共计 2208 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

背景介绍

CenterNet 是一种基于关键点检测的目标检测方法,其核心思想是将目标检测任务转化为关键点(通常是目标中心点)的预测问题。相比传统的两阶段检测器(如 Faster R-CNN)和单阶段检测器(如 YOLO、SSD),CenterNet 具有以下优势:

CenterNet 实战指南:从零开始训练自定义数据集

  • 简单高效:模型结构简洁,不需要复杂的 anchor 设计和后处理
  • 精度高:在 COCO 等基准数据集上达到 SOTA 性能
  • 灵活性强:可以方便地扩展到其他任务如姿态估计、3D 检测等

痛点分析

在训练自定义数据集时,初学者常会遇到以下问题:

  1. 标注格式转换:公开数据集标注格式(如 COCO、VOC)与模型输入要求不匹配
  2. 小目标检测:默认配置对小目标(<32×32 像素)检测效果差
  3. 训练不稳定:损失值出现 NaN 或波动剧烈
  4. 低 mAP 问题:模型收敛但验证集指标不理想
  5. 推理速度慢:实际部署时无法达到预期帧率

技术实现

数据集准备

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()

避坑指南

  1. NaN 损失问题
  2. 检查学习率是否过大(建议初始值 1e-4~5e-4)
  3. 添加梯度裁剪:torch.nn.utils.clip_grad_norm_(model.parameters(), 0.1)

  4. 低 mAP 解决方案

  5. 验证标注质量(推荐使用 LabelImg 检查)
  6. 调整损失权重(如小目标多的场景增加 wh_weight)
  7. 尝试更大 backbone(ResNet101→DLA-34)

  8. 显存不足处理

  9. 减小 batch_size(最低可到 8)
  10. 使用 --not_cuda_benchmark 禁用 cudnn 自动优化

  11. 训练震荡

  12. 启用 warmup 学习率:--lr_step 10,20 --lr_gamma 0.1
  13. 添加数据增强:--flip 0.5 --scale 0.4

  14. 推理速度优化

  15. 使用 TensorRT 加速(可提速 2 - 3 倍)
  16. 降低输入分辨率(如从 512×512→384×384)

性能优化

  1. 多尺度测试技巧

    python test.py ctdet \
    --exp_id your_experiment \
    --keep_res \
    --flip_test \
    --test_scales 0.5,0.75,1.0

  2. 模型量化

    model = torch.quantization.quantize_dynamic(model, {torch.nn.Conv2d}, dtype=torch.qint8
    )

延伸思考

  1. 如何改进 CenterNet 处理密集目标(如人群计数场景)?
  2. 当标注存在大量遮挡目标时,应该调整哪些参数?
  3. 如何设计更适合小目标检测的 head 结构?

训练曲线示例(TensorBoard):

# 启动监控
tensorboard --logdir=exp/your_experiment

经过一周的调参实战,我的自定义数据集 mAP@0.5 从 0.42 提升到了 0.68。最关键的是找到了合适的学习率衰减策略,并在第 3 个 epoch 后启用了更强的数据增强。建议新手一定要耐心做消融实验,CenterNet 对超参数其实相当敏感。

正文完
 0
评论(没有评论)