如何修改输入加载代码解决3DCNN输入形状不匹配错误?
问题解决:调整3DCNN输入维度匹配模型要求
错误原因
你的3DCNN模型定义的输入通道数为1,但当前数据加载后输出的形状是[2,10,1,320,864]——这里的10被错误地放在了通道维度的位置,导致模型报错。需要将输入形状调整为[2,1,10,320,864],也就是把**通道维度(1)**放在第2位,**帧序列维度(10)**放在第3位。
修改方案(两种任选其一)
方案一:调整单帧处理逻辑,再添加通道维度
- 循环内去掉单帧的通道维度:将
frame = frame.reshape(1, frame.shape[0], frame.shape[1])改为:frame = frame.reshape(frame.shape[0], frame.shape[1]) # 形状变为(320,864) - 堆叠帧后添加通道维度:将
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
相关产品推荐
相关产品推荐

