You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何基于自定义数据集微调DiT(Transformer骨干)实现目标检测?

解决DiT作为目标检测骨干的微调问题及错误修复

一、修复当前DETR+DiT的KeyError问题

1. 检查数据集输出格式

你的KeyError: 'labels'本质是数据集返回结构与collate_fn逻辑不匹配。确保数据集类的__getitem__返回**(原始图像, 目标标签字典)**,而非经过特征提取器处理后的结果:

class CustomDataset(Dataset):
    def __init__(self, img_dir, ann_file):
        self.img_dir = img_dir
        self.annotations = json.load(open(ann_file))

    def __len__(self):
        return len(self.annotations)

    def __getitem__(self, idx):
        # 加载原始PIL图像
        img_path = os.path.join(self.img_dir, self.annotations[idx]["image_id"])
        image = Image.open(img_path).convert("RGB")
        # 构建符合DETR要求的目标标签
        target = {
            "boxes": torch.tensor(self.annotations[idx]["boxes"], dtype=torch.float32),
            "labels": torch.tensor(self.annotations[idx]["labels"], dtype=torch.int64)
        }
        return image, target

2. 验证collate_fn逻辑

现有collate_fn逻辑正确,只要数据集返回(图像, 标签)元组,就能正常生成包含pixel_values、pixel_mask、labels的batch,不会触发KeyError。


二、将DiT作为骨干接入目标检测框架的通用方案

方案1:基于Hugging Face Transformers接入DETR

DiT是BEiT v3的变种,可封装为DETR兼容的骨干网络:

步骤1:封装DiT骨干

from transformers import BeitModel
import torch.nn as nn

class DiTBackbone(nn.Module):
    def __init__(self, model_name="microsoft/dit-large", out_channels=256):
        super().__init__()
        self.dit = BeitModel.from_pretrained(model_name)
        # 投影到DETR要求的输出通道数
        self.proj = nn.Conv2d(self.dit.config.hidden_size, out_channels, kernel_size=1)

    def forward(self, pixel_values, pixel_mask=None):
        # 获取DiT特征并移除CLS token
        outputs = self.dit(pixel_values=pixel_values, attention_mask=pixel_mask)
        last_hidden_state = outputs.last_hidden_state[:, 1:, :]
        # 将序列特征转为2D特征图(输入224x224对应14x14特征图)
        batch_size, seq_len, hidden_size = last_hidden_state.shape
        h = w = int(seq_len ** 0.5)
        feature_map = last_hidden_state.permute(0, 2, 1).view(batch_size, hidden_size, h, w)
        # 投影到目标通道数
        return {"last_hidden_state": self.proj(feature_map)}

步骤2:替换DETR骨干并初始化

from transformers import DetrForObjectDetection, DetrConfig

# 初始化DETR配置,禁用默认骨干
detr_config = DetrConfig.from_pretrained("facebook/detr-resnet-50")
detr_config.num_labels = 你的类别数  # 如文本、图表等类别总数
detr_config.backbone_config = None

model = DetrForObjectDetection(detr_config)
# 替换为自定义DiT骨干
model.backbone = DiTBackbone(model_name="microsoft/dit-large")
# 可选:加载DETR头部预训练权重
pretrained_detr = DetrForObjectDetection.from_pretrained("facebook/detr-resnet-50")
model.class_labels_classifier.load_state_dict(pretrained_detr.class_labels_classifier.state_dict())
model.bbox_predictor.load_state_dict(pretrained_detr.bbox_predictor.state_dict())

步骤3:启动训练

使用标准DETR训练流程,确保数据加载器输出符合要求即可。


方案2:基于Detectron2接入DiT作为骨干

Detectron2支持自定义骨干,可通过注册类实现:

步骤1:注册DiT骨干

from detectron2.modeling.backbone import Backbone, ShapeSpec
from detectron2.modeling.backbone.build import BACKBONE_REGISTRY
from transformers import BeitModel
import torch.nn as nn

@BACKBONE_REGISTRY.register()
class DiTBackboneDetectron2(Backbone):
    def __init__(self, cfg, input_shape):
        super().__init__()
        self.dit = BeitModel.from_pretrained("microsoft/dit-large")
        self.out_channels = cfg.MODEL.DiT.OUT_CHANNELS
        self.proj = nn.Conv2d(self.dit.config.hidden_size, self.out_channels, kernel_size=1)
        # 定义输出特征的下采样率和名称
        self._out_feature_strides = {"res5": 16}
        self._out_features = ["res5"]

    def forward(self, x):
        # 处理输入图像,生成特征图
        outputs = self.dit(pixel_values=x)
        last_hidden_state = outputs.last_hidden_state[:, 1:, :]
        batch_size, seq_len, hidden_size = last_hidden_state.shape
        h = w = int(seq_len ** 0.5)
        feature_map = last_hidden_state.permute(0, 2, 1).view(batch_size, hidden_size, h, w)
        return {"res5": self.proj(feature_map)}

    def output_shape(self):
        return {
            "res5": ShapeSpec(
                channels=self.out_channels, stride=self._out_feature_strides["res5"]
            )
        }

步骤2:修改Detectron2配置文件

创建自定义yaml配置:

_BASE_: "Base-RCNN-FPN.yaml"
MODEL:
  BACKBONE:
    NAME: "DiTBackboneDetectron2"
  DiT:
    OUT_CHANNELS: 256
  ROI_HEADS:
    NUM_CLASSES: 你的类别数
INPUT:
  MIN_SIZE_TRAIN: (224,)
  MAX_SIZE_TRAIN: 224
  MIN_SIZE_TEST: 224
  MAX_SIZE_TEST: 224

步骤3:启动训练

使用Detectron2标准训练流程:

from detectron2.engine import DefaultTrainer
from detectron2.config import get_cfg
import os

cfg = get_cfg()
cfg.merge_from_file("path/to/your/config.yaml")
cfg.DATASETS.TRAIN = ("your_train_dataset",)
cfg.DATASETS.TEST = ("your_val_dataset",)
cfg.SOLVER.IMS_PER_BATCH = 4
cfg.SOLVER.BASE_LR = 1e-4
cfg.OUTPUT_DIR = "./dit_detection_output"
os.makedirs(cfg.OUTPUT_DIR, exist_ok=True)

trainer = DefaultTrainer(cfg)
trainer.resume_or_load(resume=False)
trainer.train()

关键注意事项

  • DiT默认输入尺寸为224x224,训练时需统一数据集图像尺寸,或修改特征图转换逻辑适配可变尺寸。
  • 建议先冻结DiT骨干前几层,训练检测头部,再逐步解冻微调,避免过拟合。
  • 标签格式需严格匹配框架要求:DETR要求boxes为xyxy格式;Detectron2需符合COCO或自定义注册的数据集格式。

内容的提问来源于stack exchange,提问作者Deshwal

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.21 05:36:12