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

基于Torchvision预训练模型的目标检测模型训练及多图像训练定制方法咨询

基于Torchvision预训练模型的目标检测模型训练及多图像训练定制方法咨询

你好呀!你贴的代码其实是用预训练Faster RCNN模型做单图推理(预测)的流程,要改成多图像训练的话,需要完全调整整个 pipeline——因为训练和推理的逻辑差异很大:推理是用预训练模型直接输出预测结果,而训练需要喂给模型带标注的批量图像,并通过反向传播更新模型权重(一般是微调预训练模型,而非从头训练)。

下面我一步步给你拆解实现方法:


1. 核心概念先理清

目标检测模型的训练需要:

  • 批量的输入图像
  • 每个图像对应的标注信息:包括目标的边界框(xmin, ymin, xmax, ymax)和类别标签
  • 训练专用的数据管道(加载、增强、批量处理)
  • 训练循环(前向传播→计算损失→反向传播→更新权重)

2. 第一步:准备带标注的多图像数据集

首先你需要有带标注的数据集,标注格式可以是自定义JSON、VOC XML或COCO JSON。这里我们用自定义Dataset类来加载数据,适配Torchvision的训练要求:

import torch
from torch.utils.data import Dataset
from torchvision.io import read_image
import json

class CustomObjDetDataset(Dataset):
    def __init__(self, img_dir, annotation_path, transforms=None):
        self.img_dir = img_dir
        self.annotation_data = json.load(open(annotation_path, "r"))
        self.transforms = transforms

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

    def __getitem__(self, idx):
        # 1. 读取单张图像
        img_info = self.annotation_data[idx]
        img_path = f"{self.img_dir}/{img_info['image_name']}"
        image = read_image(img_path)  # 返回shape为(C, H, W)的tensor

        # 2. 读取对应标注(Torchvision要求格式)
        # 边界框:必须是(xmin, ymin, xmax, ymax)的float32 tensor
        boxes = torch.tensor(img_info["boxes"], dtype=torch.float32)
        # 类别标签:Torchvision要求从1开始(0为背景类)
        labels = torch.tensor(img_info["labels"], dtype=torch.int64)
        
        # 3. 打包成模型需要的target字典(必须包含boxes和labels键)
        target = {
            "boxes": boxes,
            "labels": labels,
            # 可选:添加area和iscrowd,部分模型会用到
            "area": (boxes[:, 3] - boxes[:, 1]) * (boxes[:, 2] - boxes[:, 0]),
            "iscrowd": torch.zeros_like(labels, dtype=torch.int64)
        }

        # 4. 应用数据增强/变换(训练时常用)
        if self.transforms:
            image, target = self.transforms(image, target)

        return image, target

如果你用的是VOC/COCO标准数据集,可以直接用Torchvision内置的VOCDetection/CocoDetection类,不用自己写Dataset。


3. 第二步:用DataLoader批量加载多图像

因为每个图像的目标数量不同,Torchvision要求用自定义collate_fn来处理批量数据:

from torch.utils.data import DataLoader

# 自定义collate_fn:把每个样本的图像和标注分别打包成列表
def custom_collate_fn(batch):
    return tuple(zip(*batch))

# 实例化数据集和数据加载器
train_dataset = CustomObjDetDataset(
    img_dir="path/to/your/train_images",
    annotation_path="path/to/your/train_annotations.json",
    transforms=your_training_transforms  # 后面会讲训练用的变换
)

train_dataloader = DataLoader(
    train_dataset,
    batch_size=4,  # 一次喂4张图,可根据GPU内存调整
    shuffle=True,  # 训练时打乱数据
    num_workers=4,  # 多进程加载数据
    collate_fn=custom_collate_fn
)

4. 第三步:准备可训练的预训练模型

如果你要微调的数据集类别数和预训练模型(COCO 91类)不同,需要替换模型的分类头:

from torchvision.models.detection import fasterrcnn_resnet50_fpn_v2
from torchvision.models.detection.faster_rcnn import FastRCNNPredictor

# 加载预训练模型
model = fasterrcnn_resnet50_fpn_v2(weights="DEFAULT")

# 替换分类头:假设你的数据集有10个目标类别(+1个背景类,总共11类)
num_classes = 11
in_features = model.roi_heads.box_predictor.cls_score.in_features
model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes)

5. 第四步:设置训练优化器与调度器

目标检测常用SGD优化器,配合学习率调度器:

import torch.optim as optim

# 仅训练可学习的参数
params = [p for p in model.parameters() if p.requires_grad]
optimizer = optim.SGD(
    params,
    lr=0.005,
    momentum=0.9,
    weight_decay=0.0005
)

# 学习率调度器:每3轮学习率减半
lr_scheduler = optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.5)

6. 第五步:多图像训练循环

最后编写训练循环,把模型切换到训练模式,批量喂入数据:

import time

# 选择训练设备(GPU优先)
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
model.to(device)

# 训练轮数
num_epochs = 10

for epoch in range(num_epochs):
    model.train()
    total_loss = 0.0
    start_time = time.time()

    # 遍历批量数据
    for images, targets in train_dataloader:
        # 把图像和标注移到对应设备
        images = [img.to(device) for img in images]
        targets = [{k: v.to(device) for k, v in t.items()} for t in targets]

        # 前向传播:模型自动计算损失
        loss_dict = model(images, targets)
        losses = sum(loss for loss in loss_dict.values())
        total_loss += losses.item()

        # 反向传播与权重更新
        optimizer.zero_grad()
        losses.backward()
        optimizer.step()

    # 更新学习率
    lr_scheduler.step()

    # 打印训练日志
    avg_loss = total_loss / len(train_dataloader)
    print(f"Epoch [{epoch+1}/{num_epochs}] | Avg Loss: {avg_loss:.4f} | Time: {time.time()-start_time:.2f}s")

# 保存训练后的模型
torch.save(model.state_dict(), "faster_rcnn_finetuned.pth")

额外实用提示

  1. 训练用的数据增强:可以用Torchvision的torchvision.transforms.v2模块,它能同步变换图像和边界框(比如随机翻转、裁剪),示例:
    from torchvision.transforms import v2 as transforms
    
    your_training_transforms = transforms.Compose([
        transforms.RandomHorizontalFlip(p=0.5),
        transforms.ColorJitter(brightness=0.2, contrast=0.2),
        transforms.ToDtype(torch.float32, scale=True)
    ])
    
  2. 标注格式兼容:如果你的标注是VOC XML,可以用torchvision.datasets.VOCDetection直接加载;如果是COCO JSON,用torchvision.datasets.CocoDetection。
  3. GPU内存优化:如果出现OOM(显存不足),可以调小batch_size,或启用混合精度训练(torch.cuda.amp)。

内容来源于stack exchange

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.07 10:20:31