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

如何修改输入加载代码解决3DCNN输入形状不匹配错误?

问题解决:调整3DCNN输入维度匹配模型要求

错误原因

你的3DCNN模型定义的输入通道数为1,但当前数据加载后输出的形状是[2,10,1,320,864]——这里的10被错误地放在了通道维度的位置,导致模型报错。需要将输入形状调整为[2,1,10,320,864],也就是把**通道维度(1)**放在第2位,**帧序列维度(10)**放在第3位。

修改方案(两种任选其一)

方案一:调整单帧处理逻辑,再添加通道维度

  1. 循环内去掉单帧的通道维度:将frame = frame.reshape(1, frame.shape[0], frame.shape[1])改为:
    frame = frame.reshape(frame.shape[0], frame.shape[1])  # 形状变为(320,864)
    
  2. 堆叠帧后添加通道维度:将frames = torch.stack(frames)改为:
    frames = torch.stack(frames).unsqueeze(0)  # 堆叠后是(10,320,864),加通道后变为(1,10,320,864)
    

方案二:直接调换堆叠后的维度顺序

保持单帧处理逻辑不变,仅修改堆叠后的维度排列:将frames = torch.stack(frames)改为:

frames = torch.stack(frames).permute(1, 0, 2, 3)  # 原堆叠后是(10,1,320,864),调换后变为(1,10,320,864)

修改后的完整代码示例(方案二)

def __getitem__(self, idx):
    video_idx = idx // 260
    frame_idx = idx % 260 + 41
    video_dir = self.video_dirs[video_idx]
    frames = []

    # Get all files in the directory
    all_files = os.listdir(video_dir)

    # Select only .jpg files
    jpg_files = [file for file in all_files if file.endswith('.jpg')]

    # Extract the number from the file name and sort
    numbered_files = sorted(jpg_files, key=lambda x: int(re.findall(r'\d+', x)[-1]))

    for i in range(frame_idx, frame_idx + 10):
        # Get the file with the corresponding number
        frame_file = numbered_files[i-1]  # -1 because indexing starts from 0
        frame_path = os.path.join(video_dir, frame_file)
        print(f"Loading image from {frame_path}")
        frame = cv2.imread(frame_path)
        if frame is None:
            raise ValueError(f"Could not load image at {frame_path}")
        frame = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)  # Convert to grayscale
        frame = frame.reshape(1, frame.shape[0], frame.shape[1])  # Add channel dimension
        frame = torch.tensor(frame, dtype=torch.float32)
        print(frame.shape)
        frames.append(frame)

    # 调整维度顺序,匹配模型输入要求
    frames = torch.stack(frames).permute(1, 0, 2, 3)
    label = self.labels[idx]
    return frames, label

效果验证

修改后,单样本输出形状为(1,10,320,864),经过DataLoader的batch处理(batch_size=2)后,会得到你需要的(2,1,10,320,864),完全匹配模型的输入要求,即可解决报错问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 12:55:36