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)) ---> 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_name | tag |
|---|---|
| v_CricketShot_g01_c01.avi | CricketShot |
| v_CricketShot_g01_c02.avi | CricketShot |
| v_TennisSwing_g07_c07.avi | TennisSwing |
错误原因
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
相关产品推荐
相关产品推荐

