🍨本文为🔗365天深度学习训练营中的学习记录博客
- 🍖原作者:
学习目的:采用CNN实现多云、下雨、晴、日出四种天气状态的识别
一、 前期准备
关于环境
- 语言环境:Python3.6
- 编译器:vsCode
- 深度学习环境:TensorFlow 2.6.2
1.数据导入
import tensorflow as tf import os,PIL,pathlib import matplotlib.pyplot as plt import numpy as np from tensorflow import keras from tensorflow.keras import layers,modelsdata_dir = "D:/jupyter notebook/DL-100-days/datasets/weather_photos/" data_dir = pathlib.Path(data_dir)2.查看数据
数据集一共分为cloudy、rain、shine、sunrise四类,分别存放于weather_photos文件夹中以各自名字命名的子文件夹中。
image_count = len(list(data_dir.glob('*/*.jpg'))) print("图片总数为:",image_count)运行结果:
roses = list(data_dir.glob('sunrise/*.jpg')) PIL.Image.open(str(roses[0]))运行结果:
二、数据预处理
1.加载数据
使用image_dataset_from_directory方法将磁盘中的数据加载到tf.data.Dataset中
batch_size = 32 img_height = 180 img_width = 180img_height/width:将图像统一缩放到 180×180 像素。
""" 关于image_dataset_from_directory()的详细介绍可以参考文章:https://mtyjkh.blog.csdn.net/article/details/117018789 """ train_ds = tf.keras.preprocessing.image_dataset_from_directory( data_dir, validation_split=0.2, subset="training", seed=123, image_size=(img_height, img_width), batch_size=batch_size)运行结果:
""" 关于image_dataset_from_directory()的详细介绍可以参考文章:https://mtyjkh.blog.csdn.net/article/details/117018789 """ val_ds = tf.keras.preprocessing.image_dataset_from_directory( data_dir, validation_split=0.2, subset="validation", seed=123, image_size=(img_height, img_width), batch_size=batch_size)运行结果:
通过class_names输出数据集的标签。标签将按字母顺序对应于目录名称。
class_names = train_ds.class_names print(class_names)运行结果:
2. 可视化数据
plt.figure(figsize=(20, 10)) for images, labels in train_ds.take(1): for i in range(20): ax = plt.subplot(5, 10, i + 1) plt.imshow(images[i].numpy().astype("uint8")) plt.title(class_names[labels[i]]) plt.axis("off")运行结果:
3. 再次检查数据
for image_batch, labels_batch in train_ds: print(image_batch.shape) print(labels_batch.shape) break运行结果:
- 这是一批形状180x180x3的32张图片。
Label_batch是形状(32,),这些标签对应32张图片
4. 配置数据集
AUTOTUNE = tf.data.AUTOTUNE train_ds = train_ds.cache().shuffle(1000).prefetch(buffer_size=AUTOTUNE) val_ds = val_ds.cache().prefetch(buffer_size=AUTOTUNE)三、构建CNN网络
num_classes = 4 """ 关于卷积核的计算不懂的可以参考文章:https://blog.csdn.net/qq_38251616/article/details/114278995 layers.Dropout(0.4) 作用是防止过拟合,提高模型的泛化能力。 在上一篇文章花朵识别中,训练准确率与验证准确率相差巨大就是由于模型过拟合导致的 关于Dropout层的更多介绍可以参考文章:https://mtyjkh.blog.csdn.net/article/details/115826689 """ model = models.Sequential([ layers.experimental.preprocessing.Rescaling(1./255, input_shape=(img_height, img_width, 3)), layers.Conv2D(16, (3, 3), activation='relu', input_shape=(img_height, img_width, 3)), # 卷积层1,卷积核3*3 layers.AveragePooling2D((2, 2)), # 池化层1,2*2采样 layers.Conv2D(32, (3, 3), activation='relu'), # 卷积层2,卷积核3*3 layers.AveragePooling2D((2, 2)), # 池化层2,2*2采样 layers.Conv2D(64, (3, 3), activation='relu'), # 卷积层3,卷积核3*3 layers.Dropout(0.3), # 让神经元以一定的概率停止工作,防止过拟合,提高模型的泛化能力。 layers.Flatten(), # Flatten层,连接卷积层与全连接层 layers.Dense(128, activation='relu'), # 全连接层,特征进一步提取 layers.Dense(num_classes) # 输出层,输出预期结果 ]) model.summary() # 打印网络结构
Rescaling
将像素值从 [0, 255] 缩放到 [0, 1],有助于加快收敛。Dropout(0.3)
随机丢弃 30% 的神经元,防止过拟合
运行结果:
四、编译
# 设置优化器 opt = tf.keras.optimizers.Adam(learning_rate=0.001) model.compile(optimizer=opt, loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=['accuracy'])五、训练模型
epochs = 10 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()运行结果:
七、总结
本周学习中,使用了image_dataset_from_directory实现高效的数据加载,并结合.cache()、.prefetch()等机制优化了数据流水线的读取性能;其次,在建模中加入了dropout层防止过拟合。