Keras ImageDataGenerator加载numpy数组图像方法及flow报错解决
问题原因
ImageDataGenerator.flow()的传参逻辑和flow_from_directory完全不同,不支持直接传入两类样本列表自动生成标签:
- 第一个位置参数
x要求传入全部样本组成的numpy数组,标准形状为(样本总数, 图像高度, 图像宽度, 通道数),不接受拆分后的多组样本列表 - 第二个位置参数
y要求传入与x顺序一一对应的标签数组,二分类场景下一般由0(负样本,无系外行星)、1(正样本,有系外行星)组成
你当前代码将有行星的光变曲线列表作为x传入、无行星的列表作为y传入,传入的x结构不符合方法内部的格式校验逻辑,因此触发了索引越界报错。
修正代码
你需要先合并两类样本、手动生成对应标签,再传入flow()方法,修正示例如下:
import numpy as np from tensorflow.keras.preprocessing.image import ImageDataGenerator def predictExo(exotrainfile,noexotrainfile,testfile): train = ImageDataGenerator(rescale=1/255) test = ImageDataGenerator(rescale=1/255) batchsize = 32 # 读取两类训练数据,补充单通道维度(光变曲线一般为单通道灰度图) exo_train = np.array(get_lightcurves(exotrainfile))[..., np.newaxis] noexo_train = np.array(get_lightcurves(noexotrainfile))[..., np.newaxis] # 拼接所有训练样本 train_x = np.concatenate([exo_train, noexo_train], axis=0) # 生成对应标签:存在系外行星标记为1,不存在标记为0 train_y = np.array([1]*len(exo_train) + [0]*len(noexo_train)) train_ds = train.flow( train_x, train_y, batch_size = batchsize, shuffle=True # 训练集建议打乱顺序,避免模型学习到样本排列规律 ) # 测试集无标签时可仅传入样本数组 test_x = np.array(get_lightcurves(testfile))[..., np.newaxis] test_ds = test.flow( test_x, batch_size = batchsize, shuffle=False )
注意:请确保
get_lightcurves()返回的所有光变曲线图像尺寸完全一致,否则拼接数组、输入CNN时都会触发维度错误。
可选替代方案
如果不想手动处理数组拼接和标签生成,也可以选择更灵活的TensorFlow输入管道实现:
- 使用
tf.data.DatasetAPI直接从numpy数组构造数据集,自带批处理、打乱、预加载优化,自定义数据处理逻辑的自由度更高,不受ImageDataGenerator的参数规则限制 - 针对天文光变曲线的数据增强(如加观测噪声、时间轴微偏移)可以写成自定义函数,直接集成到Dataset的处理流程中,比
ImageDataGenerator预设的自然图像增强方法更贴合系外行星检测的场景需求
内容的提问来源于stack exchange,提问作者Sophie Frankel
相关产品推荐
相关产品推荐

