尧图网站建设 尧图网络
  • 首页
  • 关于我们
  • 服务项目
  • 案例展示
  • 建站流程
  • 资讯中心
  • 联系我们
首页/资讯中心/详情

第 T10 周:数据增强

第 T10 周:数据增强
📅 发布时间:2026/8/1 3:55:43

声明

  • 本文为「365 天深度学习训练营」内部学习记录。
  • 本文参考 K 同学啊课程内容完成,仅用于个人学习与交流。
  • 猫狗识别数据集仅用于学习交流,请勿对外分享数据集。
  • 本篇为个人在 T10 关卡上的实践记录。

第 T10 周:数据增强

在本教程中,你将学会如何进行数据增强,并通过数据增强用少量数据达到非常棒的识别准确率。本文将展示两种数据增强方式,以及如何自定义数据增强方式并将其放到代码当中。

难度:夯实基础
语言:Python3、TensorFlow2

本周要求

  • 学会在代码中使用数据增强手段来提高 acc
  • 请探索更多的数据增强手段并记录

我的环境

  • 语言环境:Python 3.6.5(实践环境可用更新版本)
  • 编译器:Jupyter Notebook
  • 深度学习框架:TensorFlow 2.4.1(实践环境:TensorFlow 2.11)
  • 实验数据:34-data.zip(25.2 MB)
  • 数据目录:34-data/(2 类:cat/dog,各300张,共600张)

两种增强接入方式:

方式做法特点
方法一增强模块嵌入model可走 GPU 加速;仅Model.fit时生效
方法二在Dataset中map在数据流水线里做增强,灵活可组合

一、前期准备工作

1. 设置 GPU

如果使用的是 CPU 可以注释掉这部分代码。

importtensorflowastffromtensorflow.kerasimportlayersimportmatplotlib.pyplotaspltimportnumpyasnpimportrandomimportwarnings warnings.filterwarnings('ignore')gpus=tf.config.list_physical_devices("GPU")ifgpus:tf.config.experimental.set_memory_growth(gpus[0],True)# 设置 GPU 显存用量按需使用tf.config.set_visible_devices([gpus[0]],"GPU")print(gpus)

2. 加载数据

关于tf.keras.preprocessing.image_dataset_from_directory的介绍可参考:CSDN 讲解。

由于原始数据集不包含测试集,因此需要创建一个。使用tf.data.experimental.cardinality确定验证集中有多少批次,然后将其中的20%移至测试集。

一共有猫、狗两类。

data_dir="./34-data"img_height=224img_width=224batch_size=32train_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.
class_names=train_ds.class_namesprint(class_names)

预期输出:

['cat', 'dog']

从验证集再拆出测试集:

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))

归一化 + 性能配置:

AUTOTUNE=tf.data.AUTOTUNEdefpreprocess_image(image,label):returnimage/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().shuffle(1000).prefetch(buffer_size=AUTOTUNE)val_ds=val_ds.cache().prefetch(buffer_size=AUTOTUNE)test_ds=test_ds.cache().prefetch(buffer_size=AUTOTUNE)


二、数据增强

我们可以使用tf.keras.layers.experimental.preprocessing.RandomFlip与tf.keras.layers.experimental.preprocessing.RandomRotation进行数据增强(新版本也可直接写tf.keras.layers.RandomFlip/RandomRotation)。

  • RandomFlip:水平和垂直随机翻转每个图像
  • RandomRotation:随机旋转每个图像
data_augmentation=tf.keras.Sequential([layers.experimental.preprocessing.RandomFlip("horizontal_and_vertical"),layers.experimental.preprocessing.RandomRotation(0.2),])

第一个层表示进行随机的水平和垂直翻转,第二个层表示按0.2的因子进行随机旋转。

可视化增强效果:

plt.figure(figsize=(8,8))forimages,labelsintrain_ds.take(1):image=tf.expand_dims(images[0],0)foriinrange(9):augmented_image=data_augmentation(image,training=True)ax=plt.subplot(3,3,i+1)plt.imshow(augmented_image[0])plt.axis("off")

更多数据增强方式可参考:RandomRotation 文档。

探索记录:更多增强手段

方式API(TF 2.4+)作用
水平/垂直翻转RandomFlip镜像扩充视角
随机旋转RandomRotation角度扰动
随机缩放RandomZoom模拟远近
随机对比度RandomContrast/tf.image.stateless_random_contrast光照变化
随机亮度RandomBrightness/tf.image.random_brightness明暗变化
随机平移RandomTranslation位置偏移
随机裁剪RandomCrop局部裁切后再用
饱和度tf.image.random_saturation色彩浓淡

三、增强方式

方法一:将其嵌入 model 中

这样做的好处是:数据增强这块的工作可以得到GPU 加速(如果你使用了 GPU 训练的话)。

注意:只有在模型训练时(Model.fit)才会进行增强,在模型评估(Model.evaluate)以及预测(Model.predict)时并不会进行增强操作。

model=tf.keras.Sequential([layers.Input(shape=(img_height,img_width,3)),data_augmentation,# ← 嵌进 modellayers.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.Dropout(0.2),# 小数据易过拟合,略加 Dropoutlayers.Flatten(),layers.Dense(128,activation='relu'),layers.Dense(len(class_names))])

本周 notebook 主跑方案即为方法一。

方法二:在 Dataset 数据集中进行数据增强

defprepare(ds,shuffle=False,augment=False):ifshuffle:ds=ds.shuffle(1000)ifaugment:ds=ds.map(lambdax,y:(data_augmentation(x,training=True),y),num_parallel_calls=AUTOTUNE)returnds.prefetch(buffer_size=AUTOTUNE)# 示例:仅对训练集增强(此时 model 内不要再重复嵌增强层)# train_ds_aug = prepare(train_ds, shuffle=True, augment=True)

方法二适合把增强放在 CPU 数据流水线里,与 GPU 训练并行;自定义增强函数也更容易挂到map上。


四、训练模型

在准备对模型进行训练之前,还需要再对其进行一些设置。以下内容是在模型的编译步骤中添加的:

  • 损失函数(loss):用于衡量模型在训练期间的误差
  • 优化器(optimizer):决定模型如何根据数据与损失函数更新参数
  • 评价函数(metrics):用于监控训练和测试步骤(本例用准确率)
model.compile(optimizer='adam',loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),metrics=['accuracy'])epochs=30history=model.fit(train_ds,validation_data=val_ds,epochs=epochs)

开始训练后,用测试集评估:

loss,acc=model.evaluate(test_ds)print("Accuracy",acc)

参考输出(方法一,实测,ml_env / TF 2.11,CPU):

Epoch 1/30 14/14 - loss: 1.0330 - accuracy: 0.5571 - val_loss: 0.6756 - val_accuracy: 0.6824 Epoch 8/30 14/14 - loss: 0.2823 - accuracy: 0.8690 - val_loss: 0.3675 - val_accuracy: 0.8649 Epoch 16/30 14/14 - loss: 0.1371 - accuracy: 0.9571 - val_loss: 0.2234 - val_accuracy: 0.9392 Epoch 27/30 14/14 - loss: 0.0724 - accuracy: 0.9738 - val_loss: 0.1779 - val_accuracy: 0.9459 ...... Epoch 30/30 14/14 - loss: 0.0627 - accuracy: 0.9738 - val_loss: 0.3629 - val_accuracy: 0.9122 1/1 [==============================] - loss: 0.8222 - accuracy: 0.8125 Accuracy 0.8125

本机结果摘要:最佳验证准确率约94.59%,测试集准确率81.25%(测试集约 1 个 batch / 32 张,波动较大;期末模型 val 略回落时 test 也会受影响)。课程示例可冲到更高;可继续加RandomZoom、换学习率、或按 val 最优做 checkpoint。

训练曲线

fromdatetimeimportdatetime 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(len(loss))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(str(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()

五、自定义增强函数

这是可以自由发挥的地方。课程示例用随机对比度:

importrandomdefaug_img(image):seed=(random.randint(0,9),0)# 随机改变图像对比度returntf.image.stateless_random_contrast(image,lower=0.1,upper=1.0,seed=seed)

可视化:

# 取一张已归一化图片还原到约 0~255 再增强展示forimages,labelsintrain_ds.take(1):image=tf.expand_dims(images[3]*255.0,0)print("Min and max pixel values:",image.numpy().min(),image.numpy().max())plt.figure(figsize=(8,8))foriinrange(9):augmented_image=aug_img(image)ax=plt.subplot(3,3,i+1)plt.imshow(tf.clip_by_value(augmented_image[0],0,255).numpy().astype("uint8"))plt.axis("off")


那么如何将自定义增强函数应用到数据上呢?参考上文的preprocess_image,把aug_img嵌进去即可:

defpreprocess_image(image,label):image=image/255.0image=aug_img(image)# 仅建议挂在训练集returnimage,label# train_ds = train_ds.map(preprocess_image, num_parallel_calls=AUTOTUNE)

也可组合更多变换(探索):

defaug_img_extra(image):image=tf.image.random_brightness(image,max_delta=0.2)image=tf.image.random_saturation(image,lower=0.5,upper=1.5)iftf.random.uniform([])>0.5:image=tf.image.flip_left_right(image)returnimage

总结

本周在小样本猫狗数据(600 张)上完成:

经验分享

  1. 小数据先保语义,再谈花样。猫狗图默认头朝上,垂直翻转有时会把样本拧得不像真图;horizontal_and_vertical能扩充多样性,但若 val 抖、训得很快过拟合,可以先改成只做水平翻转,再逐步加旋转 / 缩放。增强幅度不是越大越好。
  2. 两种接入方式别叠着用。方法一把增强嵌进 model,方法二在Dataset.map里做;同一套RandomFlip/RandomRotation同时开两遍,等于扰动加倍,难排查。主跑选一种即可,自定义aug_img更适合挂在方法二的预处理里。
  3. 看 val,别只看 train。本周无增强时很容易训到接近 100% 的 train_acc,但 val 会掉;加上增强 + 一点Dropout后,最佳 val 能到九成以上。期末权重不一定是最好的那一轮——test 只有约一个 batch(32 张),波动大,更稳妥是按val_accuracy存ModelCheckpoint,再用最优权重去evaluate。
  4. 测试集要从验证集切,且切之前别cache错顺序。先take/skip再归一化与cache,否则流水线状态容易乱。val_batches // 5在 batch 很少时可能切出 0 或 1 个 batch,解读 test_acc 时心里要有数:数字漂亮不一定代表泛化已经稳了。
  5. 增强只在fit时生效(方法一)。evaluate/predict不会再随机翻转旋转,这是预期行为;对比「有无增强」时,应保证网络结构、epoch、随机种子尽量一致,只改增强相关代码,结论才有可比性。

相关新闻

  • 从零构建语音识别应用:百度API实战指南与性能优化
  • 2026北京装修行业获客新思路:家装/工装/设计工作室如何通过AIGEO低成本
  • 2026年8月层流净化车间/工业净化车间服务公司选哪家_优诺系统集成有限公司 - 行业平台推荐

最新新闻

  • TSB自定义技能系统:组件化开发与手机端性能优化实战
  • 新发传染病分子开关研究:从比较基因组学到宿主互作网络的全流程解析
  • CAN总线协议详解:从差分信号、仲裁机制到嵌入式系统通信实战
  • 2026年长沙地坪厂家推荐榜单,环氧地坪/密封固化剂地坪/金刚砂耐磨地坪/防静电地坪/防腐耐酸碱地坪/聚氨酯砂浆地坪匠心之选 - 优企名品
  • 大模型API调用中Token消耗异常分析与优化实战指南
  • 全国地形图DEM数据对比:SRTM15+、SRTM3、SRTM1与ASTER GDEM实战解析

日新闻

  • ClickHouse版本管理深度实战:4步构建零风险升级与回滚体系
  • Java 23 种设计模式:从踩坑到精通 | 番外:责任链模式 —— 物流审批流程实战
  • 华硕笔记本性能解放指南:G-Helper轻量级控制工具全面解析

周新闻

  • 大连理工大学与东京大学联手打造的“主动型AI助手“
  • 170.2026年国家级科研瓶颈:超精密单点金刚石切削(SPDT)光学表面生成
  • SongBloom:革命性歌曲生成框架深度解析——如何通过交织自回归与扩散模型创作完整音乐

月新闻

  • ClickHouse版本管理深度实战:4步构建零风险升级与回滚体系
  • Java 23 种设计模式:从踩坑到精通 | 番外:责任链模式 —— 物流审批流程实战
  • 华硕笔记本性能解放指南:G-Helper轻量级控制工具全面解析

关于尧图

  • 公司简介
  • 团队介绍
  • 企业文化
  • 荣誉资质

服务项目

  • 定制开发
  • 电商建站
  • UI 设计
  • 运维服务

快速链接

  • 案例展示
  • 建站流程
  • 常见问题
  • 资讯中心

联系方式

  • 📍北京市朝阳区互联网产业园 A 座 10 层
  • 📞400-888-8888
  • ✉️contact@rkmt.cn
  • 🕐周一至周日 9:00-21:00

© 2024 北京尧图网络科技有限公司 版权所有 | 京 ICP 备 XXXXXXXX 号