TensorFlow/Keras 入门指南
TensorFlow 是 Google 开源的工业级深度学习框架,而 Keras 是它内置的高层 API——用几行代码就能搭出并训练一个神经网络。本页带你从 Keras 的极简建模到完整的训练评估流程。前置阅读:NumPy 科学计算入门、反向传播。
这个库是什么
Section titled “这个库是什么”TensorFlow 于 2015 年由 Google Brain 团队开源,最初采用静态计算图(先编译后执行),适合大规模工业部署。Keras 原本是一个独立的高层封装库,由 François Chollet 创建,目标是用最简洁的 API 搭建神经网络。从 TensorFlow 2.0 起,Keras 被正式整合为 TensorFlow 的官方高层 API(tf.keras),TensorFlow 也引入了 eager execution(动态图,即代码按书写顺序立即执行,而非先编译整张图)让调试更直观。
把两者分工比作”自动挡汽车”:TensorFlow 是底盘与发动机(底层张量计算、自动微分、分布式运行时、部署工具链),Keras 是方向盘与油门(Sequential、compile、fit 这套极简接口)。日常建模你几乎只摸方向盘,需要时也能切到手动模式(tf.GradientTape 自定义训练循环)。TensorFlow 在工业部署上长期领先——从服务器(TF Serving)到手机(TFLite)到浏览器(TF.js)到嵌入式(TF Micro),一套模型打通全端。
2025-2026 最新进展:Keras 3 多后端时代
Section titled “2025-2026 最新进展:Keras 3 多后端时代”2025-2026 年是 TensorFlow 和 Keras 生态发生重大变革的时期:
TensorFlow 2.20(2025 年 8 月):
- tf.lite → LiteRT:TFLite(TensorFlow Lite,移动端推理框架)正式迁移到独立的 LiteRT 项目。LiteRT 在 NPU(神经网络处理器)和 GPU 硬件加速方面大幅改进,提供统一的硬件接口,零拷贝缓冲区传递,特别适合大模型和实时推理的端侧部署。
tf.lite模块将从未来的 TensorFlow Python 包中移除,建议尽早迁移到 LiteRT。 - tf.data 预热加速:新增
autotune.min_parallelism选项,减少数据管道首次处理数据的延迟(冷启动优化)。
Keras 3.15(2025-2026 年)——多后端架构成熟:
Keras 3 是一次根本性的架构升级——它不再只绑定 TensorFlow,而是同时支持 TensorFlow、PyTorch、JAX、NumPy、OpenVINO 五个后端。这意味着你可以用同一套 Keras 代码,在 PyTorch 上训练、在 JAX 上做高性能推理、在 OpenVINO 上部署到 Intel 硬件。
关键新特性包括:
- Keras-to-Torch 导出:
model.export(format="torch")将 Keras 模型导出为原生 PyTorchnn.Module,方便与 PyTorch 生态共享。 - Sliding Window Attention(滑动窗口注意力):
MultiHeadAttention和GroupedQueryAttention新增sliding_window参数,高效处理长序列。 - Flash Attention 自动调度:因果注意力(causal-only attention,即只看当前位置及之前 token 的注意力)自动使用 cuDNN 的 Flash SDPA 算子,显著提速。
- MultiOptimizer:支持为模型的不同子网络分配不同的优化器(如编码器用大学习率、解码器用小学习率)。
- ScheduleFreeAdamW:无需学习率调度的 AdamW 变体,省去复杂的 LR schedule 配置。
- AWQ 量化:支持 Activation-aware Weight Quantization(激活感知权重量化),将模型压缩到 INT4 精度。
- CLAHE 层:内置对比度受限自适应直方图均衡化预处理层。
- 全面安全加固:HDF5 文件防路径遍历、档案解压炸弹防护、反序列化安全——Keras 对模型文件安全性做了系统性加固。
安装与环境配置
Section titled “安装与环境配置”# 方式一:pip 安装 TensorFlow(含 Keras 3)pip install tensorflow
# 方式二:CPU 版(体积更小)pip install tensorflow-cpu
# 方式三:单独安装 Keras 3(多后端,不依赖 TensorFlow)pip install keras# 然后通过环境变量选择后端:# export KERAS_BACKEND="torch" # 使用 PyTorch 后端# export KERAS_BACKEND="jax" # 使用 JAX 后端# export KERAS_BACKEND="tensorflow" # 使用 TensorFlow 后端(默认)
# 验证安装python -c "import tensorflow as tf; print(tf.__version__); print('GPU:', tf.config.list_physical_devices('GPU'))"TensorFlow 2.x 已自带 Keras 3(tf.keras),无需单独安装。GPU 版在 Linux 上能自动识别 NVIDIA 显卡并调用 CUDA;Windows 的 GPU 支持需要 WSL2(Windows Subsystem for Linux 2)。导入约定写为 import tensorflow as tf。首次运行会打印一些硬件检测日志,属正常现象。
Keras 3 多后端提示:如果你只关心 Keras API 而不依赖 TensorFlow 特有功能,可以直接
pip install keras并选择 PyTorch 或 JAX 作为后端。但如果你需要 TensorFlow 的部署工具链(LiteRT、TF Serving 等),仍应安装tensorflow。
Tensor:与 PyTorch 类似的张量对象
Section titled “Tensor:与 PyTorch 类似的张量对象”TensorFlow 的 Tensor 同样是多维数组,支持 GPU 运算与自动微分。区别在于它强调”可被编译优化”——配合 tf.function 装饰器,Python 函数会被编译成高性能计算图(先编译再执行,编译器能做全局优化),既有动态图的易用,又有静态图的速度。
tf.GradientTape:磁带式自动微分
Section titled “tf.GradientTape:磁带式自动微分”TensorFlow 的自动微分机制很形象:前向传播时,所有运算像被录在”磁带”上;反向传播时倒带重放,算出梯度。你只需把前向计算包在 with tf.GradientTape() as tape: 块里,再用 tape.gradient(loss, variables) 取梯度。这是自定义训练循环的核心。
Keras 三种建模方式
Section titled “Keras 三种建模方式”Keras 提供三种由简到繁的建模方式,覆盖从入门到研究的需求:
- Sequential(顺序模型):把层像穿串一样按顺序堆叠,适合绝大多数前馈网络(数据从输入层单向流到输出层,没有回环)。最简单,一行定义一层。
- Functional API(函数式 API):把层当成函数,上一层输出作为下一层输入,支持多输入多输出与分支结构。灵活度大增,能表达任意有向无环图(DAG)。
- Subclassing(子类化):继承
tf.keras.Model自定义call方法,与 PyTorch 的nn.Module几乎一致。最灵活,适合研究型复杂模型。
统一的 compile + fit 工作流
Section titled “统一的 compile + fit 工作流”Keras 把训练抽象成三件事:compile(指定优化器、损失、指标)、fit(喂训练数据开始训练)、evaluate(在测试数据上评估)。这套”编译—训练—评估”的极简流程是 Keras 最大的魅力——一个完整模型从定义到训练评估常常不超过 15 行代码。
架构与工作流
Section titled “架构与工作流”例 1:Sequential 模型一步到位
Section titled “例 1:Sequential 模型一步到位”import tensorflow as tffrom tensorflow import kerasfrom tensorflow.keras import layers
# 用 Sequential 堆叠层:一个两层 MLP(多层感知机)# Dense = 全连接层:每个输入神经元与每个输出神经元都有连接# ReLU 激活函数:max(0, x),简单高效,是深度学习的默认选择# Softmax 激活函数:将输出归一化为概率分布(各类别概率之和为 1)model = keras.Sequential([ layers.Dense(64, activation="relu", input_shape=(20,)), # 隐藏层:64 个神经元 layers.Dropout(0.3), # 随机丢弃 30% 神经元,防过拟合 layers.Dense(3, activation="softmax"), # 输出层:3 个类别])model.summary() # 打印模型结构与参数量
# 编译:指定优化器、损失函数、评估指标# Adam:自适应矩估计优化器,自动调整每个参数的学习率# sparse_categorical_crossentropy:标签为整数时的交叉熵损失model.compile(optimizer="adam", loss="sparse_categorical_crossentropy", metrics=["accuracy"])例 2:编译、训练与评估
Section titled “例 2:编译、训练与评估”import numpy as np
# 生成随机分类数据:1000 个样本,20 维特征,3 类X = np.random.randn(1000, 20).astype(np.float32)y = np.random.randint(0, 3, (1000,))
# fit 一行启动训练:自动分批、迭代、记录历史# epochs:遍历整个训练集的次数# batch_size:每次梯度更新使用的样本数# validation_split:留出 20% 训练数据做验证,监控泛化能力history = model.fit(X, y, epochs=20, batch_size=32, validation_split=0.2, verbose=0)print(f"最终训练准确率: {history.history['accuracy'][-1]:.3f}")print(f"最终验证准确率: {history.history['val_accuracy'][-1]:.3f}")
# evaluate 在测试集上评估test_loss, test_acc = model.evaluate(X, y, verbose=0)print(f"测试准确率: {test_acc:.3f}")
# predict 做预测preds = model.predict(X[:3], verbose=0)print(f"前 3 个样本预测类别: {preds.argmax(axis=1)}")例 3:用 Functional API 构建带分支的模型
Section titled “例 3:用 Functional API 构建带分支的模型”from tensorflow.keras import layers, Model
# 函数式 API:每层是一个函数,前一层输出作为后一层输入# 适合需要多输入、多输出或跳跃连接(skip connection)的场景inputs = keras.Input(shape=(20,), name="features")x = layers.Dense(64, activation="relu")(inputs) # 第一层x = layers.BatchNormalization()(x) # 批归一化:稳定训练x = layers.Dropout(0.3)(x) # 随机丢弃防过拟合x = layers.Dense(32, activation="relu")(x) # 第二层outputs = layers.Dense(3, activation="softmax")(x) # 输出层
model = keras.Model(inputs=inputs, outputs=outputs) # 组装成模型model.compile(optimizer="adam", loss="sparse_categorical_crossentropy")model.summary() # 可看到多分支的完整结构例 4:用 GradientTape 自定义训练循环
Section titled “例 4:用 GradientTape 自定义训练循环”import tensorflow as tf
# 准备数据与模型X = tf.random.normal((500, 20))y = tf.random.uniform((500,), 0, 3, dtype=tf.int32)model = keras.Sequential([ layers.Dense(64, activation="relu"), layers.BatchNormalization(), layers.Dense(3),])loss_fn = keras.losses.SparseCategoricalCrossentropy(from_logits=True)# from_logits=True:模型输出未经 softmax,损失函数内部会自动处理optimizer = keras.optimizers.Adam(1e-2)
# 自定义训练循环:磁带录制前向,倒带算梯度# 这是最底层的训练方式,给你完全的控制权# 适合 GAN、自定义 RL 等非标准训练流程for epoch in range(20): with tf.GradientTape() as tape: # 开始录制 logits = model(X) # 前向传播 loss = loss_fn(y, logits) # 计算损失 grads = tape.gradient(loss, model.trainable_variables) # 反向求梯度 optimizer.apply_gradients(zip(grads, model.trainable_variables)) # 更新参数 if epoch % 5 == 0: acc = (tf.argmax(logits, 1) == y).numpy().mean() print(f"epoch {epoch}, loss={loss:.4f}, acc={acc:.3f}")例 5:使用 Callbacks 管理训练过程
Section titled “例 5:使用 Callbacks 管理训练过程”回调函数(Callback)是 Keras 在训练过程中的”自动助手”——在特定时机(如每个 epoch 结束)自动执行预设操作:
from tensorflow.keras import callbacks
# 定义回调cb_list = [ # 早停:验证损失连续 5 轮不改善就停止训练,避免过拟合 callbacks.EarlyStopping(patience=5, restore_best_weights=True),
# 自动保存最佳模型(按验证准确率) callbacks.ModelCheckpoint( "best_model.keras", # v3.x 推荐 .keras 格式(单文件,安全) monitor="val_accuracy", save_best_only=True, ),
# 验证损失停滞时自动降低学习率(乘以 0.5) callbacks.ReduceLROnPlateau(factor=0.5, patience=3),
# 在 TensorBoard 中记录训练曲线(可视化 loss/accuracy) # callbacks.TensorBoard(log_dir="./logs"),]
# 在 fit 中传入回调列表history = model.fit( X, y, epochs=50, batch_size=32, validation_split=0.2, callbacks=cb_list, # 自动管理训练 verbose=1,)print(f"训练在 epoch {len(history.epoch)} 处早停")例 6:用 tf.data 构建高效数据管道
Section titled “例 6:用 tf.data 构建高效数据管道”面对大规模数据,直接用 NumPy 数组传给 fit 会遇到内存瓶颈。tf.data.Dataset 提供了高效的数据管道——支持预取(prefetch,在 GPU 训练的同时 CPU 准备下一批数据)、并行映射、缓存等功能:
import tensorflow as tf
# 从 NumPy 数组创建 Datasetdataset = tf.data.Dataset.from_tensor_slices((X, y))
# 构建高效数据管道dataset = ( dataset .shuffle(buffer_size=1000) # 打乱顺序,避免模型学到数据的排列规律 .batch(32) # 分批 .prefetch(tf.data.AUTOTUNE) # 预取:CPU 和 GPU 并行工作)
# v2.20 新增:减少冷启动延迟# autotune.min_parallelism 让数据管道从一开始就并行处理# tf.data.Options().autotune.min_parallelism = 8
# 直接用 Dataset 训练model.fit(dataset, epochs=20, verbose=1)常用 API 速查
Section titled “常用 API 速查”| API | 用途 | 示例 |
|---|---|---|
keras.Sequential(layers) | 顺序堆叠模型 | Sequential([Dense(10), Dense(3)]) |
layers.Dense(n) | 全连接层 | Dense(64, activation="relu") |
layers.Conv2D(...) | 二维卷积层 | Conv2D(32, 3, activation="relu") |
layers.Dropout(rate) | 随机丢弃层 | Dropout(0.5) |
layers.BatchNormalization() | 批归一化层 | BatchNormalization() |
layers.MultiHeadAttention | 多头注意力层(含 sliding_window) | MultiHeadAttention(num_heads=8, key_dim=64) |
keras.Input(shape) | 定义输入 | Input(shape=(28,28)) |
model.compile(...) | 编译模型 | compile("adam", "mse") |
model.fit(x, y) | 训练模型 | fit(X, y, epochs=10, callbacks=[...]) |
model.evaluate(x, y) | 评估模型 | evaluate(X_test, y_test) |
model.predict(x) | 推理预测 | pred = model.predict(X) |
model.save(path) | 保存模型 | model.save("m.keras") |
model.export(format=) | 导出格式(v3.15) | model.export(format="torch") |
tf.data.Dataset | 高效数据管道 | Dataset.from_tensor_slices((X,y)) |
tf.GradientTape() | 自动微分磁带 | with GradientTape() as t: |
callbacks.EarlyStopping() | 早停回调 | EarlyStopping(patience=5) |
callbacks.ModelCheckpoint() | 模型检查点 | ModelCheckpoint("best.keras") |
- 优先用 Keras 高层 API:90% 的任务用
Sequential或 Functional API 配合fit就够了,不必手写训练循环。只有当损失或数据流特殊(如 GAN——生成对抗网络,两个模型互相博弈;或自定义强化学习)才需要GradientTape。 - 用
tf.data构建数据管道:面对大规模数据,tf.data.Dataset提供预取、缓存、并行映射,能避免 I/O 成为训练瓶颈。它是 Kerasfit的标准数据输入方式。v2.20 的autotune.min_parallelism能进一步减少冷启动延迟。 - 善用 Callbacks:
EarlyStopping(验证集不再提升时自动停止)、ModelCheckpoint(自动保存最佳权重)、ReduceLROnPlateau(自动降低学习率)——这三板斧能让训练省心且效果更好。 - 保存格式选择:新项目推荐
model.save("m.keras")(Keras 3 原生格式,单文件,安全性经过加固);需要部署到移动端则迁移到 LiteRT;需要导出给 PyTorch 生态可用model.export(format="torch")(v3.15 新增)。 - Keras 3 多后端的价值:如果你已经在用 PyTorch 或 JAX,但想用 Keras 简洁的
compile + fitAPI,Keras 3 让你鱼和熊掌兼得。研究阶段用 JAX 后端获得极致速度,部署阶段用 TF 后端获得端到端工具链。 - LiteRT 替代 TFLite:如果你在做移动端/嵌入式 AI,注意 tf.lite 正在迁移到独立的 LiteRT 项目。新项目直接用 LiteRT,能获得更好的 NPU 和 GPU 加速支持。
- 混合精度训练:
keras.mixed_precision.set_global_policy("mixed_float16")一行开启,提速省显存。详见混合精度训练。 - 安全意识:Keras 3.12-3.15 对模型文件加载做了全面安全加固(防 HDF5 路径遍历、防解压炸弹等)。加载来源不明的模型文件时,保持
safe_mode=True(默认值)。
典型应用场景
Section titled “典型应用场景”- 工业级生产部署:从训练到手机、浏览器、边缘设备的端到端部署,TensorFlow + LiteRT 工具链最成熟。
- 图像分类与视觉模型:Keras 内置常用 CNN 层与预训练模型(ResNet、EfficientNet),快速搭建视觉应用。详见卷积神经网络。
- 自然语言处理:Keras 3.15 的
MultiHeadAttention支持 Sliding Window Attention 和 Flash Attention,能高效构建 Transformer 类模型。复杂 Transformer 架构详解参见Transformer 架构。 - 推荐系统与结构化数据:Wide & Deep 等推荐模型在 TensorFlow 生态中有成熟实现。详见线性回归。
- 移动端与嵌入式 AI:LiteRT 模型转换与量化让深度学习跑在手机、树莓派甚至微控制器上。v2.20 的 LiteRT 在 NPU 加速方面有重大改进。
与同类工具对比
Section titled “与同类工具对比”| 特性 | TensorFlow 2.20 / Keras 3 | PyTorch | JAX / Flax |
|---|---|---|---|
| 高层 API | Keras 3(极简,多后端) | 无官方高层(手写循环) | Flax(较新) |
| 多后端支持 | TF / PyTorch / JAX / NumPy | 仅 PyTorch | 仅 JAX |
| 计算图 | 动态为主,tf.function 可编译 | 动态 | 静态(JIT 编译) |
| 建模速度 | 极快(Sequential + fit) | 较快 | 中等 |
| 自定义灵活性 | 中等(GradientTape) | 极高 | 极高 |
| 工业部署 | 最强(LiteRT / Serving / ONNX) | 较好(ONNX/TorchServe) | 较弱 |
| Flash Attention | v3.15 自动调度 | 手动集成 | 原生支持 |
| 安全加固 | 全面(HDF5/反序列化防护) | 基础 | 基础 |
| 学术占有率 | 下降中 | 第一 | 上升中 |
| 最佳场景 | 生产部署 / 移动端 / 快速原型 | 研究 / 大模型 | 高性能研究 |
关键洞察:需要快速搭原型或全端部署选 TensorFlow/Keras;需要极致灵活性或做前沿研究选 PyTorch。Keras 3 的多后端架构让两者的边界变得模糊——你可以在 PyTorch 后端上享受 Keras 的简洁 API。两者理念对比详见PyTorch 入门指南。深度学习训练背后的优化器与梯度下降原理是相通的,详见梯度下降与优化器。
- TensorFlow 官方教程:tensorflow.org/tutorials — 从 Keras 入门到分布式训练、模型部署的完整教程,配有大量可运行示例。
- Keras 官方网站(v3):keras.io — Keras 3 的官方文档与指南,覆盖多后端使用、Functional API、子类化模型、自定义训练循环等主题。Keras 3.0 起的所有新闻和发布都在这里。
- François Chollet,「Deep Learning with Python」(第 2 版):Keras 作者编写的经典著作,用 Keras 从零讲透深度学习的核心概念,文笔极佳,被誉为入门最佳书。
- 「TensorFlow 论文」(Abadi et al., 2016):TensorFlow 的系统设计论文,阐述其数据流图与分布式执行模型,是理解框架底层设计的关键文献。