基于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")
额外实用提示
- 训练用的数据增强:可以用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) ]) - 标注格式兼容:如果你的标注是VOC XML,可以用
torchvision.datasets.VOCDetection直接加载;如果是COCO JSON,用torchvision.datasets.CocoDetection。 - GPU内存优化:如果出现OOM(显存不足),可以调小
batch_size,或启用混合精度训练(torch.cuda.amp)。
内容来源于stack exchange
相关产品推荐
相关产品推荐

