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

运行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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 07:35:23