JAX 是 Google 主导开发的一个开源 Python 库,专为高性能数值计算和机器学习研究而设计。它常常被看作是 NumPy 的“加速器”,可以在 CPU、GPU 和 TPU 上运行相同的代码。
💡 1. 功能说明与使用说明
JAX 的核心在于它不仅仅是一个数值计算库,更是一个可组合的程序变换系统。这意味着它能将你编写的 Python 函数,通过编译和变换,转化为高性能、可并行的代码。
核心功能
-
类 NumPy 的 API (
jax.numpy):JAX 提供了一个与 NumPy 高度一致的接口 (jnp),让熟悉 NumPy 的开发者可以几乎无缝上手。大部分numpy代码只需将import numpy as np改为import jax.numpy as jnp即可运行。 -
即时编译 (JIT, Just-In-Time Compilation):通过
@jax.jit装饰器,JAX 可以将你的 Python 函数在运行时编译成针对特定硬件(如 GPU/TPU)优化的机器码,大幅提升计算速度。 -
自动微分 (Automatic Differentiation):JAX 提供了
jax.grad()等函数,可以轻松计算任何可微 Python 函数的梯度。它支持前向和反向模式的自动微分,是训练神经网络等机器学习任务的基础。 -
自动向量化 (Vectorization):通过
@jax.vmap装饰器,可以将一个处理单个样本的函数,自动转换为能高效处理批量数据的函数。 -
函数式编程范式:JAX 的核心设计哲学是函数式编程。它要求函数是纯函数(无副作用,相同输入总是产生相同输出),并且其数组 (
jax.Array) 是不可变的(Immutable)。任何修改操作都会返回一个新数组,而不是修改原数组。
使用说明
安装 JAX 非常简单,通过 Python 的包管理工具 pip 即可完成。
-
CPU 版本:
pip install jax -
NVIDIA GPU (CUDA) 版本:
pip install -U "jax[cuda13]"
具体安装方式会根据你的硬件和操作系统有所不同,建议查阅官方安装指南。
安装完成后,就可以在 Python 中使用了:
import jax.numpy as jnp import jax # 创建一个 JAX 数组 x = jnp.array([1.0, 2.0, 3.0]) # 使用 jit 加速一个函数 @jax.jit def f(x): return x ** 2 + 2 * x + 1 print(f(x)) # 计算会经过编译优化 # 计算梯度 grad_f = jax.grad(lambda x: jnp.sum(x ** 2)) print(grad_f(x)) # 输出梯度
