jax包介绍及代码示例
·
文章目录
JAX 是由 Google 开发的用于高性能机器学习研究的 Python 库,它将 NumPy 的语法与 自动微分、 向量化 和 GPU/TPU 加速 相结合,特别适合开发深度学习模型和优化算法。
核心特性
- Autograd:自动计算梯度(支持反向模式和前向模式微分)。
- JIT 编译:通过
jax.jit将 Python 函数编译为高效的 XLA 代码,加速计算。 - 并行化:通过
jax.vmap(向量化映射)和jax.pmap(并行映射)简化批量计算和分布式训练。 - 硬件加速:无缝支持 GPU/TPU,无需手动管理设备。
基础用法示例
下面通过几个例子展示 JAX 的核心功能:
1. 自动微分(Autograd)
import jax
import jax.numpy as jnp
# 定义函数
def f(x):
return jnp.sin(x)
# 计算导数(梯度)
dfdx = jax.grad(f)
print(f"f(π/2) = {f(jnp.pi/2)}") # 输出: 1.0
print(f"df/dx(π/2) = {dfdx(jnp.pi/2)}") # 输出: 6.123e-17 (接近 0)
2. JIT 编译加速
# 定义矩阵乘法函数
def matmul(a, b):
return jnp.dot(a, b)
# 创建随机矩阵
key = jax.random.PRNGKey(0)
a = jax.random.normal(key, (1000, 1000))
b = jax.random.normal(key, (1000, 1000))
# 编译函数
matmul_jit = jax.jit(matmul)
# 第一次调用会触发编译
%timeit matmul_jit(a, b) # 通常比普通 jnp.dot 快 10-100 倍
3. 向量化映射(vmap)
# 普通函数:计算单个向量的范数
def norm(x):
return jnp.sqrt(jnp.sum(x**2))
# 使用 vmap 将函数向量化,处理批量数据
batch_norm = jax.vmap(norm)
# 创建批量数据 (100 个向量,每个维度为 10)
x_batch = jax.random.normal(key, (100, 10))
# 一次性计算所有向量的范数
norms = batch_norm(x_batch)
print(f"Batch norms shape: {norms.shape}") # 输出: (100,)
4. 并行训练(pmap)
# 在多个设备(如多 GPU)上并行计算
def update(params, grads):
return params - 0.1 * grads # 简单的梯度下降更新
# 将函数编译为并行版本
update_parallel = jax.pmap(update)
# 在 8 个 TPU 核心上并行更新参数
params = jax.random.normal(key, (8, 100)) # 每个设备一个参数副本
grads = jax.random.normal(key, (8, 100)) # 每个设备一个梯度副本
updated_params = update_parallel(params, grads)
深度学习示例:训练简单神经网络
下面是一个使用 JAX 实现的手写数字识别(MNIST)的完整示例:
import jax
import jax.numpy as jnp
import numpy as np
import optax # JAX 优化器库
from tensorflow.keras.datasets import mnist
# 数据加载
def load_data():
(x_train, y_train), (x_test, y_test) = mnist.load_data()
x_train = x_train.reshape(-1, 784) / 255.0
x_test = x_test.reshape(-1, 784) / 255.0
return (x_train, y_train), (x_test, y_test)
# 初始化网络参数
def init_params(key):
key1, key2 = jax.random.split(key)
params = {
'w1': jax.random.normal(key1, (784, 256)) * 0.01,
'b1': jnp.zeros(256),
'w2': jax.random.normal(key2, (256, 10)) * 0.01,
'b2': jnp.zeros(10)
}
return params
# 前向传播
def forward(params, x):
hidden = jax.nn.relu(jnp.dot(x, params['w1']) + params['b1'])
return jnp.dot(hidden, params['w2']) + params['b2']
# 损失函数
def loss_fn(params, x, y):
logits = forward(params, x)
return optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()
# 训练步骤
@jax.jit
def train_step(params, opt_state, x, y):
grads = jax.grad(loss_fn)(params, x, y)
updates, opt_state = optimizer.update(grads, opt_state)
params = optax.apply_updates(params, updates)
return params, opt_state
# 评估模型
@jax.jit
def evaluate(params, x, y):
logits = forward(params, x)
predictions = jnp.argmax(logits, axis=1)
return jnp.mean(predictions == y)
# 主训练循环
def train():
(x_train, y_train), (x_test, y_test) = load_data()
key = jax.random.PRNGKey(0)
params = init_params(key)
# 定义优化器
optimizer = optax.adam(learning_rate=0.001)
opt_state = optimizer.init(params)
# 训练
for epoch in range(10):
for i in range(0, len(x_train), 128):
x_batch = x_train[i:i+128]
y_batch = y_train[i:i+128]
params, opt_state = train_step(params, opt_state, x_batch, y_batch)
# 评估
train_acc = evaluate(params, x_train[:1000], y_train[:1000])
test_acc = evaluate(params, x_test, y_test)
print(f"Epoch {epoch+1}, Train Acc: {train_acc:.4f}, Test Acc: {test_acc:.4f}")
# 运行训练
train()
注意事项
- 纯函数约束:JAX 要求被编译的函数必须是纯函数(无副作用,输入相同则输出相同)。
- 不可变数据:JAX 数组不可变,修改需通过函数返回新数组。
- 随机数处理:JAX 使用显式随机数生成器(PRNGKey),不同于 NumPy 的全局状态。
- 调试建议:使用
static_argnums参数标记不参与 JIT 编译的参数,避免形状变化导致的编译错误。
相关资源
- 官方文档:JAX Documentation
- 教程:JAX 101
- 生态系统:Flax(JAX 神经网络库)、Haiku(JAX 模块化神经网络)、Optax(优化器库)。
更多推荐

所有评论(0)