TensorFlow随机图像数据集创建及Dataset.from_generator迭代报错求助
TensorFlow随机图像数据集创建及Dataset.from_generator迭代报错求助
嗨,我太懂你这种想用模拟图像替代磁盘图像测试 pipeline 的需求了——读磁盘慢到让人跺脚有没有😂!看了你的代码,问题大概率出在Dataset.from_generator的迭代逻辑匹配和生成器类的实现细节上,我来给你捋捋问题出在哪,以及怎么改:
先说说你当前代码的核心问题:
- Dataset.from_generator是靠迭代器遍历来拉取数据的,不是靠
__getitem__的索引访问!你现在的FakeImageGenerator只写了__getitem__,TensorFlow没法正确识别成可迭代的生成器,自然会在迭代时报错 - 你没有给Dataset显式指定输出的数据类型和形状,TensorFlow无法自动推断数据结构时,就会抛出类型/形状不匹配的错误
- 另外你代码里
np.random.rand生成的是float64类型,和TensorFlow默认的float32类型不统一,这也容易埋下隐性bug
直接上修复后的完整可运行代码:
import numpy as np import tensorflow as tf # 统一配置参数,方便后续修改 full_image_batch_size = 1 batch_size = 1 image_height = 3000 image_width = 5328 image_channels = 3 tile_size = 256 class FakeImageGenerator: def __init__(self, number_of_images): self.generated = 0 self.number_of_images = number_of_images def __len__(self): return self.number_of_images def __iter__(self): # 重点!实现__iter__方法,返回生成器迭代器 while self.generated < self.number_of_images: # 生成随机图像,直接转成float32对齐TF默认类型 arr = np.random.rand(image_height, image_width, image_channels).astype(np.float32) # 转成TensorFlow张量 img_tensor = tf.convert_to_tensor(arr) self.generated += 1 yield img_tensor # 创建生成器实例 fake_gen = FakeImageGenerator(number_of_images=10) # 用from_generator创建数据集,必须显式指定输出类型和形状 dataset = tf.data.Dataset.from_generator( lambda: fake_gen, # 传入可迭代的生成器对象 output_types=tf.float32, output_shapes=(image_height, image_width, image_channels) ) # 接下来是你的pipeline流程:batch、切patch等 # 先做全图batch dataset = dataset.batch(full_image_batch_size) # 定义切patch的函数,用TF原生函数,速度更快 def extract_image_patches(img_batch): # img_batch形状是 (batch_size, height, width, channels) patches = tf.image.extract_patches( images=img_batch, sizes=[1, tile_size, tile_size, 1], strides=[1, tile_size, tile_size, 1], rates=[1, 1, 1, 1], padding='VALID' ) # 把patches转成 (总patch数, tile_size, tile_size, channels) patches = tf.reshape(patches, (-1, tile_size, tile_size, image_channels)) return patches # 应用patch切分,用多线程加速 dataset = dataset.map(extract_image_patches, num_parallel_calls=tf.data.AUTOTUNE) # 最终的batch dataset = dataset.batch(batch_size) # 加个预取,提升训练时的速度 dataset = dataset.prefetch(tf.data.AUTOTUNE) # 测试迭代,看看能不能正常输出 for idx, batch in enumerate(dataset.take(2)): print(f"第{idx+1}个batch的形状: {batch.shape}")
补充几个关键优化点:
- 如果你需要给图像加标签,只需要在yield的时候改成
yield (img_tensor, label),同时把output_types改成(tf.float32, tf.int32),output_shapes改成((image_height, image_width, image_channels), ()) - 尽量用TensorFlow原生的预处理函数,别在map里用numpy操作——如果实在要用,得套个
tf.py_function,但速度会慢很多 - 要是生成的图像需要rescale,直接在生成张量后加
img_tensor = img_tensor / 255.0就行,用TF的原生操作更高效
如果按照上面改完还是报错,把具体的错误信息贴出来,我再帮你细抠!
备注:内容来源于stack exchange,提问作者poilwant
相关产品推荐
相关产品推荐

