TensorFlow 高效数据读取与实时增强实战
1 从磁盘文件夹读取图像
把不同类别的图片分别放在子目录中,一行代码即可生成 tf.data.Dataset:
BATCH = 32
IMG_SIZE = 224
ds_train = tf.keras.utils.image_dataset_from_directory(
data_dir,
validation_split=0.2,
subset='training',
seed=42,
image_size=(IMG_SIZE, IMG_SIZE),
batch_size=BATCH
)
print(ds_train.class_names) # 查看类别
2 使用 TensorFlow Datasets 一键获取拆分好的集合
import tensorflow_datasets as tfds
(raw_train, raw_val, raw_test), info = tfds.load(
'tf_flowers',
split=['train[:80%]', 'train[80%:90%]', 'train[90%:]'],
as_supervised=True,
with_info=True
)
print(info.features['label'].num_classes) # 类别数
label2name = info.features['label'].int2str
3 数据归一化与基础增强
用 Rescaling 把像素值压缩到 [0,1]:
rescale = tf.keras.layers.Rescaling(1./255)
ds_train = ds_train.map(lambda img, lbl: (rescale(img), lbl))
4 自定义增强层
4.1 Lambda 层快速封装
def invert(img, prob=0.5):
if tf.random.uniform([]) < prob:
img = 255 - img
return img
invert_layer = tf.keras.layers.Lambda(lambda x: invert(x))
4.2 继承 Layer 写可复用组件
class RandomInvert(tf.keras.layers.Layer):
def __init__(self, prob=0.5):
super().__init__()
self.prob = prob
def call(self, inputs):
return invert(inputs, self.prob)
5 tf.image 原生算子组合管道
img = tf.io.read_file('flower.jpg')
img = tf.image.decode_jpeg(img, channels=3)
flipped = tf.image.flip_left_right(img)
gray = tf.squeeze(tf.image.rgb_to_grayscale(img))
saturated = tf.image.adjust_saturation(img, 3)
bright = tf.image.adjust_brightness(img, 0.4)
cropped = tf.image.central_crop(img, 0.5)
rotated = tf.image.rot90(img)
6 可重复随机增强:stateless API
seed = (1, 2) # 固定种子保证复现
img = tf.image.stateless_random_brightness(img, max_delta=0.3, seed=seed)
img = tf.image.stateless_random_contrast(img, lower=0.2, upper=0.8, seed=seed)
img = tf.image.stateless_random_crop(img, size=[200, 200, 3], seed=seed)
7 端到端增强流水线
7.1 定义缩放与增强函数
IMG_SIZE = 224
def resize_and_scale(img, lbl):
img = tf.image.resize(tf.cast(img, tf.float32), [IMG_SIZE, IMG_SIZE])
return img / 255., lbl
def augment_with_seed(img_lbl, seed):
img, lbl = img_lbl
img, lbl = resize_and_scale(img, lbl)
# 四周填充 6 像素再随机裁剪回原尺寸
img = tf.image.resize_with_crop_or_pad(img, IMG_SIZE + 6, IMG_SIZE + 6)
img = tf.image.stateless_random_crop(
img, size=[IMG_SIZE, IMG_SIZE, 3], seed=seed)
# 亮度扰动
new_seed = tf.random.experimental.stateless_split(seed, 1)[0]
img = tf.image.stateless_random_brightness(
img, max_delta=0.5, seed=new_seed)
img = tf.clip_by_value(img, 0., 1.)
return img, lbl
7.2 使用 Counter 生成唯一种子
counter = tf.data.experimental.Counter()
ds_train = tf.data.Dataset.zip((raw_train, (counter, counter)))
ds_train = (ds_train
.shuffle(1000)
.map(augment_with_seed, num_parallel_calls=tf.data.AUTOTUNE)
.batch(BATCH)
.prefetch(tf.data.AUTOTUNE))
ds_val = (raw_val
.map(resize_and_scale, num_parallel_calls=tf.data.AUTOTUNE)
.batch(BATCH)
.prefetch(tf.data.AUTOTUNE))
7.3 用 tf.random.Generator 管理状态
rng = tf.random.Generator.from_seed(123, alg='philox')
def augment_with_rng(img, lbl):
seed = rng.make_seeds(2)[0]
return augment_with_seed((img, lbl), seed)
ds_train = raw_train.map(augment_with_rng, num_parallel_calls=tf.data.AUTOTUNE)