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

ResNet50+LSTM视频帧7分类、数据加载及PyTorch转换咨询

ResNet50+LSTM 动作阶段识别方案答疑

首先说明原Keras代码的显性问题:你的任务是7分类,原代码最后输出层用1个神经元搭配sigmoid激活属于二分类配置,需要调整为7个神经元搭配softmax激活,损失函数使用多分类交叉熵。

1. 4帧输入的维度适配与数据加载实现

你查到的输入维度(B, N, 3, H, W)是PyTorch框架默认的通道在前格式,Keras默认使用通道在后格式(B, N, H, W, 3),你原来写的Input((4, 224, 224, 3))本身符合Keras的维度规则,GlobalAveragePooling2D层不需要修改,它只对单帧的特征图做空间维度池化,和时序输入长度无关。正确的实现逻辑如下:

  • 数据加载部分:不要将单帧作为独立样本,遍历每个动作阶段的抽帧文件夹时,用滑动窗口截取连续4帧作为一个样本,窗口步长可根据数据量设为1/2/4,每个样本的标签对应该段帧所属的动作阶段类别,注意加载帧时必须按文件名排序,保证时序顺序正确。
  • 维度对齐:Keras训练时保持输入维度为(batch_size, 4, 224, 224, 3)即可;PyTorch训练时调整为(batch_size, 4, 3, 224, 224)。
  • 特征提取逻辑:Keras中TimeDistributed包裹的CNN会自动将4帧拆分为单帧,逐帧送入ResNet提取特征,权重在所有时间步共享,经过GAP后输出维度为(batch_size, 4, 2048)的时序特征序列,正好匹配后续LSTM的输入要求,不需要手动调整池化层参数。

2. PyTorch版本实现代码

模型定义

import torch
import torch.nn as nn
from torchvision.models import resnet50, ResNet50_Weights

class ResNet50LSTM(nn.Module):
    def __init__(self, num_classes=7, lstm_hidden=2048, lstm_layers=1, dropout_rate=0.4):
        super().__init__()
        # 加载ImageNet预训练ResNet50,移除原网络最后的全连接分类头
        backbone = resnet50(weights=ResNet50_Weights.IMAGENET1K_V1)
        self.cnn_encoder = nn.Sequential(*list(backbone.children())[:-1])
        # 冻结ResNet所有层参数
        for param in self.cnn_encoder.parameters():
            param.requires_grad = False

        # LSTM时序编码层,输入维度为ResNet输出的单帧2048维特征
        self.lstm = nn.LSTM(
            input_size=2048,
            hidden_size=lstm_hidden,
            num_layers=lstm_layers,
            batch_first=True
        )

        self.leaky_relu = nn.LeakyReLU()
        self.dropout = nn.Dropout(dropout_rate)
        self.fc1 = nn.Linear(lstm_hidden, 2048)
        self.relu = nn.ReLU()
        # 7分类输出层
        self.cls_head = nn.Linear(2048, num_classes)

    def forward(self, x):
        # 输入x维度: (batch_size, 4, 3, 224, 224)
        batch_size, seq_len, c, h, w = x.shape
        # 合并batch和时序维度,逐帧提取特征
        x = x.reshape(batch_size * seq_len, c, h, w)
        # CNN编码后输出维度: (batch_size*4, 2048, 1, 1)
        x = self.cnn_encoder(x)
        # 移除1*1的空间维度,恢复时序拆分,输出维度: (batch_size,4,2048)
        x = x.reshape(batch_size, seq_len, -1)
        # LSTM编码,取最后一个时间步的输出作为整段4帧的特征
        lstm_out, _ = self.lstm(x)
        x = lstm_out[:, -1, :]
        # 全连接分类头
        x = self.leaky_relu(x)
        x = self.dropout(x)
        x = self.fc1(x)
        x = self.relu(x)
        x = self.cls_head(x)
        return x

数据加载器核心实现

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

class ActionPhaseDataset(Dataset):
    def __init__(self, data_root, seq_len=4, transform=None):
        self.seq_len = seq_len
        self.transform = transform
        self.samples = []
        # 按文件夹名排序映射0-6的类别标签
        self.cls_list = sorted(os.listdir(data_root))
        for cls_id, cls_name in enumerate(self.cls_list):
            cls_dir = os.path.join(data_root, cls_name)
            # 按文件名排序帧,保证时序正确
            frame_list = sorted([
                os.path.join(cls_dir, f) for f in os.listdir(cls_dir)
                if f.lower().endswith(('.jpg', '.jpeg', '.png'))
            ])
            # 滑动窗口生成连续4帧的样本
            for i in range(len(frame_list) - seq_len + 1):
                self.samples.append((frame_list[i:i+seq_len], cls_id))

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

    def __getitem__(self, idx):
        frame_paths, label = self.samples[idx]
        frame_tensor_list = []
        for p in frame_paths:
            img = Image.open(p).convert('RGB')
            if self.transform:
                img = self.transform(img)
            frame_tensor_list.append(img)
        # 堆叠后维度为(4,3,224,224),符合模型输入要求
        frames = torch.stack(frame_tensor_list)
        return frames, label

3. 参数设置逻辑与含义解释

ResNet50相关参数

  • include_top=False(PyTorch中通过移除最后一层fc实现):不需要ResNet自带的1000类ImageNet分类头,只保留卷积特征提取能力,输出通用视觉特征后接自定义的时序模块和分类头。
  • input_shape=(224,224,3):224*224是ImageNet预训练的标准输入尺寸,3对应RGB三通道,和预训练数据的输入分布对齐,能保证特征提取效果。
  • 冻结ResNet层:ImageNet预训练权重已经学习到通用的边缘、纹理、物体等视觉特征,如果你的数据集规模不大,冻结backbone可以大幅降低过拟合风险,同时加快训练速度;后续如果数据量充足,可以再解冻ResNet的顶层卷积层做微调。

LSTM相关参数

  • TimeDistributed(PyTorch中通过合并batch维度、权重共享实现):让同一个ResNet编码器对4帧的每一帧单独做特征提取,所有时间步共享CNN权重,不会为不同帧单独创建模型参数,大幅降低整体参数量。
  • LSTM输入维度2048:ResNet50经过全局平均池化后,单帧输出的特征向量维度正好是2048,和LSTM的输入尺寸严格匹配,这也是你之前调试时改到这个值能跑通的原因。
  • LSTM隐藏层维度2048:是你设置的时序特征编码维度,数值越大模型拟合能力越强,但参数量会同步上升,更容易过拟合,可以根据自己的数据集规模调整为1024或512。

原代码的无效/错误参数说明

  • 最后一层Dense(1,activation='sigmoid')是二分类配置,不适用于你的7分类任务,需要替换为7维输出+softmax激活。
  • 全连接层中input_dim=inputs是无效写法,Keras和PyTorch都会自动推断相邻层的输入维度,不需要手动传入输入层张量,这个参数是之前试错时遗留的错误配置,删掉即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 15:24:20