如何基于自定义数据集微调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
相关产品推荐
相关产品推荐

