零基础实战:基于DETR的行人检测模型训练全流程指南

行人检测作为计算机视觉的基础任务,在智能安防、客流分析、自动驾驶等领域有着广泛应用。传统方法如Faster R-CNN虽然成熟但配置复杂,而Facebook开源的DETR(Detection Transformer)通过Transformer架构实现了端到端目标检测,大幅简化了流程。本文将手把手带您完成从数据准备到模型部署的全过程,特别针对只有基础Python知识的开发者优化了操作步骤。

1. 环境配置与项目初始化

在开始之前,我们需要准备适合深度学习开发的环境。推荐使用Anaconda管理Python环境,它能有效解决依赖冲突问题。

conda create -n detr python=3.8 -y
conda activate detr
conda install pytorch torchvision torchaudio pytorch-cuda=11.7 -c pytorch -c nvidia

安装完PyTorch后,从GitHub获取DETR官方代码库:

git clone https://github.com/facebookresearch/detr.git
cd detr
pip install -r requirements.txt

常见问题处理:

  • 若pycocotools安装失败,可尝试:
    pip install pycocotools-windows
    
  • 遇到PanopticAPI错误时:
    pip install git+https://github.com/cocodataset/panopticapi.git
    

2. 数据准备:从VOC到COCO格式转换

DETR要求使用COCO格式的数据集,但很多现有数据是VOC格式。以下脚本可将VOC格式(XML标注)转换为COCO格式(JSON标注):

import os
import json
import xml.etree.ElementTree as ET
from tqdm import tqdm

def parse_voc_xml(xml_path):
    tree = ET.parse(xml_path)
    root = tree.getroot()
    
    filename = root.find('filename').text
    size = root.find('size')
    width = int(size.find('width').text)
    height = int(size.find('height').text)
    
    objects = []
    for obj in root.iter('object'):
        cls = obj.find('name').text
        bbox = obj.find('bndbox')
        xmin = float(bbox.find('xmin').text)
        ymin = float(bbox.find('ymin').text)
        xmax = float(bbox.find('xmax').text)
        ymax = float(bbox.find('ymax').text)
        objects.append({
            'class': cls,
            'bbox': [xmin, ymin, xmax-xmin, ymax-ymin]
        })
    return filename, width, height, objects

def convert_to_coco(voc_dir, output_json):
    images = []
    annotations = []
    categories = [{'id': 1, 'name': 'person'}]
    
    annotation_id = 1
    for idx, xml_file in enumerate(tqdm(os.listdir(voc_dir))):
        if not xml_file.endswith('.xml'):
            continue
            
        xml_path = os.path.join(voc_dir, xml_file)
        filename, width, height, objects = parse_voc_xml(xml_path)
        
        image_id = idx + 1
        images.append({
            'id': image_id,
            'file_name': filename,
            'width': width,
            'height': height
        })
        
        for obj in objects:
            annotations.append({
                'id': annotation_id,
                'image_id': image_id,
                'category_id': 1,
                'bbox': obj['bbox'],
                'area': obj['bbox'][2] * obj['bbox'][3],
                'iscrowd': 0
            })
            annotation_id += 1
    
    coco_dict = {
        'images': images,
        'annotations': annotations,
        'categories': categories
    }
    
    with open(output_json, 'w') as f:
        json.dump(coco_dict, f)

使用说明:

  1. 将所有VOC格式的XML文件和对应图片放入同一文件夹
  2. 运行脚本生成COCO格式的JSON标注文件
  3. 按COCO标准组织文件夹结构:
    /dataset
      /annotations
        instances_train.json
        instances_val.json
      /train2017
        *.jpg
      /val2017
        *.jpg
    

3. 模型配置与训练

DETR默认使用91类的COCO预训练权重,我们需要调整以适应单一类别(行人)检测:

import torch

# 调整预训练权重
pretrained = torch.load('detr-r50.pth')
pretrained['model']['class_embed.weight'].resize_(2, 256)  # 1类+背景
pretrained['model']['class_embed.bias'].resize_(2)
torch.save(pretrained, 'detr-r50-person.pth')

关键训练参数解析:

参数名 推荐值 说明
lr 1e-4 初始学习率
batch_size 4 根据GPU显存调整
epochs 50 训练轮次
lr_drop 40 学习率下降的epoch
weight_decay 1e-4 权重衰减系数

启动训练命令示例:

python main.py \
  --dataset_file "coco" \
  --coco_path "/path/to/your/coco_dataset" \
  --epochs 50 \
  --lr=1e-4 \
  --batch_size=4 \
  --num_workers=4 \
  --output_dir="output" \
  --resume="detr-r50-person.pth"

训练过程监控要点:

  • 关注loss变化曲线,正常情况应逐渐下降
  • mAP指标是主要评估标准
  • 如果显存不足,可减小batch_size或输入图像尺寸

4. 模型评估与可视化

训练完成后,使用以下脚本进行结果可视化:

import torch
import cv2
from models import build_model
from PIL import Image
import torchvision.transforms as T

# 加载训练好的模型
model, _, _ = build_model(args)
checkpoint = torch.load('output/checkpoint.pth', map_location='cpu')
model.load_state_dict(checkpoint['model'])
model.eval()

# 图像预处理
transform = T.Compose([
    T.Resize(800),
    T.ToTensor(),
    T.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])

def detect(image_path):
    img = Image.open(image_path)
    img_tensor = transform(img).unsqueeze(0)
    
    with torch.no_grad():
        outputs = model(img_tensor)
    
    # 解析输出结果
    probas = outputs['pred_logits'].softmax(-1)[0, :, :-1]
    keep = probas.max(-1).values > 0.7  # 置信度阈值
    
    # 绘制检测框
    img_cv = cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR)
    for p, (x, y, w, h) in zip(probas[keep], outputs['pred_boxes'][0, keep]):
        cv2.rectangle(img_cv, (int(x-w/2), int(y-h/2)), 
                     (int(x+w/2), int(y+h/2)), (0,255,0), 2)
        cv2.putText(img_cv, f'person {p[0]:.2f}',
                   (int(x-w/2), int(y-h/2)-10),
                   cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0,255,0), 2)
    
    cv2.imshow('Detection', img_cv)
    cv2.waitKey(0)

5. 性能优化与部署技巧

提升DETR在小数据集上表现的实用技巧:

  1. 数据增强策略

    • 随机水平翻转(p=0.5)
    • 颜色抖动(亮度、对比度、饱和度)
    • 随机裁剪(保持目标完整性)
  2. 模型微调建议

    • 冻结backbone的前几层参数
    • 使用更小的学习率(1e-5)微调Transformer层
    • 增加decoder层数提升检测精度
  3. 推理优化

    torchscript_model = torch.jit.script(model)  # 转换为TorchScript
    torchscript_model.save('detr_person.pt')  # 保存优化后模型
    

实际部署时,可以考虑使用Flask构建简单的API服务:

from flask import Flask, request, jsonify
import io
import torchvision.transforms as transforms

app = Flask(__name__)
model = torch.jit.load('detr_person.pt')

@app.route('/detect', methods=['POST'])
def detect_api():
    if 'image' not in request.files:
        return jsonify({'error': 'No image uploaded'}), 400
    
    image = request.files['image'].read()
    img = Image.open(io.BytesIO(image))
    img_tensor = transform(img).unsqueeze(0)
    
    with torch.no_grad():
        outputs = model(img_tensor)
    
    # 处理输出为JSON格式
    results = []
    for score, box in zip(outputs['pred_logits'], outputs['pred_boxes']):
        results.append({
            'score': score.item(),
            'bbox': box.tolist()
        })
    
    return jsonify({'detections': results})

遇到内存不足问题时,可以尝试以下解决方案:

  • 减小输入图像尺寸(如800→600)
  • 使用半精度浮点数(FP16)推理
  • 启用CUDA图形加速(CUDA Graphs)
Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐