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

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.Dataset API直接从numpy数组构造数据集,自带批处理、打乱、预加载优化,自定义数据处理逻辑的自由度更高,不受ImageDataGenerator的参数规则限制
  • 针对天文光变曲线的数据增强(如加观测噪声、时间轴微偏移)可以写成自定义函数,直接集成到Dataset的处理流程中,比ImageDataGenerator预设的自然图像增强方法更贴合系外行星检测的场景需求

内容的提问来源于stack exchange,提问作者Sophie Frankel

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 07:33:19