调用tfa.image.random_cutout触发slice index 3越界ValueError问题
错误原因
tfa.image.random_cutout 要求输入图像必须是4维张量,形状格式为 (batch_size, image_height, image_width, channels),必须包含批次维度。
你的代码中,tf.io.decode_png 解码输出的是3维张量,形状为 (image_height, image_width, 3),缺少最外层的批次维度。函数内部逻辑会尝试访问索引为3的维度,而3维张量的合法维度索引只有0、1、2,因此触发维度越界报错,和你看到的错误信息中slice index 3 of dimension 0 out of bounds的描述完全对应。
修正方案
你可以根据自己的数据处理流水线选择任意一种方案修复:
方案1:单图处理时手动增删批次维度
如果需要在组batch之前对单张图做cutout增强,只需要在调用接口前给张量临时加一个batch维度,接口执行完再移除该维度即可,修正后的代码如下:
def random_cut(image): image_string = tf.io.read_file(image) image = tf.io.decode_png(image_string, channels=3) image = tf.cast(image, tf.float32) / 255. # 临时增加batch维度,形状从(H,W,3)变为(1,H,W,3) image = tf.expand_dims(image, axis=0) image = tfa.image.random_cutout(image, (64,64), constant_values = 0) # 移除临时增加的batch维度,形状还原为(H,W,3) image = tf.squeeze(image, axis=0) return image dataset = dataset.map(random_cut)
方案2:组batch后再执行cutout增强
如果你的流水线后续本来就要做batch操作,可以把cutout步骤放到batch之后,此时数据集输出的张量天然是4维(batch_size, H, W, 3)的格式,不需要手动调整维度,参考代码如下:
# 单图阶段只做解码、归一化 def preprocess_single(image_path): image_string = tf.io.read_file(image_path) image = tf.io.decode_png(image_string, channels=3) image = tf.cast(image, tf.float32) / 255. return image dataset = dataset.map(preprocess_single) # 先组装batch dataset = dataset.batch(batch_size=32) # batch后直接调用cutout,无需调整维度 def apply_cutout(batch_images): return tfa.image.random_cutout(batch_images, (64,64), constant_values=0) dataset = dataset.map(apply_cutout)
注意:传入的cutout区域尺寸
(64,64)不能大于你数据集图像的实际高、宽值,否则会触发新的尺寸不匹配错误。
内容的提问来源于stack exchange,提问作者ChilliMayoo
相关产品推荐
相关产品推荐

