库 / SDK
keras-team/keras avatar
keras-team/keras

Keras 3:多后端深度学习框架的取舍与适配指南

多后端深度学习框架,同一套高层 API 可运行在 JAX、TensorFlow 或 PyTorch 之上,并能从笔记本扩展到 GPU 与 TPU 集群。

64,322 个 Star19,790 个 ForkPythonApache-2.0

秒懂

它是什么?
Keras 3 将高层 API 从 TensorFlow 解耦,支持 JAX、PyTorch、TensorFlow 和 OpenVINO 四种后端。本文基于官方文档与仓库信息,分析其机制、安装方式、兼容性边界,并给出适用人群与验证要点。
适合谁用?
Keras 3 适合那些希望用同一套高层 API 在不同后端间切换的团队,尤其是已有 TensorFlow 代码但想尝试 JAX 性能,或想在 PyTorch 生态中使用成熟训练 API 的开发者。不适合需要 OpenVINO 训练能力、依赖旧版 tf.keras 自定义组件且不愿迁移 .keras 格式,或要求后端在运行时动态切换的场景。
能商用吗?
可以。Apache-2.0 是宽松许可证:你可以使用、修改并销售基于它的软件,只需保留版权和许可证声明。
还在维护吗?
在维护。仓库最近一次提交在 1 天前。
用什么语言写的?
主要是 Python(依据 GitHub 的语言统计)。

以上回答依据项目的 GitHub 数据(最近同步于 2026年9月14日)和我们的分析,不构成法律意见。

开源项目深度解析

它解决的是框架锁定问题

Keras 3 要解决的问题很具体:深度学习代码被单一框架绑死。以前的 Keras 是 TensorFlow 的高层封装,模型写完后很难迁移到 JAX 或 PyTorch。Keras 3 把这一层抽出来,让你用同一套 API 在 JAX、TensorFlow、PyTorch 和 OpenVINO 上运行。目标用户是那些不想重写模型代码,又想在不同后端间灵活切换的工程师。它不解决算法设计问题,也不提供新的模型架构。它解决的是工程层面的可移植性。

后端切换的机制与硬性约束

后端选择通过环境变量 KERAS_BACKEND 或配置文件 ~/.keras/keras.json 完成。可选值有 tensorflow、jax、torch、openvino。关键约束在文档中写得很明确:必须在导入 keras 之前设置后端,导入后无法更改。这意味着后端是进程级的静态选择,不是运行时动态切换。另外,OpenVINO 后端仅支持推理,只能用 model.predict(),不能训练。这个限制对生产环境很重要,如果你打算用 OpenVINO 做训练,这条路走不通。文档还提到,在 Colab 中可以通过 os.environ 设置后再 import keras,但顺序不能颠倒。

安装步骤与 GPU 环境的坑

安装分两步:先 pip install keras --upgrade,再安装后端包 tensorflow、jax 或 torch 之一。官方明确推荐 Windows 用户使用 WSL2,因为 Keras 3 只兼容 Linux 和 macOS。本地开发安装需要先 pip install -r requirements.txt,然后从根目录运行 python pip_build.py --install。GPU 支持有单独的 requirements-{backend}-cuda.txt 文件,以 JAX 为例:conda create -n keras-jax python=3.10,激活后 pip install -r requirements-jax-cuda.txt,再执行 python pip_build.py --install。这里有个实际陷阱:requirements.txt 默认装 CPU 版本,GPU 依赖必须用对应的 cuda 文件。官方建议每个后端用独立 Python 环境,避免 CUDA 版本冲突。这个建议很实在,因为同时装 TensorFlow 和 PyTorch 的 CUDA 依赖很容易互相覆盖。

兼容性承诺与边界

官方声称 Keras 3 是 tf.keras 的 drop-in replacement,但有一个前提:模型保存必须使用最新的 .keras 格式。如果你的现有代码还在用旧的 HDF5 格式,需要先更新保存逻辑。另一个边界是自定义组件。文档说,如果模型包含自定义层或自定义 train_step(),通常可以在几分钟内转换为后端无关的实现。这个「通常」很关键,说明转换不是自动的,需要手动改代码。如果模型没有自定义组件,可以直接在 JAX 或 PyTorch 后端上运行。但要注意,数据加载是跨后端兼容的,你可以用 tf.data.Dataset 或 PyTorch DataLoader 训练同一个 Keras 模型,这减轻了迁移负担。

性能声明的来源与可信度

README 提到,选择最快的后端(常是 JAX)可以获得 20% 到 350% 的速度提升,并附有 keras.io 上的 benchmark 链接。但这里没有给出具体的测试环境、模型规模或硬件配置。作为技术编辑,我必须指出:这个数字是官方基准测试的结果,不是第三方独立验证。实际收益取决于你的模型架构、批大小和硬件。JAX 在编译优化上确实有优势,但并非所有模型都能获得 350% 的提升。如果你的模型是小型 CNN,可能差异很小。建议把官方 benchmark 当作参考,而不是承诺。在自己的数据集上跑一遍对比,比相信任何宣传数字都可靠。

与直接使用原生框架的差异

一个自然的替代方案是直接用 JAX 或 PyTorch 写模型,不经过 Keras。区别在于抽象层次。Keras 3 提供高层 API,比如内置的层、优化器、训练循环,而原生 JAX 通常需要你手动定义前向函数、损失函数和梯度更新步骤。PyTorch 有 nn.Module 和 Trainer 生态,但 Keras 3 的 API 设计更统一,特别是在跨后端时。另一种替代是继续用 tf.keras(即 tf-keras 包)。Keras 2 仍然以 tf-keras 形式存在,适合那些不需要多后端、深度绑定 TensorFlow 生态的项目。Keras 3 的价值在于你可以把模型嵌入到原生 JAX 函数或 PyTorch Module 中,这给了你底层控制权,同时保留高层 API 的便捷。

维护成本与许可要点

Keras 3 的维护成本体现在三方面。第一,后端版本有最低要求:TensorFlow 2.16.1、JAX 0.4.20、PyTorch 2.1.0、OpenVINO 2026.2.0。这意味你的环境必须满足这些版本,升级后端时可能连带升级 Keras。第二,API 生成脚本 shell/api_gen.sh 在修改公共 API 时需要运行,这对贡献者是个额外步骤,但普通用户不用管。第三,多后端意味着每个后端的行为可能有细微差异,调试时需要跨框架排查。许可方面,项目采用 Apache-2.0,这是宽松许可,允许商用和修改,但要注意不提供任何担保。它不是法律建议,但比 GPL 类许可更灵活。

编辑结论

Keras 3 适合那些希望用同一套高层 API 在不同后端间切换的团队,尤其是已有 TensorFlow 代码但想尝试 JAX 性能,或想在 PyTorch 生态中使用成熟训练 API 的开发者。不适合需要 OpenVINO 训练能力、依赖旧版 tf.keras 自定义组件且不愿迁移 .keras 格式,或要求后端在运行时动态切换的场景。采纳前应验证三点:你的自定义层或损失函数是否能在目标后端下运行,模型保存是否已迁移到 .keras 格式,以及目标后端的 CUDA 依赖是否与现有环境冲突。Keras 3 的价值在于后端可移植性,但它的上限也由最弱的后端决定,OpenVINO 仅支持推理这一点,足以让生产部署计划重新评估。

官方来源

  1. Official documentation
  2. Official README
  3. Project repository
  4. Release notes
社区笔记

社区笔记