TensorFlow实现图像cutout增强时张量赋值与tfa调用报错求解
问题原因及解决方法
1 自定义cutout函数报错解决
报错原因
- TensorFlow的常规张量是不可变类型,没有
assign方法,不能像numpy数组那样直接对切片赋值 - 直接使用
numpy的随机函数在tf.data的Graph执行模式下无法被追踪,会导致执行异常
修改后的兼容版本
def augment_cutout(image, label, size=68, n_squares=1): h = tf.shape(image)[0] w = tf.shape(image)[1] channels = tf.shape(image)[2] # 初始化掩码为全1 mask = tf.ones((h, w, channels), dtype=image.dtype) for _ in range(n_squares): # 用TF的随机算子生成坐标 y = tf.random.uniform(shape=[], minval=0, maxval=h, dtype=tf.int32) x = tf.random.uniform(shape=[], minval=0, maxval=w, dtype=tf.int32) size_half = size // 2 y1 = tf.clip_by_value(y - size_half, 0, h) y2 = tf.clip_by_value(y + size_half, 0, h) x1 = tf.clip_by_value(x - size_half, 0, w) x2 = tf.clip_by_value(x + size_half, 0, w) # 生成cut区域的全0掩码 cut_mask = tf.concat([ tf.ones((y1, w, channels), dtype=image.dtype), tf.concat([ tf.ones((y2-y1, x1, channels), dtype=image.dtype), tf.zeros((y2-y1, x2-x1, channels), dtype=image.dtype), tf.ones((y2-y1, w-x2, channels), dtype=image.dtype) ], axis=1), tf.ones((h-y2, w, channels), dtype=image.dtype) ], axis=0) mask = mask * cut_mask # 原图乘掩码得到cutout后的结果 new_image = image * mask return tf.cast(new_image, tf.float32), label
调用位置要放在batch操作之前:
train_dataset = train_ds.map(map_func = preprocess_img, num_parallel_calls=AUTOTUNE) # 先做cutout增强,再batch train_dataset = train_dataset.map(map_func = augment_cutout, num_parallel_calls=AUTOTUNE) train_dataset = train_dataset.shuffle(buffer_size=1000).batch(batch_size=BATCH_SIZE).prefetch(buffer_size=AUTOTUNE)
2 tfa.image.random_cutout报错解决
报错原因
tfa.image.random_cutout的输入要求是4维批量张量(batch_size, height, width, channels),你当前是在batch操作之前调用,输入是单张图片的3维张量(height, width, channels),接口读取维度时索引越界。
修复方案
把random_cut_out的调用放在batch操作之后即可:
import tensorflow_addons as tfa def random_cut_out(images, labels): return tfa.image.random_cutout(images, (64, 64), constant_values = 1), labels train_dataset = train_ds.map(map_func = preprocess_img, num_parallel_calls=AUTOTUNE) train_dataset = train_dataset.map(map_func = augment, num_parallel_calls=AUTOTUNE) # 先batch,再调用cutout接口 train_dataset = train_dataset.shuffle(buffer_size=1000).batch(batch_size=BATCH_SIZE) train_dataset = train_dataset.map(random_cut_out, num_parallel_calls=AUTOTUNE) train_dataset = train_dataset.prefetch(buffer_size=AUTOTUNE)
内容的提问来源于stack exchange,提问作者JKnecht
相关产品推荐
相关产品推荐

