JAX 是由 Google 开发的用于高性能机器学习研究的 Python 库,它将 NumPy 的语法与 自动微分向量化GPU/TPU 加速 相结合,特别适合开发深度学习模型和优化算法。

核心特性

  1. Autograd:自动计算梯度(支持反向模式和前向模式微分)。
  2. JIT 编译:通过 jax.jit 将 Python 函数编译为高效的 XLA 代码,加速计算。
  3. 并行化:通过 jax.vmap(向量化映射)和 jax.pmap(并行映射)简化批量计算和分布式训练。
  4. 硬件加速:无缝支持 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()

注意事项

  1. 纯函数约束:JAX 要求被编译的函数必须是纯函数(无副作用,输入相同则输出相同)。
  2. 不可变数据:JAX 数组不可变,修改需通过函数返回新数组。
  3. 随机数处理:JAX 使用显式随机数生成器(PRNGKey),不同于 NumPy 的全局状态。
  4. 调试建议:使用 static_argnums 参数标记不参与 JIT 编译的参数,避免形状变化导致的编译错误。

相关资源

  • 官方文档JAX Documentation
  • 教程JAX 101
  • 生态系统:Flax(JAX 神经网络库)、Haiku(JAX 模块化神经网络)、Optax(优化器库)。
Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐