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

PyTorch中设置batch_size大于1时VVC图像块预测模型报错的解决问询

PyTorch中设置batch_size大于1时VVC图像块预测模型报错的解决问询

问题背景

我正在构建一个神经网络,用于预测VVC(Versatile Video Coding)压缩过程中图像的分区方式。模型输入是YUV420图像的单帧Y分量,使用包含真实块位置和大小的CSV文件进行训练。

输入与真值

  • 输入:单帧10位YUV420图像的Y分量
  • 真值:包含块位置、大小和额外分区标记的CSV文件

示例(388016_320x480_37.yuv):
YUV帧示例

示例(388016_320x480_37.csv):
CSV真值示例

问题描述

我实现了train.py和dataset.py,但在DataLoader中设置batch_size > 1时遇到错误。batch_size=1时模型运行正常,但增大batch_size会导致运行时错误。

代码摘要

以下是我的custom_collate_fn和DataLoader设置的简化版本:

def custom_collate_fn(batch):
    frames = [item[0] for item in batch]  # Y-frame tensors
    blocks = [item[1] for item in batch]  # Block information
    frames = torch.stack(frames, dim=0)  # Stacking frames along batch dimension
    return frames, blocks

dataloader = DataLoader(
    dataset,
    batch_size=batch_size, 
    shuffle=True,
    collate_fn=custom_collate_fn
)

观察结果

  • 当batch_size = 1时,训练循环中的blocks_batch是包含单组块数据的列表
  • 当batch_size > 1时,blocks_batch变成列表的列表,导致索引时出错

训练循环片段:

for i, (frame, blocks_batch) in enumerate(dataloader):
    frame = frame.to(device)  # Shape: [batch_size, 1, H, W]
    blocks = blocks_batch[0]  # Works with batch_size=1 but fails with larger sizes

报错信息

File "C:\Users\Administrator\Documents\VVC_fast\test4\train.py", line 91, in <module>
    loss1 = criterion(out_split, target_split)
File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\modules\module.py", line 1518, in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\modules\module.py", line 1527, in _call_impl
    return forward_call(*args, **kwargs)
File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\modules\loss.py", line 725, in forward
    return F.binary_cross_entropy_with_logits(input, target,
File "C:\ProgramData\miniconda3\envs\newenv3\lib\site-packages\torch\nn\functional.py", line 3193, in binary_cross_entropy_with_logits
    raise ValueError(f"Target size ({target.size()}) must be the same as input size ({input.size()})")
ValueError: Target size (torch.Size([1, 1])) must be the same as input size (torch.Size([2, 1]))

问题分析与解决思路

我来帮你拆解问题本质,再给出几个实用的解决方向:

核心问题根源

直接报错的原因很明确:当batch_size=2时,模型输出是对应整个batch的张量(比如形状[2,1]),但你在训练循环里只取了blocks_batch[0]——也就是第一个样本的真值(形状[1,1]),两者维度完全不匹配,自然触发了尺寸错误。

而更深层的核心问题是:不同图像的VVC分区块数量是可变的,你当前的custom_collate_fn只是把每个样本的块数据简单打包成列表的列表,这种结构没法和带batch维度的模型输出直接对齐,必须针对性处理这些变长的真值数据。

具体解决方法

方案1:拼接所有块数据+记录batch归属(适合逐块预测的模型)

如果你的模型是对每个真实存在的分区块做预测(比如预测每个块的分裂类型),可以把整个batch的所有块数据拼在一起,同时记录每个块属于哪个batch样本,这样就能和模型输出一一对应。

修改custom_collate_fn:

def custom_collate_fn(batch):
    frames = []
    all_blocks = []
    batch_indices = []  # 标记每个块属于第几个batch样本

    for batch_idx, (frame, blocks) in enumerate(batch):
        frames.append(frame)
        # 把当前样本的块数据加入列表
        all_blocks.append(blocks)
        # 给当前样本的每个块添加batch索引
        batch_indices.extend([batch_idx] * len(blocks))
    
    # 堆叠帧张量(batch维度在前)
    frames = torch.stack(frames, dim=0)
    # 拼接所有块数据成一个大张量
    all_blocks = torch.cat(all_blocks, dim=0)
    batch_indices = torch.tensor(batch_indices, dtype=torch.long)
    
    return frames, all_blocks, batch_indices

调整训练循环:

for i, (frame, all_blocks, batch_indices) in enumerate(dataloader):
    frame = frame.to(device)
    all_blocks = all_blocks.to(device)
    batch_indices = batch_indices.to(device)
    
    # 模型前向传播
    outputs = model(frame)
    # 假设模型输出是[batch_size, 所有可能块的数量, 预测维度],先展平成[总块数, 预测维度]
    out_split = outputs.flatten(0, 1)
    # 只保留对应真实块的预测结果(用batch_indices对齐)
    out_split = out_split[batch_indices]
    
    # 计算损失
    loss1 = criterion(out_split, all_blocks)

方案2:用Padding统一块数据长度(适合块数差异不大的场景)

如果你的数据集里不同图像的块数差异不大,可以把每个样本的块数据补到当前batch的最大块数,同时用mask标记哪些是真实块、哪些是补的无效数据。

修改custom_collate_fn:

def custom_collate_fn(batch):
    frames = [item[0] for item in batch]
    blocks_list = [item[1] for item in batch]
    
    # 找到当前batch中块数最多的样本
    max_block_count = max(len(blocks) for blocks in blocks_list)
    padded_blocks = []
    
    for blocks in blocks_list:
        # 计算需要补的长度
        pad_length = max_block_count - len(blocks)
        # 补0(如果你的块数据是张量,要保证padding的维度和原数据一致)
        padding = torch.zeros(pad_length, blocks.shape[1], dtype=blocks.dtype)
        padded = torch.cat([blocks, padding], dim=0)
        padded_blocks.append(padded)
    
    # 堆叠帧和补全后的块张量
    frames = torch.stack(frames, dim=0)
    padded_blocks = torch.stack(padded_blocks, dim=0)  # shape: [batch_size, max_block_count, 块特征维度]
    
    return frames, padded_blocks

调整训练循环(需要忽略padding部分的损失):

for i, (frame, padded_blocks) in enumerate(dataloader):
    frame = frame.to(device)
    padded_blocks = padded_blocks.to(device)
    
    outputs = model(frame)  # shape: [batch_size, max_block_count, 预测维度]
    # 创建mask:标记哪些是真实块(非padding,这里假设padding是全0张量)
    mask = (padded_blocks.sum(dim=-1) != 0)
    
    # 只计算真实块的损失
    loss1 = criterion(outputs[mask], padded_blocks[mask])

方案3:把真值转换成固定尺寸的掩码(推荐,从根源解决batch问题)

如果你的模型是对图像的每个位置/每个固定网格块做预测(比如预测每个位置是否是块边界,或者属于哪种块类型),那最好的方式是把CSV里的块信息转换成和输入图像同尺寸的掩码张量。

比如,根据CSV里的块位置和大小,生成一个HxW的张量:每个像素标记对应的块ID、块大小或者分裂类型。这样每个样本的真值张量尺寸都是固定的(和输入Y帧的尺寸匹配),DataLoader可以直接用默认的collate_fn堆叠,batch_size>1时完全不会有问题。

这种方式从根源上消除了变长数据的问题,是最稳妥的方案,推荐优先考虑。


备注:内容来源于stack exchange,提问作者조동건

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 19:13:07