Keras 3 — 多后端深度学习
Keras 3 支持 JAX、TensorFlow、PyTorch、OpenVINO——一套代码,任意后端。20-350% 加速、keras.ops、KerasCV/Hub。GitHub 6.4 万星。
Keras 3 是 keras-team 推出的多后端深度学习框架——“为人类而生的深度学习”。与 Keras 2 的关键区别:不再依赖 TensorFlow。一套代码现在可以跑在 JAX、TensorFlow、PyTorch 和 OpenVINO(仅推理)上。GitHub 约 6.4 万星,超过 500 万开发者使用。
为什么多后端重要#
过去你必须选定一个框架并一直用下去。Keras 3 把 API 与底层计算解耦:
- 无锁定——改一个环境变量即可切换后端
- 最佳性能——为每个模型挑选最快的后端。Keras 基准显示 20-350% 的加速,GPU/TPU/CPU 上通常 JAX 胜出
- 生态可选性——Keras 模型即 PyTorch
Module,可导出为 TensorFlowSavedModel,也可实例化为 JAX 无状态函数
选择后端#
export KERAS_BACKEND="jax" # 或:tensorflow、torch、openvinobash或编辑 ~/.keras/keras.json。各后端最低版本(Keras 3 稳定版):TensorFlow 2.16.1、JAX 0.4.20、PyTorch 2.1.0、OpenVINO 2025.3。
| 后端 | 最擅长 |
|---|---|
| JAX | GPU 与 TPU 上最快的训练/推理、XLA |
| TensorFlow | 现有 TF 技术栈、tf.data 管道 |
| PyTorch | HF 生态、动态模型、科研 |
| OpenVINO | CPU 推理优化(Intel) |
keras.ops —— 一套 NumPy 风格 API#
Keras 3 内置 keras.ops,一个在全部后端都能用的完整 NumPy API——自定义层、损失、指标无需 import tf. / torch. 即可随处运行:
import keras
import numpy as np
x = keras.ops.ones((3, 3))
y = keras.ops.matmul(x, keras.ops.transpose(x))python跨框架数据管道#
Keras 3 模型可用任意管道训练——tf.data.Dataset、torch.utils.data.DataLoader、NumPy 数组、Pandas DataFrame 或 keras.utils.PyDataset。无需重写。
完整高层 API#
层、指标、损失、优化器、回调、训练循环、保存/序列化——一应俱全,且与后端无关。自定义 train_step() 需按后端编写,但 compute_loss() 处处可用。
示例:同一模型,任意后端#
import os
os.environ["KERAS_BACKEND"] = "jax" # 切换:"tensorflow"、"torch"
import keras
model = keras.Sequential([
keras.Input(shape=(28, 28)),
keras.layers.Flatten(),
keras.layers.Dense(128, activation="relu"),
keras.layers.Dense(10, activation="softmax"),
])
model.compile(optimizer="adam", loss="sparse_categorical_crossentropy")
model.fit(x_train, y_train, epochs=5)python同样的代码可跑在 JAX、TF 或 PyTorch 上。
预训练模型#
- Keras Applications: 全部 40 个模型(ResNet、EfficientNet、MobileNet…)支持所有后端
- KerasCV & KerasHub: BERT、OPT、Whisper、T5、StableDiffusion、YOLOv8、SegmentAnything——所有后端
从 Keras 2 / tf.keras 迁移#
仅用内置层的模型几乎零改动迁移。自定义 tf.* 代码需替换为 keras.ops,后端特定的 train_step() 覆写则按后端分别实现。迁移后,一个环境变量即可切到 JAX 或 PyTorch。
优点与缺点#
优点:
- 一次编写,四大框架运行
- 无框架锁定——2026 年实属罕见
- 通过 JAX 后端获得顶尖性能
- 出自打造 tf.keras 的同一团队(François Chollet)
缺点:
- 自定义训练循环需要按后端写代码
- OpenVINO 仅支持推理
- 新的分布式 API(device mesh)目前仅限 JAX
- 小众自定义算子可能需要后端专属回退
结论#
Keras 3 是 2026 年最接近统一深度学习 API 的存在。它不在 TensorFlow、PyTorch、JAX 之间选赢家——而是让你一次写模型、到处运行,并取各框架最佳性能。对于做库或面向多方用户的团队,它正日益成为默认推荐。