news 2026/10/10 22:02:55

T10-数据增强

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
T10-数据增强

● 🍨 本文为🔗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),另外也可以自定义函数进行数据增强,比如上文的随机调整图片的对比度。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/10/10 22:01:44

可视化运维监控实战:从故障可见到可控可定位的完整体系搭建

干了这么多年运维,我最大的感触是:系统出故障不可怕,可怕的是故障发生之后,你盯着满屏的告警邮件,却说不清楚现在到底是什么挂了、影响多大、从哪开始查。那种“明明知道出事了,却无从下手”的无力感&#…

作者头像 李华
网站建设 2026/10/10 21:59:39

恶意软件逆向全流程:从加壳样本到 Ghidra 里的真相

恶意软件逆向全流程:从加壳样本到 Ghidra 里的真相 【免费下载链接】ghidra Ghidra is a software reverse engineering (SRE) framework 项目地址: https://gitcode.com/GitHub_Trending/gh/ghidra 拿到一个加了 UPX 壳的勒索软件样本,第一反应是…

作者头像 李华
网站建设 2026/10/10 21:56:50

模板代码调试技巧:模板字符串、Twig、STM32三大场景实战

写模板代码这事,说起来有点意思。你从网上或者同事手里拿到的“模板”,本意是拿来就能跑、省得从零开始,但真到改出问题的时候,往往比直接写还难受。尤其是标题里那些关键词串起来之后——模板字符串、twig模板手册、STM32工程模板…

作者头像 李华
网站建设 2026/10/10 21:51:27

专科生论文降AI率工具避坑指南:8类工具原理与正确用法

专科生写论文最头疼的事,除了查重红标,这两年又多了个“AI率”。明明是自己一个字一个字敲的,交上去却显示“疑似AI生成”,轻则退回修改,重则影响答辩资格。于是“降AI率工具”成了热门搜索词,但市面上的工…

作者头像 李华