TensorFlow Datasets:dataset.batch()调用后按批次裁剪/调整图像尺寸
map作用范围说明
你在batch()之后调用的map()方法作用于每个单独的批次:tf.data.Dataset的map始终作用于数据集的单个元素,执行batch操作后,数据集的每个元素就是形状为(batch_size, H, W, 3)的批次张量,因此map会逐个批次处理数据,完全符合你的使用场景,不需要额外通过迭代器手动处理批次再用concatenate合并。
动态批次尺寸实现方案
你当前的代码里SIZE写死为(300, 300)才会导致所有批次尺寸相同,只要将目标尺寸改为每次处理批次时动态生成即可,以下是修改后的代码示例:
# 数据集构造部分和你原来的代码一致,只修改resize_data函数和后续map逻辑 ALLOWED_BATCH_SIZES = [224, 300, 400] # 可按需调整允许的尺寸列表 def resize_data(images, labels): tf.print('Original shape -->', tf.shape(images)) # 随机选择当前批次的目标尺寸 target_size = tf.gather(ALLOWED_BATCH_SIZES, tf.random.uniform(shape=[], minval=0, maxval=len(ALLOWED_BATCH_SIZES), dtype=tf.int32)) # 同批次所有图像都resize到目标尺寸,也可替换为tf.image.crop_and_resize实现裁剪+调整逻辑 resized_imgs = tf.image.resize(images, (target_size, target_size)) return resized_imgs, labels dataset = dataset.map(resize_data, num_parallel_calls=tf.data.experimental.AUTOTUNE) dataset = dataset.prefetch(tf.data.experimental.AUTOTUNE)
注意事项
- 如果需要目标尺寸是连续范围内的随机值,直接替换
target_size生成逻辑即可:target_size = tf.random.uniform(shape=[], minval=224, maxval=512, dtype=tf.int32) - 下游训练的模型需要支持动态输入尺寸(比如全卷积结构),如果存在固定输入维度的全连接层则无法兼容动态批次尺寸的输入。
内容的提问来源于stack exchange,提问作者Avid Learner
相关产品推荐
相关产品推荐

