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

Keras视频Transformer教程运行报错:TypeError索引问题求助

Keras视频Transformer教程中np.concatenate报错TypeError的解决方法

问题场景

运行Keras官方视频Transformer教程代码时,调用prepare_all_videos函数处理训练/测试数据集时触发TypeError,报错信息为:

TypeError: only integer scalar arrays can be converted to a scalar index

错误出现在prepare_all_videos函数内的这一行:

frames = np.concatenate(frames, padding)

相关代码片段

def prepare_all_videos(df, root_dir):
    num_samples = len(df)
    video_paths = df["video_name"].values.tolist()
    labels = df["tag"].values
    labels = label_processor(labels[..., None]).numpy()

    frame_features = np.zeros(
        shape=(num_samples, MAX_SEQ_LENGTH, NUM_FEATURES), dtype="float32"
    )

    for idx, path in enumerate(video_paths):
        frames = load_video(os.path.join(root_dir, path))

        # Pad shorter videos.
        if len(frames) < MAX_SEQ_LENGTH:
            diff = MAX_SEQ_LENGTH - len(frames)
            padding = np.zeros((diff, IMG_SIZE, IMG_SIZE, 3))
            frames = np.concatenate(frames, padding)  # 出错行

        frames = frames[None, ...]
        # 后续特征提取代码...

错误堆栈

81             diff = MAX_SEQ_LENGTH - len(frames)
     82             padding = np.zeros((diff, IMG_SIZE, IMG_SIZE, 3))
---&gt; 83             frames = np.concatenate(frames, padding)
     84 
     85         frames = frames[None, ...]

<__array_function__ internals> in concatenate(*args, **kwargs)

TypeError: only integer scalar arrays can be converted to a scalar index

数据集样例

video_nametag
v_CricketShot_g01_c01.aviCricketShot
v_CricketShot_g01_c02.aviCricketShot
v_TennisSwing_g07_c07.aviTennisSwing

错误原因

np.concatenate的参数使用完全错误:

  • 该函数第一个参数必须是数组组成的序列(列表/元组),用于指定要拼接的多个数组
  • 第二个参数是可选整数,用于指定拼接的轴(默认值为0)
    原代码直接将frames(单个数组)作为第一个参数,padding(另一个数组)作为第二个参数,不符合函数参数规则,导致类型不匹配报错。

修复方案

1. 修正np.concatenate调用方式

将出错行替换为:

frames = np.concatenate([frames, padding], axis=0)

解释:

  • 把frames和padding放在列表中作为第一个参数,明确告诉函数要拼接这两个数组
  • 指定axis=0是因为要在帧数量维度(数组第0轴)拼接,补全到MAX_SEQ_LENGTH帧

2. 补充过长视频的截断逻辑

原代码仅处理了帧数量不足的情况,未处理帧数量超过MAX_SEQ_LENGTH的场景,建议添加分支截断:

if len(frames) < MAX_SEQ_LENGTH:
    diff = MAX_SEQ_LENGTH - len(frames)
    padding = np.zeros((diff, IMG_SIZE, IMG_SIZE, 3))
    frames = np.concatenate([frames, padding], axis=0)
else:
    # 截断过长视频,保留前MAX_SEQ_LENGTH帧
    frames = frames[:MAX_SEQ_LENGTH]

修复后的完整代码片段

def prepare_all_videos(df, root_dir):
    num_samples = len(df)
    video_paths = df["video_name"].values.tolist()
    labels = df["tag"].values
    labels = label_processor(labels[..., None]).numpy()

    frame_features = np.zeros(
        shape=(num_samples, MAX_SEQ_LENGTH, NUM_FEATURES), dtype="float32"
    )

    for idx, path in enumerate(video_paths):
        frames = load_video(os.path.join(root_dir, path))

        # 处理视频帧长度:补全或截断
        if len(frames) < MAX_SEQ_LENGTH:
            diff = MAX_SEQ_LENGTH - len(frames)
            padding = np.zeros((diff, IMG_SIZE, IMG_SIZE, 3))
            frames = np.concatenate([frames, padding], axis=0)
        else:
            frames = frames[:MAX_SEQ_LENGTH]

        frames = frames[None, ...]
        # 后续特征提取代码...

内容的提问来源于stack exchange,提问作者alex-uarent-alex

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 13:56:05