运行keypoint_rotation.py遇RuntimeError:张量尺寸不匹配
解决keypoint_rotation.py中torch.stack的维度不匹配问题
你遇到的RuntimeError是因为torch.stack要求所有输入张量的维度完全一致,但当前处理的样本(或帧)特征尺寸存在差异(1024 vs 1728)。以下是具体排查和解决步骤:
第一步:定位问题根源
先在stack_features函数中添加调试代码,确认是哪种维度不匹配:
def stack_features(features, _): # 打印调试信息,定位问题 for idx, sample_frames in enumerate(features): print(f"样本{idx}的总帧数:{len(sample_frames)}") if sample_frames: base_dim = sample_frames[0].shape print(f"样本{idx}的单帧特征维度:{base_dim}") # 检查同一样本内所有帧的维度是否一致 for frame_idx, frame in enumerate(sample_frames): if frame.shape != base_dim: print(f"样本{idx}的第{frame_idx}帧维度异常:{frame.shape}") return torch.stack([torch.stack(ft, dim=0) for ft in features], dim=0)
运行后根据输出判断:
- 如果是不同样本的帧数不同:比如样本0有16帧(1664=1024),样本181有27帧(2764=1728),属于样本间序列长度不一致;
- 如果是同一样本内不同帧的特征维度不同:比如某样本内有的帧是1024维,有的是1728维,属于数据或加载代码的问题。
第二步:对应解决方法
情况1:样本间帧数不同(最可能)
这是数据增强中常见的问题,需要统一所有样本的序列长度(帧数),可以通过以下两种方式处理:
方法1:修改Field定义,让torchtext自动填充/截断
先计算数据集的最大帧数,再配置Field的固定长度和填充值:
# 先遍历数据集获取最大帧数 max_frame_count = max(len(ex.keypoints) for ex in dataset) # 修改第60行左右的keypoint_field定义 keypoint_field = Field( sequential=True, use_vocab=False, dtype=torch.float, postprocessing=stack_features, batch_first=True, fix_length=max_frame_count, # 统一到最大帧数 padding_value=0.0 # 用0填充短序列 )
如果最大帧数过大,也可以设置一个合理的阈值(比如500帧),超过的截断,不足的填充。
方法2:手动预处理样本
在加载数据集后,对每个样本的keypoints做统一长度处理:
def standardize_keypoint_length(keypoints, target_len, pad_val=0.0): current_len = len(keypoints) if current_len < target_len: # 填充短序列 pad_frame = np.full_like(keypoints[0], pad_val) keypoints += [pad_frame] * (target_len - current_len) elif current_len > target_len: # 截断长序列 keypoints = keypoints[:target_len] return keypoints # 加载数据集后调用 target_frame_len = 500 # 根据你的数据情况调整 for example in dataset: example.keypoints = standardize_keypoint_length(example.keypoints, target_frame_len)
情况2:单样本内帧维度不同
这种情况说明数据本身存在异常,或者数据加载代码有bug:
- 检查数据源:确认所有keypoint文件的格式一致,每个帧的特征维度统一;
- 检查加载逻辑:查看读取keypoint的代码,是否有部分帧被错误解析(比如读取时的维度计算错误)。
关于你怀疑的第60行和第96行
- 第60行的Field定义:如果没有配置
fix_length和padding_value,torchtext不会自动处理序列长度差异,这是导致错误的核心原因之一; - 第96行的Iterator创建:如果
batch_size>1且sort_within_batch=False,会让不同长度的样本被分到同一个batch,触发维度不匹配。建议保持batch_size=1或者开启sort_within_batch=True(按序列长度排序后组batch,减少填充量)。
内容的提问来源于stack exchange,提问作者Nashenas Agha
相关产品推荐
相关产品推荐

