当前位置:首页 > 随笔 > 正文内容

RT-DETR目标检测模型:自定义数据训练与推理全流程解析

访客 随笔 2026年8月20日 1

RT-DETR(Real-Time Detection Transformer)是由百度提出的一种高效实时目标检测架构。该模型在继承 DETR 核心编解码器设计的基础上,通过深度优化实现了精度与推理速度的卓越平衡,在自动驾驶及智能监控等领域具备显著的工程价值。

RT-DETR Architecture

Ultralytics 框架已原生集成了 RT-DETR 的实现。本文将详细解析如何基于该框架,完成从自定义数据构建到模型训练、验证及推理的全流程。

核心架构特性

  • 高效混合编码器 (Hybrid Encoder):创新性地解耦了尺度内特征交互与跨尺度特征融合。通过基于注意力的 AIFI(尺度内特征交互)模块处理特征,并利用基于 CNN 的 CCFM(跨尺度特征融合)模块,大幅降低了多尺度处理的计算复杂度。
  • IOU 感知查询选择 (IOU-aware Query Selection):优化了目标查询的初始化机制。该方法联合建模编码器特征的潜在变量,显式构建认知不确定性,从而为解码器筛选出兼具高分类置信度与高 IOU 分数的初始查询向量。
  • 计算资源优化:通过缩减特征图尺寸及注意力头数量,有效控制了模型参数量,同时引入分组注意力机制进一步提升特征表达效能。

环境初始化

在确保本地计算节点已正确安装 CUDA 与 cuDNN 的前提下,建议创建独立的 Python 虚拟环境,并安装 Ultralytics 核心依赖库:

pip install ultralytics

数据集构建与划分

RT-DETR 依赖 YOLO 格式的标注数据,即每张图片对应一个同名的 .txt 文件,内容格式为 class_id x_center y_center width height(边界框坐标均已归一化至 0-1 之间)。

Dataset Directory

以下脚本利用面向对象的思想与 pathlib 库重构了数据划分逻辑,将原始数据按 8:1:1 的比例随机分配至训练集、验证集和测试集,并自动生成对应的索引文件。

import shutil
import random
from pathlib import Path

class DatasetSplitter:
    def __init__(self, img_src, lbl_src, out_base, split_ratio=(0.8, 0.1, 0.1)):
        self.img_src = Path(img_src)
        self.lbl_src = Path(lbl_src)
        self.out_base = Path(out_base)
        self.ratios = split_ratio

    def _prepare_directories(self):
        for subset in ['train', 'val', 'test']:
            (self.out_base / 'images' / subset).mkdir(parents=True, exist_ok=True)
            (self.out_base / 'labels' / subset).mkdir(parents=True, exist_ok=True)

    def execute(self):
        self._prepare_directories()
        # 获取所有标签文件并提取文件名
        all_files = [f.stem for f in self.lbl_src.glob('*.txt')]
        random.shuffle(all_files)

        total = len(all_files)
        train_end = int(total * self.ratios[0])
        val_end = train_end + int(total * self.ratios[1])

        # 使用切片进行数据集划分
        splits = {
            'train': all_files[:train_end],
            'val': all_files[train_end:val_end],
            'test': all_files[val_end:]
        }

        for subset, files in splits.items():
            index_file = self.out_base / f'{subset}.txt'
            with open(index_file, 'w') as f:
                for name in files:
                    img_file = self.img_src / f'{name}.jpg'
                    lbl_file = self.lbl_src / f'{name}.txt'
                    
                    dst_img = self.out_base / 'images' / subset / f'{name}.jpg'
                    dst_lbl = self.out_base / 'labels' / subset / f'{name}.txt'
                    
                    if img_file.exists() and lbl_file.exists():
                        shutil.copy(img_file, dst_img)
                        shutil.copy(lbl_file, dst_lbl)
                        f.write(f"{dst_img.absolute()}\n")

if __name__ == '__main__':
    splitter = DatasetSplitter('raw_data/images', 'raw_data/labels', 'processed_datasets')
    splitter.execute()

执行后,将在目标目录下生成结构化的 imageslabels 子目录,以及用于索引的 .txt 文件。

Processed Dataset Structure

配置文件设定

1. 数据集描述文件

在项目根目录创建 custom_data.yaml,明确数据路径及类别映射关系:

path: ../processed_datasets
train: train.txt
val: val.txt
test: test.txt

names:
  0: target_object

2. 模型结构文件

复制框架自带的 ultralytics/cfg/models/rt-detr/rtdetr-l.yaml 并重命名为 rtdetr_custom.yaml。核心操作是将顶部的类别参数 nc 修改为实际任务的类别数量。

3. 训练超参数

Ultralytics 提供了灵活的参数控制接口,关键参数解析如下:

  • epochs:模型训练的迭代周期总数。
  • batch:批处理大小,需根据 GPU 显存容量动态调整以避免 OOM 异常,推荐设为 2 的幂次方。
  • imgsz:输入图像的分辨率,要求必须为 32 的整数倍。
  • workers:数据加载进程数,若系统内存受限触发异常,可将其降级为 0。

模型训练与推理工程化实践

为了提升代码的复用性与可维护性,以下示例将训练、验证与推理流程封装为一个统一的管线类,并使用字典解包的方式传递超参数。

from ultralytics import RTDETR
import logging

logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')

class DetectionPipeline:
    def __init__(self, model_config='rtdetr_custom.yaml'):
        self.detector = RTDETR(model_config)

    def train(self, data_config, **kwargs):
        default_train_cfg = {
            'epochs': 200,
            'imgsz': 640,
            'batch': 16,
            'device': 0,
            'workers': 8,
            'optimizer': 'AdamW',
            'amp': False
        }
        default_train_cfg.update(kwargs)
        logging.info(f"Initiating training phase...")
        self.detector.train(data=data_config, **default_train_cfg)

    def evaluate(self, weights_path, data_config, **kwargs):
        self.detector = RTDETR(weights_path)
        eval_cfg = {
            'batch': 32,
            'imgsz': 640,
            'device': 0,
            'split': 'test'
        }
        eval_cfg.update(kwargs)
        logging.info("Starting model evaluation...")
        self.detector.val(data=data_config, **eval_cfg)

    def run_inference(self, weights_path, source_dir, **kwargs):
        self.detector = RTDETR(weights_path)
        infer_cfg = {
            'imgsz': 640,
            'device': 0,
            'save': True
        }
        infer_cfg.update(kwargs)
        logging.info(f"Running inference on {source_dir}...")
        self.detector.predict(source=source_dir, **infer_cfg)

if __name__ == '__main__':
    pipeline = DetectionPipeline('ultralytics/cfg/models/rt-detr/rtdetr-l.yaml')
    
    # 执行训练
    pipeline.train('custom_data.yaml', epochs=100, batch=8)
    
    # 加载最优权重进行验证与推理
    best_model = 'runs/detect/train/weights/best.pt'
    pipeline.evaluate(best_model, 'custom_data.yaml')
    pipeline.run_inference(best_model, 'sample_images/')

训练任务完成后,系统会自动将表现最优的模型权重保存至 runs/detect/train/weights/best.pt 路径下,供后续部署调用。

相关文章

可以按小时收费的VPS

很多 VPS 提供商都支持 按小时计费(hourly billing),想短期试用 / 临时搭建节点、测试网络、短期项目等场景非常合适。下面是当前最主流且靠谱的按小时 VPS 选项,分别按不同需求场景整理: 1. Vultr(全球节点,包括日本) 按小时计费 可选机房:东京 / 大阪 / 洛杉矶 / 法兰克福 / 伦敦 … 支持 PayPal(部分情况),但更常用信用卡/PayPal+卡价格参考$...

在 iPhone 上下载国外App

地区/国家限制App Store 会根据 Apple ID 的国家或地区限制应用下载。如果你的 Apple ID 绑定的是中国大陆,就可能无法下载 OpenAI 官方的 ChatGPT 应用,因为它在大陆 App Store 不上架。解决办法:换成美国、加拿大、香港等地区的 Apple ID。或者在现有 Apple ID 上更改地区。注册一个国外 Apple ID(推荐)比如注册 美国区 Appl...

Node.js 中的异步编程:回调与 Promise

Node.js 是一个基于 JavaScript 构建的单线程、非阻塞运行环境,它通过异步编程机制来高效处理多个操作。在执行如文件读取、API 请求或数据库查询等任务时,Node.js 不会等待这些操作完成,而是使用回调函数和 Promise 来避免阻塞主线程。 回调方式实现异步 那么当异步操作完成后,Node.js 如何知道接下来要做什么呢?这就要用到 回调函数(callback)。 回调本质上...

MariaDB Galera集群故障快速恢复指南

OpenStack控制节点采用三节点MariaDB Galera集群架构。当数据库集群因故障重启时,有时会出现Galera集群无法正常启动的问题。虽然有多种方法可以恢复数据库服务,但如何实现快速启动同时确保数据完整性呢? 通过分析日志发现,MariaDB Galera集群节点宕机时会在日志中输出以下信息: [Note] WSREP: 新集群视图:全局状态: 874d8e7e-5980-11e8-8...

Android 中 EventBus 的通信机制与实现原理深度解析

EventBus 核心设计思想 EventBus 是一个基于观察者模式的事件总线框架,广泛应用于 Android 平台以实现组件解耦。它通过中心化的消息分发机制,使不同层级、不同线程的对象能够以"发布-订阅"方式通信,避免了传统接口回调或广播带来的强依赖问题。 核心角色说明 事件(Event):任意 Java 对象,作为数据载体,如网络状态变更通知、用户登录信息等。 发布者(Publi...

二叉树基础操作实现(C语言)

二叉树基础操作实现(C语言)

二叉树遍历方法 以下为二叉树结构示例 前序访问实现: void traversePreOrder(TreeNode* node) { if (node == NULL) { printf("N "); return; } printf("%d ", node->value); trave...

发表评论

访客

◎欢迎参与讨论,请在这里发表您的看法和观点。