keras
Keras 3 — 支持 JAX/TensorFlow/PyTorch 多后端切换的深度学习框架,用
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
Keras 3 — 支持 JAX/TensorFlow/PyTorch 多后端切换的深度学习框架,用
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象一下这样的场景:一位生物医学研究员想要用深度学习分析病理切片图像,却不想花三个月学习 GPU 集群配置和底层张量运算;一位独立游戏开发者希望为自己的游戏添加 AI 角色对话能力,却不知道该选 PyTorch 还是 JAX。在这种"人人都有 AI 需求"的时代,Keras 3 正是为解决这一痛点而生——它用最少的代码、最低的门槛,让任何人都能快速构建和训练深度学习模型。
Keras 的诞生可以追溯到 2015 年,由 Google 工程师 Francois Chollet(同时也是著名的 XCeptor 图像分类模型和《Python 深度学习》一书作者)创建。初代 Keras 以 Theano 为后端,后来相继支持 TensorFlow、CNTK、MXNet 等。2019 年 Keras 被正式纳入 TensorFlow 核心,成为 TF.keras 模块,吸引了数百万开发者。
然而,随着 PyTorch 在研究领域的崛起和 JAX 在 Google 内部的广泛使用,单一后端依赖的 Keras 逐渐暴露局限性。于是 2023 年发布的 Keras 3 带来了革命性改变——彻底摆脱单一后端,同时原生支持 JAX、TensorFlow、PyTorch 三大主流框架,以及仅用于推理的 OpenVINO。用户可以在不同后端之间自由切换,甚至在同一个训练循环中混用不同框架的操作符,享受各后端的最优性能。
Keras 3 的设计哲学高度统一,可以用一个词概括:简洁。它提供了从输入到输出的完整高层 API 抽象,开发者无需关心底层张量运算的细节,只需几行代码就能定义模型、配置训练流程。
这是 Keras 3 最大的技术亮点。通过设置环境变量 KERAS_BACKEND,用户可以在 jax、tensorflow、torch 之间一键切换,后端会在运行时自动选择对应的实现。以一个图像分类任务为例,模型定义保持完全不变,但底层可以用 PyTorch 的 eager 模式调试(方便排错),也可以切换到 JAX 后端获得 20%~350% 的训练速度提升。
keras.applications 模块提供了大量开箱即用的预训练模型(ResNet、VGG、EfficientNet、MobileNet 等),keras.layers 涵盖了从全连接层、卷积层、注意力机制到循环网络的全套构建块。开发者可以直接复用这些组件快速搭建自定义模型。
Keras 3 原生支持多 GPU 和 TPU 分布式训练。keras.distribution 模块提供了简洁的 API 来配置数据并行和模型并行策略,配合官方提供的 requirements-jax-cuda.txt、requirements-torch-cuda.txt 等依赖文件,用户可以快速在单机多卡或大规模集群上扩展训练规模。
Keras 3 的安装非常简单,核心只需一条命令:
pip install keras
pip install torch # 或 jax 或 tensorflow(至少安装一个后端)
安装完成后,验证安装:
import os
os.environ["KERAS_BACKEND"] = "torch" # 可选:jax / tensorflow / torch
import keras
print(keras.__version__) # 3.x
构建一个简单的 MNIST 分类模型:
import keras
from keras import layers
model = keras.Sequential([
layers.Dense(512, activation="relu"),
layers.Dropout(0.5),
layers.Dense(10, activation="softmax")
])
model.compile(optimizer="adam", loss="sparse_categorical_crossentropy", metrics=["accuracy"])
(x_train, y_train), (x_test, y_test) = keras.datasets.mnist.load_data()
x_train = x_train.reshape(60000, 784).astype("float32") / 255
model.fit(x_train, y_train, epochs=5, batch_size=32)
这段代码几乎与伪代码无异,却完成了从数据加载、模型构建到训练的完整流程。Keras 3 的 Functional API 还支持构建多输入多输出、共享层等复杂拓扑结构,满足从科研原型到生产级模型的一切需求。
Keras 3 本身是纯 Python 库,基础安装仅需 2GB 磁盘空间和 8GB 内存。但作为深度学习框架,GPU 是实际训练模型的必需品。官方建议 NVIDIA GPU(RTX 3080 及以上),显存 8GB 以上。需要提前安装 CUDA 11.8+/12.x 和 cuDNN 8.x。
值得注意的是,Keras 3 官方仅提供源代码安装和 pip 安装两种方式,不支持 Docker 一键部署,也没有 Web 界面。对于需要快速原型验证的场景,推荐使用 Google Colab 或 Kaggle Notebook(两者均已预装 TensorFlow/PyTorch 后端);对于生产级部署,建议将训练好的模型导出为 ONNX 或 TensorFlow SavedModel 格式,再通过 TensorFlow Serving 或 TorchScript 进行服务化。
尽管 Keras 3 功能强大,但它并非银弹。首先,多后端支持是有代价的:为了兼容不同底层框架,Keras API 无法直接调用任何单一后端的全部高级特性(如 PyTorch 的某些自定义 CUDA 扩展)。其次,对于追求极致性能的场景,Keras 3 的抽象层会带来一定的运行时开销,在超大规模训练中可能不如直接使用 JAX 或 PyTorch 原生代码高效。最后,官方文档目前仍在完善中,部分高级用法(如自定义分布式策略)的示例资料相对稀缺。
Keras 3 的出现标志着深度学习框架从"百花齐放、互不兼容"走向"统一抽象、灵活切换"的新阶段。凭借 64000+ 的 GitHub Stars 和近 300 万开发者社区,Keras 已成为连接学术研究和工业应用的桥梁。随着 JAX 后端带来的性能提升和 OpenVINO 对边缘推理的支持,Keras 3 有望在云端训练、端侧部署、科研原型等多个场景中继续保持其独特的生态位。对于希望快速入门深度学习、又不想被底层细节束缚的开发者而言,Keras 3 仍然是目前最友好的选择之一。