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

基于提取特征训练YOLOv5m模型:多相机检测特征输入训练方法咨询

基于YOLOv5中间特征替代图像输入的训练方案

一、数据对齐与验证

  • 确保每个特征文件和原图像的标注文件严格一一对应(比如用相同文件名后缀替换,如img1.jpg对应img1_feat.pt和img1.txt),避免训练时标签错配。
  • 加载任意几个特征文件验证维度:用torch.load()或numpy.load()读取后检查shape是否为(768, 20, 20),同时确认特征的 dtype 为float32(和YOLOv5训练时一致),避免存储时的维度或精度错乱。

二、模型输入层与结构调整

结合你具备修改多输入模型的能力,核心调整点如下:

  1. 截断原模型:保留YOLOv5m第23层之后的所有网络(从第24层到检测头),将这部分作为新模型的主体。
  2. 适配特征输入:新模型的输入层直接接受(batch_size, 768, 20, 20)的特征张量,无需保留原模型的前23层(因为你已提前提取这部分的输出)。
  3. 多相机融合模块(可选):如果是多相机输入,可在模型最前端加一个融合模块,示例代码如下:
import torch.nn as nn
import torch

class MultiCamFusion(nn.Module):
    def __init__(self, in_channels=768, num_cams=2):
        super().__init__()
        # 简单注意力融合,可根据需求替换为concat/直接相加等方式
        self.attn = nn.Sequential(
            nn.AdaptiveAvgPool2d(1),
            nn.Conv2d(in_channels * num_cams, num_cams, kernel_size=1),
            nn.Softmax(dim=1)
        )
        self.num_cams = num_cams
        self.in_ch = in_channels

    def forward(self, cam_features):
        # cam_features: 列表,每个元素是单相机特征 [B,768,20,20]
        concat_feat = torch.cat(cam_features, dim=1)  # [B, 768*N,20,20]
        weights = self.attn(concat_feat)  # [B,N,1,1]
        # 拆分特征并加权融合
        split_feats = torch.split(concat_feat, self.in_ch, dim=1)
        fused = 0.0
        for i in range(self.num_cams):
            fused += split_feats[i] * weights[:, i:i+1, :, :]
        return fused  # [B,768,20,20]

将这个模块的输出接入原模型第24层的输入即可。

三、自定义数据加载器

构建PyTorch Dataset类,核心是加载特征和对应标签:

import os
import torch
from torch.utils.data import Dataset

class FeatDetDataset(Dataset):
    def __init__(self, feat_dir, label_dir):
        self.feat_paths = [os.path.join(feat_dir, f) for f in os.listdir(feat_dir) if f.endswith('.pt')]
        self.label_dir = label_dir

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

    def __getitem__(self, idx):
        # 加载特征
        feat = torch.load(self.feat_paths[idx])
        # 加载对应标签
        base_name = os.path.basename(self.feat_paths[idx]).replace('_feat.pt', '.txt')
        label_path = os.path.join(self.label_dir, base_name)
        labels = []
        if os.path.exists(label_path):
            with open(label_path, 'r') as f:
                for line in f.readlines():
                    cls, x, y, w, h = map(float, line.strip().split())
                    labels.append([cls, x, y, w, h])
        # 转换为tensor,YOLO标签格式无需修改(特征尺度20×20对应原图640×640,检测头已适配)
        labels = torch.tensor(labels) if labels else torch.empty((0,5))
        return feat, labels

之后用DataLoader封装,注意根据显存调整batch size(特征单张大小约3MB,比图像略大,可适当减小batch size)。

四、训练配置调整

  • 学习率:基于高层特征训练的任务难度比图像输入低,建议初始学习率设为原图像训练的1/101/5(比如原LR=0.01,现在用0.0010.002),避免震荡。
  • 损失函数:直接沿用YOLOv5的原生损失函数即可,检测头的输出格式和标签格式完全匹配。
  • 正则化:可在融合模块后加一层nn.Dropout(p=0.1),防止过拟合(高层特征泛化性相对弱一些)。
  • 训练模式:确保模型处于train()模式,注意如果原模型第23层之后有BatchNorm层,训练时要让其正常更新均值和方差(提取特征时用eval()模式,避免干扰)。

五、训练与验证

  • 训练阶段:监控box loss、obj loss、cls loss的下降趋势,和原图像训练的曲线对比,若损失下降缓慢,可微调学习率或优化器(比如用AdamW替代SGD)。
  • 验证阶段:先提取验证集所有图像的第23层特征,用训练好的模型推理,计算mAP指标,验证效果是否符合预期。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 11:30:47