● 🍨 本文为🔗365天深度学习训练营中的学习记录博客
● 🍖 原作者:K同学啊
一、前期准备
1.设置GPU
import matplotlib.pyplot as plt import numpy as np import warnings warnings.filterwarnings('ignore') from tensorflow.keras import layers import tensorflow as tf gpus = tf.config.list_physical_devices("GPU") if gpus: tf.config.experimental.set_memory_growth(gpus[0],True) tf.config.set_visible_devices([gpus[0]],"GPU")2.导入数据
data_dir = "D:/新建文件夹/34-data" img_height = 224 img_width = 224 batch_size = 32 train_ds = tf.keras.preprocessing.image_dataset_from_directory( data_dir, validation_split = 0.3, subset = 'training', seed = 12, image_size = (img_height,img_width), batch_size = batch_size) val_ds = tf.keras.preprocessing.image_dataset_from_directory( data_dir, validation_split = 0.3, subset = 'validation', seed = 12, image_size = (img_height,img_width), batch_size = batch_size)Found 600 files belonging to 2 classes. Using 420 files for training. Found 600 files belonging to 2 classes. Using 180 files for validation.3.划分测试集
val_batches = tf.data.experimental.cardinality(val_ds) #验证集数据批次 test_ds = val_ds.take(val_batches // 5) #前五分之一的数据作为测试集 val_ds = val_ds.skip(val_batches // 5)#跳过前五分之一的数据,后面的是验证集 print('Number of validation batches: %d' % tf.data.experimental.cardinality(val_ds)) print('Number of test batches: %d' % tf.data.experimental.cardinality(test_ds))Number of validation batches: 5 Number of test batches: 1将验证集中的一部分划为测试集
使用tf.data.experimmental.cardinality()得出验证集的批次,前五分之一为测试集,剩余为验证集
4.输出分类名称
class_names = train_ds.class_names print(class_names)['cat', 'dog']5.归一化处理并配置数据集
AUTOTUNE = tf.data.AUTOTUNE def preprocess_image(image,label): return (image/255.0,label) train_ds = train_ds.map(preprocess_image,num_parallel_calls=AUTOTUNE) val_ds = val_ds.map(preprocess_image,num_parallel_calls=AUTOTUNE) test_ds = test_ds.map(preprocess_image,num_parallel_calls=AUTOTUNE) train_ds = train_ds.cache().prefetch(buffer_size=AUTOTUNE) val_ds = val_ds.cache().prefetch(buffer_size=AUTOTUNE)6.数据可视化
plt.figure(figsize=(15,10)) for images,labels in train_ds.take(1): for i in range(8): ax = plt.subplot(1,8,i+1) plt.imshow(images[i]) plt.title(class_names[labels[i]]) plt.axis("off")二、数据增强
data_augmentation = tf.keras.Sequential([ tf.keras.layers.experimental.preprocessing.RandomFlip("horizontal_and_vertical"), tf.keras.layers.experimental.preprocessing.RandomRotation(0.2), ])RandomFlip可以让图片水平或竖直翻转
RandomRotation可以让图片旋转
下面取一张图片进行数据增强
image = tf.expand_dims(images[i],0) plt.figure(figsize=(8, 8)) for i in range(9): augmented_image = data_augmentation(image) ax = plt.subplot(3, 3, i + 1) plt.imshow(augmented_image[0]) plt.axis("off")方法一:将数据增强放在model中,优点是可以得到GPU的加速,但数据增加只会在模型训练时生效,在模型评估(evaluate)和预测(predict)时不会生效
model = tf.keras.Sequential([ data_augmentation, layers.Conv2D(16, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Conv2D(32, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Conv2D(64, 3, padding='same', activation='relu'), layers.MaxPooling2D(), layers.Flatten(), layers.Dense(128, activation='relu'), layers.Dense(len(class_names)) ])方法二:在dataset中进行数据增强
batch_size = 32 AUTOTUNE = tf.data.AUTOTUNE def prepare(ds): ds = ds.map(lambda x, y: (data_augmentation(x, training=True), y), num_parallel_calls=AUTOTUNE) return ds train_ds = prepare(train_ds)三、编译
model.compile(optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=['accuracy'])设置优化器,损失函数,指标
四、模型训练
epochs=20 history = model.fit( train_ds, validation_data=val_ds, epochs=epochs )from datetime import datetime current_time = datetime.now() # 获取当前时间 acc = history.history['accuracy'] val_acc = history.history['val_accuracy'] loss = history.history['loss'] val_loss = history.history['val_loss'] epochs_range = range(epochs) plt.figure(figsize=(12, 4)) plt.subplot(1, 2, 1) plt.plot(epochs_range, acc, label='Training Accuracy') plt.plot(epochs_range, val_acc, label='Validation Accuracy') plt.legend(loc='lower right') plt.title('Training and Validation Accuracy') plt.xlabel(current_time) plt.subplot(1, 2, 2) plt.plot(epochs_range, loss, label='Training Loss') plt.plot(epochs_range, val_loss, label='Validation Loss') plt.legend(loc='upper right') plt.title('Training and Validation Loss') plt.show()五、自定义数据增强
import random def aug_img(image): seed = (random.randint(0,9), 0) # 随机改变图像对比度 stateless_random_brightness = tf.image.stateless_random_contrast(image, lower=0.1, upper=1.0, seed=seed) return stateless_random_brightness个人总结:本周学习了数据增强,数据增强可以让小数据集的训练结果更好,有两种实现方式,一种是嵌入model中,另一种是在数据集中进行数据增强,常用的数据增强的方法是随机水平或竖直翻转(RandomFlip)和让图片随机旋转(RandomRotation),另外也可以自定义函数进行数据增强,比如上文的随机调整图片的对比度。