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

如何重新训练PyTorchVideo的slowfast_r50_detection预训练模型?

基于PyTorchVideo的SlowFast R50 Detection重训练指南

一、数据集准备

  • 私有数据集需转换成COCO格式目标检测标注,即包含images文件夹(存放训练/验证视频帧或视频文件)、annotations文件夹(内有instances_train.json/instances_val.json标注文件)。若标注为VOC等其他格式,需先转成COCO格式。
  • 确保标注文件的categories字段与你的自定义类别完全匹配,类别ID要连续无冲突。

二、环境配置与模型导入

  • 先安装必备依赖:
    pip install torch torchvision pytorchvideo pycocotools
    
  • 加载预训练模型:
    import torch
    import pytorchvideo.models as models
    
    # 加载带预训练权重的模型
    model = models.slowfast.slowfast_r50_detection(pretrained=True)
    

三、修改模型头部适配自定义类别

  • 原模型头部针对COCO的90类设计,需替换为你的类别数:
    num_classes = 你的自定义类别数量  # 比如3类就填3
    # 获取原全连接层的输入特征维度
    in_features = model.head.fc.in_features
    # 替换最后的分类层
    model.head.fc = torch.nn.Linear(in_features, num_classes)
    

四、数据加载与预处理

  • 使用PyTorchVideo的CocoDetection类加载数据集,搭配对应预处理:
    from pytorchvideo.data import CocoDetection, make_clip_sampler
    from pytorchvideo.transforms import (
        ApplyTransformToKey, Normalize, RandomShortSideScale,
        UniformTemporalSubsample, ShortSideScale
    )
    from torchvision.transforms import Compose, Lambda, RandomCrop, RandomHorizontalFlip
    
    # 训练集预处理
    train_transform = Compose([
        ApplyTransformToKey(
            key="video",
            transform=Compose([
                UniformTemporalSubsample(32),  # 对应SlowFast的采样规则,适配2秒视频
                Lambda(lambda x: x / 255.0),
                Normalize((0.45, 0.45, 0.45), (0.225, 0.225, 0.225)),
                RandomShortSideScale(min_size=256, max_size=320),
                RandomCrop(224),
                RandomHorizontalFlip(p=0.5),
            ])
        )
    ])
    
    # 训练集加载
    train_dataset = CocoDetection(
        data_path="你的数据集路径/images/train",
        annotation_path="你的数据集路径/annotations/instances_train.json",
        clip_sampler=make_clip_sampler("random", clip_duration=2),
        transform=train_transform
    )
    
    # 验证集预处理(移除随机增强)
    val_transform = Compose([
        ApplyTransformToKey(
            key="video",
            transform=Compose([
                UniformTemporalSubsample(32),
                Lambda(lambda x: x / 255.0),
                Normalize((0.45, 0.45, 0.45), (0.225, 0.225, 0.225)),
                ShortSideScale(size=256),
            ])
        )
    ])
    
    # 验证集加载
    val_dataset = CocoDetection(
        data_path="你的数据集路径/images/val",
        annotation_path="你的数据集路径/annotations/instances_val.json",
        clip_sampler=make_clip_sampler("uniform", clip_duration=2),
        transform=val_transform
    )
    
    # 数据加载器(batch_size根据GPU显存调整)
    train_loader = torch.utils.data.DataLoader(
        train_dataset, batch_size=2, shuffle=True, num_workers=4
    )
    val_loader = torch.utils.data.DataLoader(
        val_dataset, batch_size=2, shuffle=False, num_workers=4
    )
    

五、训练循环编写

  • 设置优化器与损失函数,启动训练:
    from pytorchvideo.losses import DetectionLoss
    
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = model.to(device)
    
    optimizer = torch.optim.SGD(model.parameters(), lr=0.001, momentum=0.9, weight_decay=0.0001)
    loss_fn = DetectionLoss()
    
    num_epochs = 15  # 根据任务复杂度调整轮数
    for epoch in range(num_epochs):
        # 训练阶段
        model.train()
        total_train_loss = 0.0
        for batch in train_loader:
            video = batch["video"].to(device)
            labels = [label.to(device) for label in batch["label"]]
    
            optimizer.zero_grad()
            outputs = model(video, labels)
            loss = outputs["loss"]
            loss.backward()
            optimizer.step()
    
            total_train_loss += loss.item()
    
        avg_train_loss = total_train_loss / len(train_loader)
        print(f"Epoch {epoch+1} | 训练损失: {avg_train_loss:.4f}")
    
        # 验证阶段
        model.eval()
        total_val_loss = 0.0
        with torch.no_grad():
            for batch in val_loader:
                video = batch["video"].to(device)
                labels = [label.to(device) for label in batch["label"]]
                outputs = model(video, labels)
                total_val_loss += outputs["loss"].item()
    
        avg_val_loss = total_val_loss / len(val_loader)
        print(f"Epoch {epoch+1} | 验证损失: {avg_val_loss:.4f}\n")
    

六、模型保存与推理复用

  • 训练完成后保存权重:
    torch.save(model.state_dict(), "slowfast_custom_det.pth")
    
  • 沿用官方示例的推理方式加载模型:
    # 重新初始化模型(不加载预训练权重)
    model = models.slowfast.slowfast_r50_detection(pretrained=False)
    # 替换头部为自定义类别数
    model.head.fc = torch.nn.Linear(model.head.fc.in_features, num_classes)
    # 加载自定义权重
    model.load_state_dict(torch.load("slowfast_custom_det.pth"))
    model = model.to(device)
    model.eval()
    
    # 后续即可按照官方示例的流程输入视频,获取检测结果
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 11:02:09