☰
TensorFlow 2 图片数据建模全流程实战:基于 tf.data 管道构建 Cifar2 飞机与汽车图像分类模型
2026/9/25 1:23:27 网站建设 项目流程
  • 教程
  • 深度学习
  • 机器学习

【免费下载链接】eat_tensorflow2_in_30_days

Tensorflow2.0 🍎🍊 is delicious, just eat it! 😋😋

项目地址:https://gitcode.com/gh_mirrors/ea/eat_tensorflow2_in_30_days
点击查看免费下载

本文来自开源仓库 eat_tensorflow2_in_30_days(30 天吃掉那个 TF2.0)系列教程的第 1-2 章,完整演示从图片文件到可部署模型的端到端建模流程:使用 TensorFlow 原生tf.data.Dataset搭配tf.image构建高性能图片数据管道,用函数式 API 搭建卷积神经网络,并通过fit内置训练、TensorBoard 可视化、模型评估与 SavedModel 导出完成一个 Cifar2 二分类任务。读完本文,你将掌握图片分类任务中"数据准备 → 模型定义 → 训练 → 评估 → 预测 → 部署"六个环节的完整实操方案,以及num_parallel_calls、prefetch、AUTOTUNE等数据管道性能优化手段。

一、任务与数据集:什么是 Cifar2

Cifar2 数据集是 Cifar10 数据集的子集,只保留前两种类别airplane(飞机)和automobile(机动车)。本仓库已预先准备好该数据集的目录结构(见 data/cifar2):

  • 训练集:airplane 与 automobile 图片各 5000 张(共 10000 张);
  • 测试集:airplane 与 automobile 图片各 1000 张(共 2000 张);
  • 图片均为 32×32 的 JPEG 格式。

任务目标非常明确:训练一个二分类模型,对输入的飞机 / 汽车图片给出正确的类别判别。由于每个类别单独放在一个子目录中,文件路径本身就携带了标签信息,这为后续用tf.strings.regex_full_match从路径解析标签提供了天然便利。

二、准备数据:两种常用图片数据方案对比

在 TensorFlow 中准备图片数据通常有两种方案:

  1. ImageDataGenerator 方案:使用tf.keras.preprocessing.image.ImageDataGenerator构建图片数据生成器,配置简单、上手快,适合快速原型验证。例如仓库 5-1,数据管道Dataset.md 中演示的flow_from_directory方式,通过rescale=1.0/255归一化并从./data/cifar2/test/按目录自动生成标签:

    image_generator = ImageDataGenerator(rescale=1.0/255).flow_from_directory( "./data/cifar2/test/", target_size=(32, 32), batch_size=20, class_mode='binary') classdict = image_generator.class_indices # 如 {'airplane': 0, 'automobile': 1}
  2. tf.data + tf.image 原生方案:使用tf.data.Dataset搭配tf.image中的图片处理方法构建数据管道。这是 TensorFlow 的原生方法,更加灵活,使用得当可以获得更好的性能,也是本章主角。

本章选择第二种方案。它之所以性能更优,是因为tf.data管道具备多进程并行转换、预取、缓存等能力,详见后文"数据管道性能优化"小节。

三、构建图片数据管道:从文件路径到张量批次

3.1 图片加载与标签解析函数

首先定义load_image函数,完成"读取文件 → 解码 → 缩放 → 归一化"的完整转换,并从文件路径中解析出标签:

import tensorflow as tf from tensorflow.keras import datasets, layers, models BATCH_SIZE = 100 def load_image(img_path, size=(32, 32)): label = tf.constant(1, tf.int8) if tf.strings.regex_full_match(img_path, ".*automobile.*") \ else tf.constant(0, tf.int8) img = tf.io.read_file(img_path) img = tf.image.decode_jpeg(img) # 注意此处为 jpeg 格式 img = tf.image.resize(img, size) / 255.0 return (img, label)

逐行解读这个函数的关键设计:

  • tf.strings.regex_full_match(img_path, ".*automobile.*"):用正则匹配路径中是否含automobile,从而把标签确定为 1(机动车)或 0(飞机)。标签来自路径而非额外标注文件,这是目录式数据集常见的做法;
  • tf.io.read_file:在 TensorFlow 图内直接读取文件内容,返回字符串张量;
  • tf.image.decode_jpeg:将 JPEG 字节解码为像素张量,注意此处数据为 jpeg 格式;若数据集是 PNG 或其他格式,需相应换成decode_png等解码器;
  • tf.image.resize(img, size) / 255.0:统一缩放为 32×32 并将像素值归一化到 [0, 1] 区间,缩小数据量级、加速收敛;
  • tf.constant(1, tf.int8):标签以tf.int8常量返回,避免 Python 布尔/整数与图执行之间的类型转换开销。

补充:仓库 5-1,数据管道Dataset.md 中给出了这个函数的另一种写法:label = 1 if tf.strings.regex_full_match(img_path, ".*/automobile/.*") else 0。两种正则都能工作,前者更宽松(只要路径出现 automobile 即视为正类),后者限定目录层级。可依据实际目录命名灵活选择。

3.2 构建训练集与测试集管道

# 使用并行化预处理 num_parallel_calls 和预存数据 prefetch 来提升性能 ds_train = tf.data.Dataset.list_files("./data/cifar2/train/*/*.jpg") \ .map(load_image, num_parallel_calls=tf.data.experimental.AUTOTUNE) \ .shuffle(buffer_size=1000).batch(BATCH_SIZE) \ .prefetch(tf.data.experimental.AUTOTUNE) ds_test = tf.data.Dataset.list_files("./data/cifar2/test/*/*.jpg") \ .map(load_image, num_parallel_calls=tf.data.experimental.AUTOTUNE) \ .batch(BATCH_SIZE) \ .prefetch(tf.data.experimental.AUTOTUNE)

管道各环节的作用如下:

环节作用
list_files("./data/cifar2/train/*/*.jpg")按通配符递归收集所有 JPEG 文件路径,*/*.jpg表示"每个类别子目录下的所有 jpg"
.map(load_image, num_parallel_calls=AUTOTUNE)对每个路径执行加载函数;设置num_parallel_calls让转换多进程并行执行
.shuffle(buffer_size=1000)用容量 1000 的缓冲区打乱样本顺序(仅训练集需要,测试集不需要)
.batch(BATCH_SIZE)将样本聚合成大小为 100 的批次
.prefetch(AUTOTUNE)让数据准备与参数迭代两个过程相互并行,提前预取下一批数据

tf.data.experimental.AUTOTUNE让 TensorFlow自动选择合适的并行度与预取缓冲大小,无需手工调参。验证管道形状:

for x, y in ds_train.take(1): print(x.shape, y.shape) # 输出:(100, 32, 32, 3) (100,)

即一个批次包含 100 张 32×32×3 的图片与 100 个标签。可视化部分样本:

%matplotlib inline %config InlineBackend.figure_format = 'svg' from matplotlib import pyplot as plt plt.figure(figsize=(8, 8)) for i, (img, label) in enumerate(ds_train.unbatch().take(9)): ax = plt.subplot(3, 3, i + 1) ax.imshow(img.numpy()) ax.set_title("label = %d" % label) ax.set_xticks([]) ax.set_yticks([]) plt.show()

3.3 数据管道性能优化原理(仓库源码级佐证)

模型训练耗时主要来自两个部分:数据准备与参数迭代。参数迭代依赖 GPU 加速,而数据准备则可以通过高效的数据管道来提速。仓库 5-1,数据管道Dataset.md 针对本章用到的优化手段给出了底层原理与基准实验:

  1. prefetch并行化:数据准备与参数迭代默认是串行的(总耗时 ≈ 准备耗时 + 迭代耗时);加上prefetch后二者并行执行,总耗时约为二者最大值。文档中用一个"准备每步 2s、训练每步 1s、共 10 步"的模拟实验证明:串行约 30s,prefetch后约 20s;
  2. num_parallel_calls多进程转换:map(load_image)是单进程转换,map(load_image, num_parallel_calls=tf.data.experimental.AUTOTUNE)则让转换过程多进程执行,对图片解码、缩放这类 CPU 密集操作收益尤其明显(文档对 cifar2 10000 张图片做了单进程 vs 多进程转换的耗时对比实验);
  3. cache缓存:如果数据集不大,可在第一个 epoch 后把数据缓存到内存,后续 epoch 直接复用。仓库 6-2,训练模型的3种方法.md 中的 reuters 示例即采用prefetch(...).cache()的组合写法;
  4. 先 batch 再向量化 map:转换时优先对批次做向量化运算而非逐样本标量运算,可显著降低 Python 层开销。

这些手段正是本章管道写法"map(并行) → shuffle → batch → prefetch"的性能依据,也是"tf.data 方案性能更优"的根源。

四、定义模型:用函数式 API 构建卷积神经网络

Keras 接口构建模型有三种方式:

  1. Sequential 顺序模型:按层顺序堆叠,适合纯线性结构(示例见 6-1,构建模型的3种方法.md);
  2. 函数式 API:构建任意结构模型(多输入、多输出、共享权重、残差连接等),本章采用此方式;
  3. Model 子类化:继承Model基类自定义模型,灵活性最高但出错概率也更高。

本章模型结构为「卷积-池化-卷积-池化-丢弃-展平-全连接-输出」的经典小 CNN:

tf.keras.backend.clear_session() # 清空会话,避免多轮实验间的计算图污染 inputs = layers.Input(shape=(32, 32, 3)) x = layers.Conv2D(32, kernel_size=(3, 3))(inputs) x = layers.MaxPool2D()(x) x = layers.Conv2D(64, kernel_size=(5, 5))(x) x = layers.MaxPool2D()(x) x = layers.Dropout(rate=0.1)(x) x = layers.Flatten()(x) x = layers.Dense(32, activation='relu')(x) outputs = layers.Dense(1, activation='sigmoid')(x) model = models.Model(inputs=inputs, outputs=outputs) model.summary()

各层职责与输出尺寸推演:

  • Input(shape=(32,32,3)):声明输入为 32×32×3 的图片张量;
  • Conv2D(32, kernel_size=(3,3)):32 个 3×3 卷积核提取局部特征,32×32 输入经无填充 3×3 卷积后变为 30×30,输出(None, 30, 30, 32),参数量 896 = 3×3×3×32 + 32(偏置);
  • MaxPool2D():默认 2×2 最大池化下采样,30×30 → 15×15;
  • Conv2D(64, kernel_size=(5,5)):64 个 5×5 卷积核,15×15 → 11×11,参数量 51264 = 5×5×32×64 + 64;
  • MaxPool2D():11×11 → 5×5;
  • Dropout(rate=0.1):以 10% 概率随机丢弃神经元,缓解过拟合;
  • Flatten():5×5×64 展平为 1600 维向量;
  • Dense(32, activation='relu'):32 个神经元全连接层,参数量 51232 = 1600×32 + 32;
  • Dense(1, activation='sigmoid'):输出 1 个值并经 sigmoid 映射到 [0,1],作为二分类正类概率,参数量 33 = 32 + 1。

模型总参数量103,425,全部可训练:

Model: "model" _________________________________________________________________ Layer (type) Output Shape Param # ================================================================= input_1 (InputLayer) [(None, 32, 32, 3)] 0 conv2d (Conv2D) (None, 30, 30, 32) 896 max_pooling2d (MaxPooling2D) (None, 15, 15, 32) 0 conv2d_1 (Conv2D) (None, 11, 11, 64) 51264 max_pooling2d_1 (MaxPooling2 (None, 5, 5, 64) 0 dropout (Dropout) (None, 5, 5, 64) 0 flatten (Flatten) (None, 1600) 0 dense (Dense) (None, 32) 51232 dense_1 (Dense) (None, 1) 33 ================================================================= Total params: 103,425 Trainable params: 103,425 Non-trainable params: 0

tf.keras.backend.clear_session()是每次建模前的好习惯:它清空当前会话与计算图,避免在多次运行 notebook 时出现变量重复定义、维度冲突等"脏会话"问题。

五、训练模型:内置 fit 方法与回调机制

训练模型通常有三种方法(详见 6-2,训练模型的3种方法.md):

  1. 内置fit方法:最常用最简单,支持 numpy array、tf.data.Dataset、Python generator 三种数据源,并可通过回调函数实现复杂训练控制;本章采用此法;
  2. 内置train_on_batch方法:更灵活,可在批次级别精细控制训练过程,比如中途调整学习率(仓库示例在第 5 个 epoch 将学习率减半);
  3. 自定义训练循环:无需编译模型,直接利用优化器根据损失函数反向传播迭代参数,灵活性最高。

注:fit_generator方法在 tf.keras 中已不推荐使用,其功能已被fit完全包含。

5.1 配置 TensorBoard 回调

import datetime import os stamp = datetime.datetime.now().strftime("%Y%m%d-%H%M%S") logdir = os.path.join('data', 'autograph', stamp) # 在 Python3 下建议使用 pathlib 修正各操作系统的路径 # from pathlib import Path # stamp = datetime.datetime.now().strftime("%Y%m%d-%H%M%S") # logdir = str(Path('./data/autograph/' + stamp)) tensorboard_callback = tf.keras.callbacks.TensorBoard(logdir, histogram_freq=1)

要点说明:

  • 用%Y%m%d-%H%M%S时间戳生成唯一日志目录,避免多次实验的日志互相覆盖(仓库 data/autograph 下的20200218-132438等目录即这种命名方式的实际产物);
  • histogram_freq=1表示每个 epoch 后记录权重直方图,便于在 TensorBoard 中观察参数分布变化;
  • 注释中给出了使用pathlib.Path的跨平台写法,在 Windows 等系统上比os.path.join更稳妥。

回调函数是model.fit训练控制的核心机制:BaseLogger(收集各批次指标均值)与History(记录到 history 并作为 fit 返回值)被所有模型默认添加;此外还有EarlyStopping(指标停滞时提前终止)、ModelCheckpoint(每 epoch 保存模型)、ReduceLROnPlateau(指标停滞时降低学习率)等常用内置回调,详见 5-8,回调函数callbacks.md。

5.2 编译与训练

model.compile( optimizer=tf.keras.optimizers.Adam(learning_rate=0.001), loss=tf.keras.losses.binary_crossentropy, metrics=["accuracy"] ) history = model.fit(ds_train, epochs=10, validation_data=ds_test, callbacks=[tensorboard_callback], workers=4)
  • 优化器:Adam(learning_rate=0.001),自适应矩估计优化器,是当下 CNN 训练的默认选择之一,默认学习率 0.001 通常无需改动;
  • 损失函数:binary_crossentropy(二元交叉熵),与"单输出 + sigmoid"的二分类输出层天然匹配;若输出层改为双神经元 softmax,则需改用SparseCategoricalCrossentropy;
  • 评估指标:accuracy(准确率),训练与验证阶段都会计算;
  • workers=4:使用 4 个并行工作线程执行数据管道,与管道内的AUTOTUNE并行转换叠加进一步压榨数据吞吐。

训练过程日志(部分):

Train for 100 steps, validate for 20 steps Epoch 1/10 100/100 [==============================] - 16s 156ms/step - loss: 0.4830 - accuracy: 0.7697 - val_loss: 0.3396 - val_accuracy: 0.8475 ... Epoch 10/10 100/100 [==============================] - 14s 142ms/step - loss: 0.1006 - accuracy: 0.9617 - val_loss: 0.1614 - val_accuracy: 0.9345
  • 10000 张训练图 / batch 100 = 100 个训练步;2000 张测试图 = 20 个验证步;
  • 10 个 epoch 后训练准确率升至 96.17%,验证准确率升至 93.45%,且验证损失持续下降,模型收敛良好、未见明显过拟合(Dropout 与数据量充足功不可没)。

六、评估模型:TensorBoard 与指标曲线

6.1 在 TensorBoard 中查看训练过程

%load_ext tensorboard # %tensorboard --logdir ./data/keras_model from tensorboard import notebook notebook.list() # 在 tensorboard 中查看模型 notebook.start("--logdir ./data/keras_model")

TensorBoard 回调把计算图、损失/指标曲线、权重直方图等写入logdir,之后可通过 Jupyter 的%tensorboard魔法命令或notebook.start拉起可视化面板(注意示例中notebook.list()返回的是训练日志目录,如./data/keras_model下按时间戳命名的各子目录)。

6.2 用 DataFrame 查看训练历史

model.fit返回的history.history是包含loss、accuracy、val_loss、val_accuracy的字典,转为 DataFrame 后按 epoch 索引,可直观浏览每个 epoch 的表现:

import pandas as pd dfhistory = pd.DataFrame(history.history) dfhistory.index = range(1, len(dfhistory) + 1) dfhistory.index.name = 'epoch' dfhistory

6.3 绘制 Loss 与 Accuracy 曲线

%matplotlib inline %config InlineBackend.figure_format = 'svg' import matplotlib.pyplot as plt def plot_metric(history, metric): train_metrics = history.history[metric] val_metrics = history.history['val_' + metric] epochs = range(1, len(train_metrics) + 1) plt.plot(epochs, train_metrics, 'bo--') plt.plot(epochs, val_metrics, 'ro-') plt.title('Training and validation ' + metric) plt.xlabel("Epochs") plt.ylabel(metric) plt.legend(["train_" + metric, 'val_' + metric]) plt.show() plot_metric(history, "loss") plot_metric(history, "accuracy")

6.4 用 evaluate 做最终评估

val_loss, val_accuracy = model.evaluate(ds_test, workers=4) print(val_loss, val_accuracy) # 输出:0.16139143370091916 0.9345

evaluate直接对ds_test全量测试集计算损失与指标,最终验证损失约 0.161、验证准确率约 93.45%,与训练日志中最后一个 epoch 的val_accuracy完全一致,印证了训练的稳定复现性。

七、使用模型:predict 与 predict_on_batch

训练好的模型有两种预测入口:

  1. model.predict(ds_test):对整个数据集(或任意可迭代数据源)逐批预测,返回概率数组;
  2. model.predict_on_batch(x):对一个批次张量直接预测,适合在线服务或逐批推理场景。
model.predict(ds_test) # 输出(节选): array([[9.9996173e-01], [9.5104784e-01], [2.8648047e-04], ..., [1.1484033e-03], [3.5589080e-02], [9.8537153e-01]], dtype=float32)
for x, y in ds_test.take(1): print(model.predict_on_batch(x[0:20])) # 输出(节选):shape=(20, 1) 的 tf.Tensor,每个元素为该样本属于 automobile 的概率

输出值为 sigmoid 概率(越接近 1 越可能是 automobile,越接近 0 越可能是 airplane)。由于二分类输出是单神经元 sigmoid,实际类别可通过predict >= 0.5阈值化得到。

八、保存模型:权重、SavedModel 与跨平台部署

推荐使用TensorFlow 原生方式保存模型,有两种粒度:

8.1 仅保存权重(Checkpoint)

# 保存权重,该方式仅仅保存权重张量 model.save_weights('./data/tf_model_weights.ckpt', save_format="tf")

仅持久化权重张量,不包含模型结构;恢复时需要用相同的代码重建模型结构再load_weights。仓库根目录下的 tf_model_weights.ckpt.data-00000-of-00001 与 tf_model_weights.ckpt.index 即该方式的产物(TF 2.x 的 checkpoint 采用"数据文件 + 索引文件"的 Sharded 存储格式)。

8.2 保存完整模型(SavedModel,跨平台可部署)

# 保存模型结构与模型参数到文件, 该方式保存的模型具有跨平台性便于部署 model.save('./data/tf_model_savedmodel', save_format="tf") print('export saved model.') model_loaded = tf.keras.models.load_model('./data/tf_model_savedmodel') model_loaded.evaluate(ds_test) # 输出:[0.16139124035835267, 0.9345]

save_format="tf"导出的是SavedModel 格式,同时保存模型结构与参数,加载后无需重建代码即可直接evaluate、predict。仓库 data/tf_model_savedmodel 目录完整展示了 SavedModel 的标准产物:saved_model.pb(计算图与元数据)、variables/(权重数据分片,如variables.data-00000-of-00001与variables.index)、assets/(附加资产文件,如 saved_model.json)。

SavedModel 的跨平台价值在部署环节体现得最充分:仓库 6-6,使用tensorflow-serving部署模型.md 展示了用model.save(export_path + version, save_format="tf")按版本号导出模型,并通过saved_model_cli show --dir {export_path} --all检查模型签名,进而交给 TensorFlow Serving 提供在线推理服务;6-7,使用spark-scala调用tensorflow模型.md 则演示了 Scala/Spark 侧加载同一格式模型进行预测。可见「训练导出 SavedModel → 多端复用」是本章保存方式选择的核心考量。

九、小结与扩展路径

至此,一个完整的图片二分类建模流程闭环完成:文件路径 → tf.data 高性能管道 → 函数式 API 卷积网络 → fit 内置训练 → TensorBoard / 曲线可视化评估 → predict 推理 → SavedModel 导出部署,最终在 Cifar2 测试集上达到约 93.45% 的验证准确率。

如果希望继续深入本主题,仓库内配套资源可以按需研读:

  • 数据管道性能调优全解析:5-1,数据管道Dataset.md(prefetch、interleave、num_parallel_calls、cache的基准实验);
  • 三种建模方式对比(Sequential / 函数式 API / Model 子类化):6-1,构建模型的3种方法.md;
  • 三种训练方式对比(fit / train_on_batch / 自定义循环):6-2,训练模型的3种方法.md;
  • 回调函数机制与常用回调: 5-8,回调函数callbacks.md;
  • 模型跨平台部署(TensorFlow Serving、Spark-Scala 调用):6-6,使用tensorflow-serving部署模型.md、6-7,使用spark-scala调用tensorflow模型.md。

实践提示:直接运行本仓库中的示例代码时,请将./data/cifar2/...等路径替换为仓库实际位置(如 data/cifar2),并确保已安装 TensorFlow 2.x 及 matplotlib、pandas、tensorboard 等依赖。

  • 教程
  • 深度学习
  • 机器学习

【免费下载链接】eat_tensorflow2_in_30_days

Tensorflow2.0 🍎🍊 is delicious, just eat it! 😋😋

项目地址:https://gitcode.com/gh_mirrors/ea/eat_tensorflow2_in_30_days
点击查看免费下载

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询