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

CNN+LSTM视频二分类模型损失不下降、指标无提升问题排查求助

视频二分类模型训练失败问题排查与修复方案

问题概述

构建视频二分类网络,每个视频加载16/32帧及对应标签,模型采用预训练ResNet101+LSTM+全连接层结构,计划使用SGD优化器(lr=0.01)+BCELoss。但训练50+epoch后损失始终徘徊在0.7左右,模型全预测0或1,各项指标停留在随机水平(约50%),尝试在小数据集上过拟合验证学习能力也无效果。


一、数据集类问题修复

1. 帧读取逻辑冗余且易出错

当前逐帧遍历视频判断采样帧的方式效率低,且易因视频读取失败导致帧数量不足;同时缓存目录每次初始化都删除重建,存在不必要的IO开销。

修改代码:

# 替换get_total_frame_count中的逐帧循环逻辑
def get_total_frame_count(self):
    for vid_name in os.listdir(self.vids_dir):
        curr_sample_frame_label_map = {}
        curr_label = self.vid_name_label_map[vid_name.split('.')[0]]
        vid_path = os.path.join(self.vids_dir, vid_name)
        cap = cv2.VideoCapture(vid_path)
        vid_frames = int(cap.get(7))
        frames_to_capture = np.linspace(0, vid_frames-1, self.sequence_length, dtype=np.int16)
        curr_sample_frame_label_map[curr_label] = []
        
        # 直接跳转到目标帧读取,避免逐帧遍历
        for frame_ind in frames_to_capture:
            cap.set(cv2.CAP_PROP_POS_FRAMES, frame_ind)
            success, frame = cap.read()
            if not success:
                # 读取失败用最后成功读取的帧填充
                if len(curr_sample_frame_label_map[curr_label]) == 0:
                    raise ValueError(f"无法读取视频{vid_name}的任何帧")
                frame = cv2.imread(curr_sample_frame_label_map[curr_label][-1])
            save_as = os.path.join(self.cache_dir, f'{vid_name}_{frame_ind}.jpg')
            if not os.path.exists(save_as):
                cv2.imwrite(save_as, frame)
            curr_sample_frame_label_map[curr_label].append(save_as)
        
        # 补全帧数量
        while len(curr_sample_frame_label_map[curr_label]) < self.sequence_length:
            curr_sample_frame_label_map[curr_label].append(curr_sample_frame_label_map[curr_label][-1])
        self.frames_label_map.append(curr_sample_frame_label_map.copy())

# 修改缓存目录创建逻辑,避免重复删除重建
def create_cache_dir(self):
    os.makedirs(self.cache_dir, exist_ok=True)

2. 样本格式与标签处理优化

当前__getitem__返回的帧序列可能为空,标签处理步骤冗余,且未强制统一图像通道数(易出现灰度图通道不匹配问题)。

修改代码:

def __getitem__(self, idx):
    label_img_paths_map = self.frames_label_map[idx]
    label = list(label_img_paths_map.keys())[0]
    path2imgs = label_img_paths_map[label]
    frames_tr = []
    
    for p2i in path2imgs:
        # 强制转RGB,避免灰度图导致通道数错误
        frame = Image.open(p2i).convert('RGB')
        frame_tr = self.transform(frame)
        frames_tr.append(frame_tr)
    
    # 确保帧数量符合要求
    assert len(frames_tr) == self.sequence_length, f"样本{idx}帧数量异常"
    frames_tr = torch.stack(frames_tr)
    # 直接返回tensor格式标签
    return frames_tr, torch.tensor(label, dtype=torch.float32)

# 简化collate_fn,无需过滤空样本
def collate_fn(batch):
    imgs_batch, label_batch = list(zip(*batch))
    imgs_tensor = torch.stack(imgs_batch)
    labels_tensor = torch.stack(label_batch).view(-1, 1)
    return imgs_tensor, labels_tensor

二、模型类核心问题修复

1. ResNet梯度被完全冻结

模型中对ResNet前向传播使用with torch.no_grad(),导致预训练权重完全无法更新,仅LSTM和全连接层训练,特征提取能力严重受限。

修改代码:

def forward(self, x_3d):
    hidden = None
    for t in range(x_3d.size(1)):
        # 移除torch.no_grad(),允许ResNet权重更新
        x = self.resnet(x_3d[:, t, :, :, :])  
        out, hidden = self.lstm(x.unsqueeze(0), hidden)         

    x = self.fc1(out[-1, :, :])
    x = F.relu(x)
    x = self.fc2(x)
    x = self.sig(x)
    return x

可选优化:分层冻结ResNet
如果不想完全解冻ResNet,可仅训练后几层:

def __init__(self, num_classes=1):
    super(CNNLSTM, self).__init__()
    self.resnet = resnet101(pretrained=True)
    self.resnet.fc = nn.Sequential(nn.Linear(self.resnet.fc.in_features, 300))
    
    # 冻结ResNet前几层,仅解冻最后一个卷积层和fc层
    for param in self.resnet.parameters():
        param.requires_grad = False
    for param in self.resnet.layer4.parameters():
        param.requires_grad = True
    for param in self.resnet.fc.parameters():
        param.requires_grad = True
    
    self.lstm = nn.LSTM(input_size=300, hidden_size=256, num_layers=3)
    self.fc1 = nn.Linear(256, 128)
    self.fc2 = nn.Linear(128, num_classes)
    self.sig = nn.Sigmoid()

2. LSTM时序建模逻辑错误

当前将单帧特征作为长度为1的序列输入LSTM,完全无法学习帧间时序关系,等价于仅使用最后一帧的特征输出。

修改代码:

def forward(self, x_3d):
    # x_3d shape: [batch_size, seq_len, C, H, W]
    batch_size, seq_len = x_3d.size(0), x_3d.size(1)
    
    # 合并batch和seq维度,一次性输入ResNet提取特征
    x = x_3d.view(-1, 3, 224, 224)  # shape: [batch*seq_len, C, H, W]
    x = self.resnet(x)
    x = x.view(batch_size, seq_len, -1)  # shape: [batch_size, seq_len, 300]
    
    # 转换LSTM输入格式为(seq_len, batch_size, input_size)
    x = x.transpose(0, 1)
    out, hidden = self.lstm(x)
    
    # 取最后一个时间步的输出
    x = out[-1, :, :]
    x = self.fc1(x)
    x = F.relu(x)
    x = self.fc2(x)
    x = self.sig(x)
    return x

三、训练脚本问题修复

1. 学习率与优化器配置优化

原计划使用lr=0.01,但脚本实际用0.001,且SGD在预训练模型微调时需分层设置学习率,避免破坏预训练权重。

修改代码:

# 分层设置学习率,ResNet用较小学习率,新增层用较大学习率
optimizer = torch.optim.SGD([
    {'params': cnnlstm_model.resnet.parameters(), 'lr': 1e-4},
    {'params': cnnlstm_model.lstm.parameters(), 'lr': 1e-3},
    {'params': cnnlstm_model.fc1.parameters(), 'lr': 1e-3},
    {'params': cnnlstm_model.fc2.parameters(), 'lr': 1e-3}
], momentum=0.9, weight_decay=1e-4)

2. 添加训练数据增强

当前训练集使用测试变换,无任何数据增强,小数据集下模型无法有效学习特征。

修改代码:

# 新增训练变换
train_transformer = transforms.Compose([
    transforms.Resize((h,w)),
    transforms.RandomHorizontalFlip(),
    transforms.RandomRotation(10),
    transforms.ColorJitter(brightness=0.2, contrast=0.2),
    transforms.ToTensor(),
    transforms.Normalize(mean, std),
])

# 训练集使用训练变换
train_ds = VideoDataset(
    vids_dir='dataset_merged_clips/videos',
    labels_path='dataset_merged_clips/labels.csv',
    transform=train_transformer,
    sequence_length=16
)
train_dl = DataLoader(train_ds, batch_size=batch_size, shuffle=True, collate_fn=rnn_collate_fn, pin_memory=True)

3. 处理类别不平衡问题

若数据集存在类别不平衡,模型会偏向预测多数类,导致指标停留在随机水平,可使用加权BCELoss:

# 假设正样本占比为pos_ratio,需根据实际数据集计算
pos_ratio = 0.3
weight = torch.tensor([1/(1-pos_ratio)]).to(device)
criterion = BCELoss(weight=weight)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 02:22:04