ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TensorFlow 2.x+Keras从零入门:环境搭建到工程实践

TensorFlow 2.x+Keras从零入门:环境搭建到工程实践 在实际深度学习项目中TensorFlow 仍然是一个绕不开的名字。很多新手在 2025 年学习 TensorFlow 时会被安装过程、版本差异、Keras API 选择、GPU 配置和大量低质量教程搞得一头雾水。真正有用的一篇入门教程应该带读者从环境准备开始一步步跑通一个最小可训练的模型并解释每一步背后的原理而不是只展示一段能运行的代码。本文就以 TensorFlow 2.x 和 Keras 为主线从虚拟环境搭建、安装验证、张量概念、模型构建、数据管道、训练验证到常见错误排查和工程化实践完整走一遍入门到可上手的路径。这篇文章适合以下几类读者第一次接触 TensorFlow、准备在本地跑通第一个深度学习模型的新手已经熟悉 PyTorch但因为项目或工作需要了解 TensorFlow 的开发者以及想从零开始搭建一套可维护的训练代码而不只是复制笔记本示例教程的工程技术人员。阅读并复现本文示例后你会得到一个可运行、可扩展的 TensorFlow 项目雏形也能在遇到安装报错、训练不收敛、显存不足等问题时按照清晰的排查链路定位原因。1. 为什么 2025 年还要认真学 TensorFlow1.1 从“框架之争”看 TensorFlow 的实际位置TensorFlow 和 PyTorch 的使用率对比最近几年一直有各种讨论。PyTorch 在研究社区中确实很流行动态图和灵活的表达方式让算法验证非常方便。但从工程角度看TensorFlow 在服务端推理、移动端部署、TFLite、TF Serving、存量生产系统、模型版本管理和多语言推理等领域仍有大量真实项目依赖。很多公司现有的图像分类、推荐排序、文本匹配系统底层跑的还是 TensorFlow。因此选型时不应该只听“哪个框架更流行”而要看项目所在的生态。如果你的目标是快速验证研究想法PyTorch 可能顺手如果你的任务是把模型接进已有的推荐系统、部署到移动端或者公司技术栈里已经有 TensorFlow 服务那学习 TensorFlow 就是刚需。作为学习者理解 TensorFlow 2.x 的编程模型也能反过来帮助你更深入理解 PyTorch、JAX 等框架的共同设计思路因为它们都在解决张量计算、自动微分和模型训练这三个核心问题。1.2 本文采用的技术路线TensorFlow 2.x KerasTensorFlow 最初版本 1.x 的编程方式比较复杂需要手动构建计算图、定义 Session 再执行学习曲线很陡。TensorFlow 2.x 最大的变化是默认开启 Eager Execution动态执行编程体验更接近普通 Python同时把 Keras 作为高层 API 集成进来。现在的常见开发方式不再需要手写繁琐的占位符、会话和图结构可以直接用tf.keras.Sequential或函数式 API 搭建模型。本文的所有示例都基于这条路线使用 Python 环境管理工具创建独立虚拟环境安装 TensorFlow 2.x用 Keras 构建模型用model.fit完成训练用model.evaluate评估用model.predict做推理。这个流程是 TensorFlow 目前最常见的入门路径也是生产代码的最小骨架。理解这套流程后再去看复杂的自定义训练循环、分布式策略和模型服务就不会觉得无从下手。1.3 学习环境与生产环境从第一天就要分开很多入门教程只教你安装 TensorFlow然后直接运行一段示例没有区分“本地学习环境”和“生产运行环境”。这个习惯会在项目落地时付出代价。学习环境追求快速跑通可以使用 CPU 版 TensorFlow、小数据集、简单模型生产环境则要考虑 GPU 或 TPU 加速、稳定版本锁定、训练与推理分离、模型版本管理、日志监控、异常回滚等。从第一天就应该养成隔离环境的习惯。不要把所有依赖都装到系统 Python 里也不要用pip install tensorflow装完就忘记记录版本。建议每个项目使用独立虚拟环境并维护requirements.txt或pyproject.toml。下面从安装环境的准备开始讲这是整个学习路径里最常见也最容易出错的一步。对比维度学习环境生产环境核心目标快速验证代码逻辑稳定运行、可监控、可回滚硬件选择CPU 可以先跑通按业务需求选择 GPU/CPU版本管理最新稳定版即可锁定精确版本并验证兼容性数据处理直接加载内置数据集构建可复现的数据管道模型存储保存权重文件导出为 SavedModel配合版本服务监控要求观察 loss/accuracy记录指标、日志、异常告警2. 安装 TensorFlow 之前先把环境隔离这件事做对2.1 环境检查和版本选择安装 TensorFlow 前先确认当前机器上的 Python 版本、操作系统和硬件情况。TensorFlow 对 Python 版本有明确支持范围通常不会支持所有最新版本的 Python。千万不要安装完毕后出现“找不到匹配版本”再回头检查 Python正确顺序是先确认支持矩阵再创建虚拟环境再安装。在 Windows、Linux 和 macOS 上安装方式会有些差异。Windows 上通常使用原生 PythonLinux 上要区分系统 Python 和虚拟环境macOS 还需要注意 Apple Silicon 和 Intel 芯片的差异。TensorFlow 官方会为不同平台提供不同的 wheel 包所以遇到安装失败时先检查自己是否下载了对应平台的包。这里推荐使用 Python 自带的venv模块不额外安装第三方虚拟环境工具减少出问题的环节。如果你已经在使用 Anaconda 或 Miniconda也可以用 conda 创建环境但核心原则相同不要污染全局 Python。2.2 用 Python 虚拟环境隔离项目依赖创建虚拟环境并安装 TensorFlow命令如下python -m venv tf-env在 Windows 上激活环境tf-env\Scripts\activate在 Linux 或 macOS 上激活环境source tf-env/bin/activate激活后命令行前缀会显示当前环境名。这时再用pip安装任何 Python 包都会安装到这个虚拟环境中而不会影响系统其他 Python 项目。安装 TensorFlow 时建议先升级 pippython -m pip install --upgrade pip然后是核心安装命令pip install tensorflow如果你网络较慢可以使用国内镜像源。以清华 PyPI 镜像为例pip install tensorflow -i https://pypi.tuna.tsinghua.edu.cn/simple安装过程中尽量不要混用多个镜像源也不要遇到一个包装不上就反复强制安装容易导致依赖版本错乱。安装完成后先不要急着写代码先验证安装是否成功。2.3 CPU 版和 GPU 版的取舍TensorFlow 分为 CPU 版本和 GPU 版本。在 TensorFlow 2.x 中CPU 版和 GPU 版默认使用同一个tensorflow包名系统检测到可用 GPU 和相关依赖后会自动加速不再像 1.x 时期需要单独安装tensorflow-gpu包。因此很多人不需要特意安装 GPU 版本。但 GPU 能否被 TensorFlow 使用并不仅仅取决于包名还取决于显卡驱动、CUDA、cuDNN 的版本是否匹配。TensorFlow 官方会给出对应版本的 CUDA 和 cuDNN 要求安装前一定要查阅对应版本的兼容性列表。不要凭猜测安装最新版 CUDA最新版不一定被当前 TensorFlow 支持。在本地学习时建议先用 CPU 版跑通所有代码。CPU 训练虽然慢但对于 MNIST 这样的小数据集足够完成任务。等需要训练更大模型时再考虑 GPU 环境。GPU 环境遇到问题时先用nvidia-smi查看显卡驱动是否正常再确认 TensorFlow 是否检测到 GPU这在后面会进一步展开。2.4 验证安装是否可用的三个命令安装完成后进入 Python 解释器依次执行以下检查import tensorflow as tf print(tf.__version__)正常会输出 TensorFlow 的版本号例如2.18.0。接下来检查 Keras 是否可用print(tf.keras.__version__)如果你安装的是 TensorFlow 2.18这里的输出可能与 TensorFlow 版本不同因为 Keras 已经作为独立包发布。只要导入不报错就说明高层 API 可用。再检查能否进行最基本的张量计算a tf.constant([[1.0, 2.0], [3.0, 4.0]]) b tf.constant([[2.0, 0.0], [0.0, 2.0]]) print(a b)这是一个 2x2 矩阵乘法正常会得到一个 2x2 的张量。如果这三个检查都通过说明 TensorFlow 核心功能已经可用。此时如果机器有 NVIDIA GPU可以继续检查 GPU 是否可见print(tf.config.list_physical_devices(GPU))如果列表为空说明 TensorFlow 没有检测到可用于加速的 GPU需要检查驱动、CUDA 版本或安装方式。这个检查在后续训练大模型时非常重要。安装场景推荐方式注意事项Windows 学习环境python -m venv CPU 版Python 版本务必在支持范围内Linux 生产环境容器或虚拟环境 锁定版本GPU 需要匹配 CUDA/cuDNN 版本macOS Apple Silicon查看官方是否有对应 wheel部分版本依赖 Rosetta 或 arm64 支持已有 Anacondaconda create -n tf-env python3.11安装后仍用 pip 装 tensorflow3. 先建立最小可运行的模型从张量到训练闭环3.1 张量是什么为什么 TensorFlow 叫 TensorFlowTensorFlow 的核心处理对象是张量Tensor。你可以把张量理解为一个多维数组0 维是标量1 维是向量2 维是矩阵3 维及以上可以表示图像、视频、序列等更复杂的数据。深度学习中几乎所有的计算都是在张量上完成的例如一张 28x28 的灰度图片可以表示成一个[28, 28]的张量一批 32 张图片可以表示成[32, 28, 28]。TensorFlow 的名字中“Tensor” 指张量而 “Flow” 指数据流动。在 1.x 时代开发者需要先定义计算图再让数据在图里流动。2.x 默认使用 Eager Execution计算立即执行结果立即可见理解起来更容易。但底层仍然有计算图的概念Keras 会自动把模型操作转化为计算图这样能够自动求导、优化和部署。理解张量后你就能理解为什么很多深度学习的代码都在做形状变换、维度扩展和归一化。例如全连接层输入要求二维[batch_size, features]而卷积层输入要求四维[batch_size, height, width, channels]。如果形状不对TensorFlow 会直接报错。3.2 用 Keras Sequential 构建一个分类模型构建一个最小可运行的图像分类模型最常用的数据集是 Fashion MNIST。先加载数据并做简单预处理import tensorflow as tf (x_train, y_train), (x_test, y_test) tf.keras.datasets.fashion_mnist.load_data() # 把像素值从 0-255 缩放到 0-1有利于梯度下降收敛 x_train x_train.astype(float32) / 255.0 x_test x_test.astype(float32) / 255.0 # 增加通道维度把 [28, 28] 变成 [28, 28, 1] x_train x_train[..., tf.newaxis] x_test x_test[..., tf.newaxis] print(x_train.shape, y_train.shape)Fashion MNIST 每张图片是 28x28 的灰度图共有 10 个类别。这里先不直接使用卷积层而是先使用全连接网络减少初学时的概念负担。构建模型model tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape(28, 28, 1)), tf.keras.layers.Dense(128, activationrelu), tf.keras.layers.Dense(10, activationsoftmax) ])Flatten把 28x28x1 的图片展平成 784 维向量第一个Dense有 128 个神经元激活函数使用 ReLU最后一个Dense有 10 个神经元对应 10 个类别激活函数使用 softmax输出每个类别的概率。3.3 编译、训练、评估三个 API 完成闭环模型刚构建时权重是随机初始化的必须通过编译和训练来更新。编译时指定优化器、损失函数和评价指标model.compile( optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy] )这里选择优化器adam它适合大多数中小型模型。损失函数使用sparse_categorical_crossentropy因为标签是整数例如 0 到 9如果标签已经做了 one-hot 编码就需要用categorical_crossentropy。这个区别是新手最容易搞混的地方。训练模型history model.fit( x_train, y_train, epochs10, batch_size32, validation_data(x_test, y_test) )epochs表示整个数据集遍历多少轮batch_size表示每次计算梯度使用多少条样本。训练过程中会输出每个 epoch 的 loss 和 accuracy以及验证集上的对应指标。训练完成后评估test_loss, test_acc model.evaluate(x_test, y_test) print(fTest accuracy: {test_acc:.4f})到这里一个最小训练闭环已经完成。很多问题会出现在这个阶段比如 loss 不下降、精度不变、显存不足等后面会单独分析。compile 参数作用常见选择optimizer控制权重更新方式adam、sgd、rmsproploss度量模型输出与真实标签差距sparse_categorical_crossentropy、msemetrics训练时额外展示的指标accuracy、mae、自定义指标4. 核心 API 和技术细节模型构建不再停留在示例4.1 从 Sequential 到函数式 API什么时候需要换Sequential适合线性堆叠的模型也就是一层接一层的结构。但真实项目经常有多输入、多输出、分支结构、残差连接等需求。例如一个模型同时输入图片和文本或者一个网络在一定层之后分支出两个输出头。这时继续使用Sequential会非常吃力应该使用函数式 API。函数式 API 的核心思路是把每一层当作一个函数调用后返回一个新的张量最后用tf.keras.Model把输入和输出组合成模型。示例inputs tf.keras.Input(shape(28, 28, 1)) x tf.keras.layers.Flatten()(inputs) x tf.keras.layers.Dense(128, activationrelu)(x) x tf.keras.layers.Dropout(0.2)(x) outputs tf.keras.layers.Dense(10, activationsoftmax)(x) model tf.keras.Model(inputsinputs, outputsoutputs)这段代码比 Sequential 更灵活因为可以在任意位置插入分支、拼接多个输入也可以将某一层的输出额外送到另一个分支。函数式 API 仍然是 Keras 的声明式写法适合大多数相对固定的网络结构。如果需要在训练时动态改变结构或者研究完全自定义的训练逻辑可以使用 Keras 的Model子类化方式继承tf.keras.Model并重写call方法。这种写法更接近 PyTorch 的编程体验但代价是模型序列化、部署时可能遇到额外限制。实际项目中优先用函数式 API不要一开始就上子类化。4.2 Dataset 数据管道把数据加载和预处理标准化实际项目中数据不会像内置数据集那样一次性全部载入内存。更常见的做法是有一个文件目录里面是图片文件文件名或子目录名对应标签。tf.data.Dataset是 TensorFlow 推荐的数据管道工具它可以把文件读取、解码、预处理、打乱、批量、预取等操作组织起来。一个典型流程如下def load_image(image_path, label): image tf.io.read_file(image_path) image tf.image.decode_jpeg(image, channels3) image tf.image.resize(image, [224, 224]) image tf.cast(image, tf.float32) / 255.0 return image, label dataset tf.keras.utils.image_dataset_from_directory( data/train, label_modeint, image_size(224, 224), batch_size32, shuffleTrue )或者手动构建image_paths [data/train/1.jpg, data/train/2.jpg] labels [0, 1] dataset tf.data.Dataset.from_tensor_slices((image_paths, labels)) dataset dataset.map(load_image, num_parallel_callstf.data.AUTOTUNE) dataset dataset.batch(32).prefetch(tf.data.AUTOTUNE)map用于逐样本做预处理batch把数据打包成批次prefetch让数据加载与模型训练并行。实际训练中prefetch能明显减少数据加载造成的等待。理解tf.data的关键是不要在每个 epoch 里手动写 for 循环读文件而是把数据转换流程声明好让框架自动调度。4.3 回调函数模型检查点、早停和学习率调整训练时只调用model.fit还不够。模型可能训练到第 5 轮就开始过拟合或者中途停电导致训练白跑。这时需要回调函数Callback来干预训练过程。常用的回调包括ModelCheckpoint、EarlyStopping、ReduceLROnPlateau。示例callbacks [ tf.keras.callbacks.ModelCheckpoint( model.weights.h5, save_best_onlyTrue, save_weights_onlyTrue, monitorval_loss ), tf.keras.callbacks.EarlyStopping( monitorval_loss, patience3, restore_best_weightsTrue ), tf.keras.callbacks.ReduceLROnPlateau( monitorval_loss, factor0.5, patience2 ) ] model.fit( train_dataset, epochs50, validation_dataval_dataset, callbackscallbacks )ModelCheckpoint在验证集 loss 最好时保存权重避免最后一轮模型过拟合导致最佳权重丢失。EarlyStopping在验证集指标连续多个 epoch 不再提升时停止训练节省时间。ReduceLROnPlateau在指标卡住时自动降低学习率比手动调整更稳定。回调作用推荐场景ModelCheckpoint保存最优或定期权重训练时间较长、需要恢复训练EarlyStopping防止过拟合节省算力模型已经收敛但 epochs 设置过大ReduceLROnPlateau自动降低学习率loss 下降趋缓时TensorBoard记录训练指标和计算图需要可视化监控训练过程5. 跑通验证用真实小数据集完成训练与预测5.1 用 MNIST 还是 Fashion MNIST很多教程使用 MNIST 手写数字数据集但它过于简单随便一个线性模型都能达到 90% 以上的准确率难以体现特征学习的作用。这里推荐使用 Fashion MNIST它是 MNIST 的现代替代品同样是 28x28 灰度图但分类目标是服装类别难度更高更有实际意义。Fashion MNIST 有 10 个类别T 恤、裤子、套衫、裙子、外套、凉鞋、衬衫、运动鞋、包、短靴。类别之间有一些相似之处例如 T 恤和衬衫容易混淆因此模型需要学习更有区分度的特征。5.2 完整训练脚本与预期输出为了便于复现下面给出一个完整脚本包含数据加载、模型构建、训练、评估和预测import tensorflow as tf # 1. 加载数据 (x_train, y_train), (x_test, y_test) tf.keras.datasets.fashion_mnist.load_data() # 2. 预处理 x_train x_train.astype(float32) / 255.0 x_test x_test.astype(float32) / 255.0 x_train x_train[..., tf.newaxis] x_test x_test[..., tf.newaxis] # 3. 构建模型 model tf.keras.Sequential([ tf.keras.layers.Flatten(input_shape(28, 28, 1)), tf.keras.layers.Dense(128, activationrelu), tf.keras.layers.Dropout(0.2), tf.keras.layers.Dense(10, activationsoftmax) ]) # 4. 编译 model.compile( optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy] ) # 5. 训练 model.fit( x_train, y_train, epochs10, batch_size32, validation_split0.2 ) # 6. 评估 test_loss, test_acc model.evaluate(x_test, y_test) print(fTest accuracy: {test_acc:.4f}) # 7. 预测 predictions model.predict(x_test[:5]) predicted_classes tf.argmax(predictions, axis-1).numpy() print(Predicted:, predicted_classes) print(Ground truth:, y_test[:5])运行这个脚本每个 epoch 会输出类似下面的信息Epoch 1/10 1500/1500 [] - 3s 2ms/step - loss: 0.5326 - accuracy: 0.8122 - val_loss: 0.4231 - val_accuracy: 0.8511 ... Epoch 10/10 1500/1500 [] - 3s 2ms/step - loss: 0.3155 - accuracy: 0.8872 - val_loss: 0.3621 - val_accuracy: 0.8769不同硬件和 TensorFlow 版本下数字会有波动但趋势应该是 loss 逐步下降accuracy 逐步上升最终测试准确率通常在 0.86 到 0.89 之间。如果你的结果远低于这个范围需要检查数据处理是否正确、网络结构是否合理、学习率是否合适。5.3 如何判断模型训练是否正常训练过程中不能只看最终准确率还要观察 loss 曲线的形态。正常情况下训练 loss 和验证 loss 都应该逐步下降然后趋于平缓。如果验证 loss 先下降后上升而训练 loss 仍然下降说明模型开始过拟合训练轮数过多或模型容量太大。如果训练 loss 几乎不变可能学习率太低、数据归一化有问题、初始化不稳定或梯度消失。另一个需要留意的现象是 loss 为NaN。常见原因是学习率过大导致梯度爆炸或者输入数据里有NaN。可以先检查输入数据是否包含空值再尝试降低学习率。训练初期出现轻微波动是正常的但几个 epoch 后仍然不下降就应该停下来调整。建议在训练脚本里加入TensorBoard回调保存训练曲线到日志目录然后在浏览器中观察 loss 和 accuracy。这样可以更直观地判断模型是否收敛。6. 常见问题排查安装、训练、版本不匹配6.1 安装时报错找不到匹配版本或下载慢现象是执行pip install tensorflow后报ERROR: Could not find a version that satisfies the requirement tensorflow或者下载速度很慢。首先确认 Python 版本。TensorFlow 对 Python 版本有支持范围如果 Python 版本太新可能还没有对应的 wheel 包。此时建议使用官方支持范围内的 Python 版本例如 Python 3.10 或 3.11通常兼容性更好。接着确认当前的 pip 是否可用python -m pip install --upgrade pip如果网络慢使用国内镜像源例如清华源或阿里源。不要在暂时失败后反复重新安装同一命令先看错误信息是网络问题还是版本问题。版本问题的提示通常是No matching distribution found网络问题通常是超时或连接错误。区分这两类错误能节省大量时间。6.2 启动时报错Could not load dynamic library 和 OneDNN 信息很多时候安装成功但导入 TensorFlow 时出现红色警告例如Could not load dynamic library libcudart.so.11.0如果你的机器没有独立显卡或者显卡驱动不匹配TensorFlow 会尝试加载 CUDA 相关库但失败。这个警告在 CPU 环境下不影响使用只是说明 TensorFlow 没有找到 GPU 加速库。此时可以忽略或者继续检查 GPU 配置。另一种常见输出是OneDNN custom operations are on. You may see slightly different numerical results...这也不是错误而是 TensorFlow 启用了 OneDNN 优化库的提示。OneDNN 是数学计算优化库能提升 CPU 上部分算子执行效率。看到这些信息不要紧张继续看后续代码是否能正常运行。如果确实需要 GPU 加速检查步骤是先运行nvidia-smi确认驱动能识别显卡再对照 TensorFlow 官方兼容列表确认 CUDA 和 cuDNN 版本最后在 Python 中运行tf.config.list_physical_devices(GPU)如果返回空列表说明 TensorFlow 仍未正确找到 GPU重点检查 CUDA 版本和 PATH 环境变量。6.3 训练时 loss 不下降或为 NaN这是训练阶段最常见的两类问题。loss 不下降的排查顺序是检查数据预处理。像素是否归一化到 0-1标签是否正确输入形状是否匹配网络输入。检查损失函数。整数标签用sparse_categorical_crossentropyone-hot 标签用categorical_crossentropy。检查学习率。Adam 默认学习率是 0.001如果数据量很小可以尝试 0.0001 或 0.0003。检查模型结构。层数过深或激活函数选择不当也可能导致训练困难。loss 为NaN时优先降低学习率比如从 0.001 降到 0.0001。再检查输入数据里是否有inf或NaN可以用tf.debugging.check_numerics来定位。某些数据集图片解码失败会输出异常值也需要在数据管道中过滤。6.4 显存不足与 CPU 训练慢训练大模型或调大 batch size 时GPU 显存不足会直接报ResourceExhaustedError。这时的解决方法是减小 batch size或者减小输入图片分辨率。如果模型太大也可以考虑减少网络宽度或层数。实际项目中不要盯着一个超大批次训练应该在显存允许的范围内选择合适的 batch size。CPU 训练慢不是错误但会明显影响调试效率。建议在开发阶段使用小数据集、少 epochs先跑通流程再全量训练。比如可以只加载原始数据的前 2000 条做模型调试x_train x_train[:2000] y_train y_train[:2000]这样可以快速定位代码逻辑错误而不是花半小时等一个 epoch 跑完后才发现前面有问题。问题现象常见原因检查方式处理建议安装时找不到包Python 版本不被支持python --version换用支持范围内的 Python 版本下载特别慢网络原因观察 pip 输出使用镜像源导入时提示 CUDA 库缺失GPU 依赖不匹配或没有 GPUnvidia-smi、官方兼容表CPU 环境可忽略GPU 环境对齐版本loss 不下降数据或超参数问题打印输入范围、标签分布归一化、调整学习率loss 为 NaN学习率过大或数据异常检查训练数据数值降低学习率、过滤异常值显存不足batch size 过大观察错误信息减小 batch size7. 从入门到可维护工程化实践和后续学习路线7.1 训练代码的文件目录结构入门阶段可以只写一个脚本但到项目阶段把数据加载、模型构建、训练、评估都塞进一个文件里会越来越难维护。推荐使用一个简单但清晰的目录结构project/ ├── config.py # 超参数和路径配置 ├── data_loader.py # 数据读取和预处理 ├── model.py # 模型结构定义 ├── train.py # 训练入口 ├── evaluate.py # 评估和推理脚本 └── models/ # 保存模型文件和日志config.py适合放学习率、batch size、epochs、数据路径等参数。data_loader.py负责返回已经预处理好的tf.data.Dataset。model.py只负责定义模型结构。train.py负责加载数据、构建模型、执行训练。这样做的好处是改网络结构时不需要改动数据代码改数据增强时不影响模型代码做实验时可以单独调整参数而不是在训练脚本里到处搜索变量。7.2 模型保存、导出与加载训练好的模型需要保存和加载。Keras 提供多种保存方式。最简单的保存整个模型model.save(fashion_mnist_model.keras)加载loaded_model tf.keras.models.load_model(fashion_mnist_model.keras)如果只保存权重不保存结构和优化器状态model.save_weights(model_weights.weights.h5)要加载权重到重新定义的模型中需要先构建相同的模型结构new_model tf.keras.Sequential([...]) # 与原始结构一致 new_model.load_weights(model_weights.weights.h5)生产部署场景更推荐导出为SavedModel格式因为它不依赖 Python 环境可以被 TensorFlow Serving、Java、Go 等推理环境使用。导出方式model.export(saved_model_dir)或使用传统方式tf.saved_model.save(model, saved_model_dir)在本地学习阶段至少掌握save和load_model能避免每次重新训练。7.3 超参数不要散落在代码里超参数是模型训练效果的重要变量。推荐把超参数集中在一个配置类或配置文件中。简单场景可以用 Python 的dataclassesfrom dataclasses import dataclass dataclass class Config: batch_size: int 32 epochs: int 50 learning_rate: float 1e-3 image_size: int 224 num_classes: int 10 config Config()训练脚本读取config.epochs和config.learning_rate而不是直接在model.fit里写数字。这样在做实验对比时可以通过修改配置快速尝试不同组合也便于记录每次实验所用参数。7.4 迁移学习在预训练模型基础上微调当数据量不足时从零训练一个大模型往往会过拟合。此时可以使用迁移学习加载在 ImageNet 上预训练过的模型冻结底部特征提取层只训练顶部分类层。TensorFlow 内置了ResNet50V2、MobileNetV2、EfficientNetV0等预训练模型。示例base_model tf.keras.applications.MobileNetV2( weightsimagenet, include_topFalse, input_shape(224, 224, 3) ) base_model.trainable False model tf.keras.Sequential([ base_model, tf.keras.layers.GlobalAveragePooling2D(), tf.keras.layers.Dense(10, activationsoftmax) ])这里weightsimagenet表示加载在 ImageNet 上训练好的权重include_topFalse去掉预训练模型的分类头然后添加自己的分类层。冻结底部层可以有效减少训练参数量也能借用预训练模型已经学到的通用特征。当模型在新数据集上有一定表现后可以选择解冻部分底层用较小学习率做微调进一步提升精度。7.5 发布前检查清单无论是本地作业还是生产项目在训练任务发布前都值得检查以下内容虚拟环境是否隔离依赖版本是否锁定。数据读取路径是否正确数据预处理是否一致。训练集、验证集、测试集是否切分合理有没有数据泄漏。是否存在明显类别不均衡是否需要调整损失函数权重。超参数是否集中在配置文件中是否记录了实验版本。回调函数是否保存了最优模型是否使用了早停。模型导出格式是否满足后续推理需求。是否在 CPU 小数据上先跑通过一个最小值再上全量数据。日志和监控是否到位训练中断能否恢复。代码和模型是否有统一版本记录方便回滚到之前的效果。这个清单可以避免很多“训练了一整晚结果发现数据加载错了”的惨痛教训。TensorFlow 的学习路径并不复杂先理解张量和自动微分再掌握 Keras 的基础建模方式然后学习数据管道和回调函数最后把代码工程化。核心不是背 API而是在一次次跑通训练和解决报错中建立直觉。沿着本文给出的示例往下走你可以在本地完成第一个 TensorFlow 训练任务再继续往函数式 API、迁移学习、模型部署方向扩展就能逐步进入真实项目开发场景。
返回列表