对图像执行reshape操作报错:BatchDataset对象无reshape属性

在适配GAN模型输入形状、对图像做reshape操作时触发如下运行错误:
AttributeError: 'BatchDataset' object has no attribute 'reshape'
错误原因
reshape是TensorFlow张量、NumPy数组的内置形状调整方法,你当前调用方法的对象是tf.dataAPI生成的BatchDataset批量数据集实例——这是封装了数据加载、预处理、批量迭代逻辑的流水线对象,本身没有提供reshape接口,直接对整个数据集对象调用reshape就会触发这个属性错误。
解决方法
不要试图直接对数据集对象做reshape,把形状调整逻辑通过.map()方法嵌入到数据预处理流水线中即可,根据你reshape的时机分两种写法:
- 单样本维度reshape(batch操作前处理,推荐)
对每一条单独的样本做形状调整,处理完再组装批量,逻辑更清晰:import tensorflow as tf def preprocess_reshape(image, label): # 替换成你GAN模型需要的输入形状,比如宽28、高28、单通道的手写数字图形状 image = tf.reshape(image, [28, 28, 1]) # 需要归一化像素值也可以在这里加逻辑:image = tf.cast(image, tf.float32) / 255.0 return image, label # 把reshape逻辑挂载到数据集,再做批量、预取等后续操作 dataset = dataset.map(preprocess_reshape, num_parallel_calls=tf.data.AUTOTUNE).batch(batch_size=32).prefetch(tf.data.AUTOTUNE) - 批量维度reshape(已经做了batch操作后的处理)
如果你已经提前调用了.batch()生成了批量数据集,就在map逻辑里保留第一维的batch维度,用-1自动推导批量大小即可:def reshape_batch_data(batch_imgs, batch_labels): batch_imgs = tf.reshape(batch_imgs, [-1, 28, 28, 1]) return batch_imgs, batch_labels dataset = dataset.batch(32).map(reshape_batch_data, num_parallel_calls=tf.data.AUTOTUNE)
不要为了调用reshape直接把整个数据集转成NumPy数组加载到内存,数据量较大时会直接触发内存溢出,也会丢失tf.data流水线并行加载、预取的性能优势。
内容的提问来源于stack exchange,提问作者yungdenzel
相关产品推荐
相关产品推荐

