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

TensorFlow数据生成器读取独立文件输入失败问题求助

问题:tf.data.Dataset.from_generator读取单个npy文件时报错,内存加载则正常

我尝试用Python生成器结合tf.data.Dataset.from_generator()创建tf.data数据集,输入到Keras神经网络。输入数据是目录下的单个npy文件,每个文件是815×815的NumPy数组,对应一个样本;标签是每个样本的独热向量,共54个类别。

读取单个文件的生成器代码

# wavelet_filenames是输入样本的文件名数组,格式为"1.npy"到"1259.npy"
def shuffle_generator(wavelet_filenames, label, seed): 
    idx = np.arange(len(wavelet_filenames))
    np.random.default_rng(seed).shuffle(idx)
    path = "C:/Users/tslee/Python/whited_wavelets/" # 存放所有输入样本的文件夹
    for i in idx:
        data_array = np.load((path + wavelet_filenames[i]), allow_pickle=True)
        yield data_array, label[i]

创建数据集代码

dataset = tf.data.Dataset.from_generator( # 用于训练
    shuffle_generator,
    args=[X_train_filenames, Y_train, 66], # 随机种子66
    output_signature=(
        tf.TensorSpec(shape=(815,815), dtype=tf.float16),
        tf.TensorSpec(shape=(num_classes), dtype=tf.bool)))

报错信息

File c:\Users\tslee\Python\IRP\Lib\site-packages\keras\utils\traceback_utils.py:70, in filter_traceback..error_handler(*args, **kwargs)
     67     filtered_tb = _process_traceback_frames(e.__traceback__)
     68     # To get the full stack trace, call:
     69     # `tf.debugging.disable_traceback_filtering()`
---> 70     raise e.with_traceback(filtered_tb) from None
     71 finally:
     72     del filtered_tb

File c:\Users\tslee\Python\IRP\Lib\site-packages\tensorflow\python\eager\execute.py:52, in quick_execute(op_name, num_outputs, inputs, attrs, ctx, name)
     50 try:
...
TypeError: can only concatenate str (not "bytes") to str


     [[{{node PyFunc}}]]
     [[IteratorGetNext]] [Op:__inference_train_function_37747]

可行的内存加载方案

将所有样本合并到一个npy文件中,加载到内存后使用生成器,代码可正常运行:

内存加载的生成器代码

def shuffle_generator_mem(wavelets, label, seed): 
    idx = np.arange(len(wavelets))
    np.random.default_rng(seed).shuffle(idx)
    for i in idx:
        yield wavelets[i], label[i] # 返回小波(输入样本)和标签(独热向量)

对应数据集创建代码

dataset = tf.data.Dataset.from_generator(
    shuffle_generator_mem,
    args=[wavelets, Y_train, 66], # 随机种子66
    output_signature=(
        tf.TensorSpec(shape=(815,815), dtype=tf.float16),
        tf.TensorSpec(shape=(num_classes), dtype=tf.bool)))

补充说明

内存加载时的npy文件维度为(n×815×815),每行对应一个样本,加载代码:

wavelets = np.load("C:/Users/tslee/IRP Data/Python/DATA/whited_hf_wavelet_samples.npy", allow_pickle=True)

我不清楚为什么读取单个文件时会报错,求解决方法。


内容的提问来源于stack exchange,提问作者Thomas Lee Young

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 19:23:18