tf.data.Dataset.from_generator不支持uint16、uint32类型问题求助
解决TensorFlow from_generator读取uint16数据的类型不支持问题
我之前也碰到过类似的坑,核心问题在于TensorFlow的from_generator需要显式指定输出数据类型——哪怕你的numpy数组是uint16类型,自动类型推断有时候会出错或者不兼容。下面我结合你的场景给出完整的解决思路和代码:
1. 先复现问题(补全你的测试代码)
import tensorflow as tf import numpy as np def hdf5_data_generator(): # 模拟从HDF5读取uint16数据的逻辑 for _ in range(10): # 生成随机uint16数据,模拟HDF5读取结果 yield np.random.randint(0, 65535, size=(32, 32), dtype=np.uint16) # 错误用法:不指定output_types时,可能触发类型不支持错误 # ds = tf.data.Dataset.from_generator(hdf5_data_generator)
2. 正确的解决方案
只需要在from_generator中明确指定output_types(以及可选的output_shapes),让TensorFlow清楚知道生成数据的类型和形状:
# 正确用法:显式指定输出类型为tf.uint16 ds = tf.data.Dataset.from_generator( hdf5_data_generator, output_types=tf.uint16, # 关键:对应numpy的uint16 output_shapes=tf.TensorShape([32, 32]) # 可选:如果形状固定可以指定,提升性能 ) # 测试迭代数据集 for sample in ds.take(1): print(f"数据类型:{sample.dtype}") print(f"数据形状:{sample.shape}")
3. 额外注意事项
- 如果你从HDF5读取的是带标签的成对数据(比如
(图像数据, 标签)),需要把output_types设为元组:def labeled_generator(): for _ in range(10): img = np.random.randint(0, 65535, size=(32,32), dtype=np.uint16) label = np.random.randint(0, 10, dtype=np.int32) yield img, label ds = tf.data.Dataset.from_generator( labeled_generator, output_types=(tf.uint16, tf.int32), # 对应输入和标签的类型 output_shapes=(tf.TensorShape([32,32]), tf.TensorShape([])) ) - 先确认HDF5中的数据确实是
uint16类型:用h5py打开文件后,通过h5_file['你的数据集名称'].dtype检查,确保是numpy.uint16,避免源数据类型本身就不对。 - 如果还是有问题,可以在生成器内部直接转换为TensorFlow张量返回:
def generator_with_conversion(): data = np.random.randint(0, 65535, size=(32,32), dtype=np.uint16) yield tf.convert_to_tensor(data, dtype=tf.uint16)
内容的提问来源于stack exchange,提问作者mikkola
相关产品推荐
相关产品推荐

