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

能否修改YOLOv8作为其他任务的特征提取器并实现多任务学习?

基于YOLOv8构建多任务自定义模型的实现方案

需求背景

我正在阅读YOLOv8官方文档,但没找到实现以下需求的简便方法:

  • 加载预训练YOLOv8模型,将其作为子模块构建更大的自定义模型
  • 修改YOLOv8的前向传播,同时输出目标检测损失和卷积特征,供后续自定义任务层使用
  • 实现多任务学习:基于YOLOv8特征添加额外预测层,同时保留YOLOv8的检测损失

YOLOv8原生Python用法示例:

from ultralytics import YOLO

# 从零创建YOLO模型
model = YOLO('yolov8n.yaml')

# 加载预训练模型(训练推荐)
model = YOLO('yolov8n.pt')

# 用coco128数据集训练3轮
results = model.train(data='coco128.yaml', epochs=3)

# 在验证集上评估性能
results = model.val()

# 检测图片
results = model('https://ultralytics.com/images/bus.jpg')

# 导出为ONNX格式
success = model.export(format='onnx')

我期望的自定义模型大致结构(示例代码无法直接运行):

import torch
import torch.nn as nn
from ultralytics import YOLO

class Yolov8Wrapper(nn.Module):
    
    def __init__(self, yolov8_feature_dim, n1, n2, n3):
        super().__init__()
        self.yolov8 = YOLO('yolov8n.pt')
        self.fc1 = nn.Linear(yolov8_feature_dim, n1)
        self.fc2 = nn.Linear(yolov8_feature_dim, n2)
        self.fc3 = nn.Linear(yolov8_feature_dim, n3)
    
    def forward(self, images, gt_boxes):
        features, loss = self.yolov8(images, gt_boxes)
        logits1 = self.fc1(features)
        logits2 = self.fc2(features)
        logits3 = self.fc3(features)
        return {
            'logits1': logits1,
            'logits2': logits2,
            'logits3': logits3,
            'yolov8_loss': loss,
        }

注:构建自定义包装器后,无法使用YOLO自带的训练/验证/预测功能,需要编写自定义数据加载器,提供YOLOv8所需输入和自定义任务的额外数据。


可行性与实现方案

可行性结论

完全可行。Ultralytics YOLOv8的代码架构具备足够灵活性,支持提取中间特征、自定义前向传播逻辑,以及多任务损失组合。

具体实现步骤

1. 正确加载YOLOv8的PyTorch核心模型

YOLO类是高层API封装,直接作为nn.Module子模块会有问题,需提取其内部的核心模型:

from ultralytics import YOLO

yolo_api = YOLO('yolov8n.pt')
# 提取底层PyTorch模型(真正的nn.Module实例)
yolo_core = yolo_api.model

2. 修改YOLOv8前向传播,输出特征与损失

YOLOv8的核心模型在训练模式下返回损失,推理模式返回检测结果。可以通过两种方式提取中间特征:

方法一:重写前向传播函数

直接继承YOLOv8的检测模型类,自定义前向逻辑:

from ultralytics.models.yolo.detect import DetectionModel

class CustomYOLO(DetectionModel):
    def __init__(self, cfg='yolov8n.yaml', ch=3, nc=None, verbose=True):
        super().__init__(cfg, ch, nc, verbose)
        # 指定要提取的特征层索引(根据模型结构调整,比如backbone最后一层)
        self.target_feature_idx = -3

    def forward(self, x, augment=False, visualize=False, targets=None):
        # 执行原模型的特征提取
        feats = self.model(x)
        # 检测头输出预测结果
        pred = self.head(feats)

        if targets is not None:
            # 训练模式:计算检测损失并返回特征+损失
            loss, _ = self.loss(pred, targets)
            return feats[self.target_feature_idx], loss
        else:
            # 推理模式:返回特征+检测预测结果
            return feats[self.target_feature_idx], pred
方法二:使用PyTorch钩子函数提取特征

如果不想修改原模型代码,可注册钩子获取中间特征:

feature_cache = None

def feature_hook(module, input, output):
    global feature_cache
    feature_cache = output

# 绑定钩子到目标特征层(比如backbone最后一层)
target_layer = yolo_core.model[-3]
hook_handle = target_layer.register_forward_hook(feature_hook)

# 前向传播时,feature_cache会被赋值为目标层特征
_, yolo_loss = yolo_core(images, targets=gt_boxes)
extracted_feats = feature_cache

# 使用后移除钩子
hook_handle.remove()

3. 构建多任务自定义模型包装器

基于修改后的YOLO模型,整合自定义任务层:

import torch
import torch.nn as nn
from ultralytics.models.yolo.detect import DetectionModel

class MultiTaskYOLO(nn.Module):
    def __init__(self, yolo_cfg='yolov8n.yaml', pretrained_weights='yolov8n.pt', n1=10, n2=20, n3=30):
        super().__init__()
        # 初始化自定义YOLO模型
        self.yolo = CustomYOLO(cfg=yolo_cfg)
        # 加载预训练权重
        self.yolo.load(pretrained_weights)
        
        # 自动获取特征维度(通过dummy输入测试)
        with torch.no_grad():
            dummy_img = torch.randn(1, 3, 640, 640)
            feat, _ = self.yolo(dummy_img)
            # 全局池化后展平的维度
            self.feature_dim = feat.shape[1]
        
        # 自定义任务全连接层
        self.fc1 = nn.Linear(self.feature_dim, n1)
        self.fc2 = nn.Linear(self.feature_dim, n2)
        self.fc3 = nn.Linear(self.feature_dim, n3)
        self.global_pool = nn.AdaptiveAvgPool2d(1)

    def forward(self, images, gt_boxes=None):
        if self.training and gt_boxes is not None:
            # 训练模式:获取特征和YOLO检测损失
            feats, yolo_loss = self.yolo(images, targets=gt_boxes)
            # 特征全局池化并展平
            feats_flat = self.global_pool(feats).flatten(1)
            # 自定义任务输出
            logits1 = self.fc1(feats_flat)
            logits2 = self.fc2(feats_flat)
            logits3 = self.fc3(feats_flat)
            return {
                'logits1': logits1,
                'logits2': logits2,
                'logits3': logits3,
                'yolov8_loss': yolo_loss
            }
        else:
            # 推理模式:获取特征和检测结果
            feats, det_pred = self.yolo(images)
            feats_flat = self.global_pool(feats).flatten(1)
            logits1 = self.fc1(feats_flat)
            logits2 = self.fc2(feats_flat)
            logits3 = self.fc3(feats_flat)
            return {
                'logits1': logits1,
                'logits2': logits2,
                'logits3': logits3,
                'detection_pred': det_pred
            }

4. 自定义数据加载器

需同时提供YOLO检测标注和自定义任务标签,示例结构:

from torch.utils.data import Dataset, DataLoader
import cv2
import torch

class MultiTaskDataset(Dataset):
    def __init__(self, img_paths, det_labels, custom_labels, img_size=640):
        self.img_paths = img_paths
        self.det_labels = det_labels  # YOLO格式标注:[class_id, x_center, y_center, w, h](归一化到0-1)
        self.custom_labels = custom_labels  # 自定义任务标签(如分类/回归标签)
        self.img_size = img_size

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

    def __getitem__(self, idx):
        # 图像加载与预处理(对齐YOLO逻辑)
        img = cv2.imread(self.img_paths[idx])
        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
        # 图像resize、归一化(可复用YOLO的letterbox函数)
        img = torch.from_numpy(img).permute(2, 0, 1).float() / 255.0
        
        # 转换标注为张量
        det_label = torch.tensor(self.det_labels[idx], dtype=torch.float32)
        custom_label = torch.tensor(self.custom_labels[idx], dtype=torch.float32)
        
        return img, det_label, custom_label

# 构建DataLoader
dataset = MultiTaskDataset(img_paths_list, det_labels_list, custom_labels_list)
dataloader = DataLoader(dataset, batch_size=8, shuffle=True, collate_fn=lambda x: tuple(zip(*x)))

5. 自定义训练循环

组合YOLO检测损失与自定义任务损失,执行反向传播:

import torch.optim as optim

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = MultiTaskYOLO().to(device)
optimizer = optim.Adam(model.parameters(), lr=1e-4)

# 自定义任务损失函数
criterion1 = nn.CrossEntropyLoss()
criterion2 = nn.MSELoss()
criterion3 = nn.BCELoss()

model.train()
for epoch in range(10):
    total_epoch_loss = 0.0
    for imgs, det_labels, custom_labels in dataloader:
        imgs = torch.stack(imgs).to(device)
        det_labels = [label.to(device) for label in det_labels]
        # 拆分自定义任务标签
        lab1, lab2, lab3 = zip(*custom_labels)
        lab1 = torch.stack(lab1).to(device)
        lab2 = torch.stack(lab2).to(device)
        lab3 = torch.stack(lab3).to(device)

        optimizer.zero_grad()
        outputs = model(imgs, gt_boxes=det_labels)

        # 组合总损失(可根据任务重要性调整权重)
        total_loss = outputs['yolov8_loss'] + \
                     criterion1(outputs['logits1'], lab1) + \
                     criterion2(outputs['logits2'], lab2) + \
                     criterion3(outputs['logits3'], lab3)
        
        total_loss.backward()
        optimizer.step()
        total_epoch_loss += total_loss.item()
    
    print(f"Epoch {epoch+1} | Total Loss: {total_epoch_loss/len(dataloader):.4f}")

关键注意事项

  • 特征层选择:根据自定义任务需求选择不同层级特征(早期层适合细粒度任务,晚期层适合语义任务)
  • 权重冻结:若无需微调YOLOv8,可冻结其参数:for param in model.yolo.parameters(): param.requires_grad = False
  • 数据对齐:必须严格遵循YOLOv8的图像预处理逻辑,否则会严重影响检测性能
  • 损失权重:多任务训练时需调整各损失项权重,避免单一任务主导训练过程

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 11:54:57