You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

为什么我的JAX + Haiku代码无法在GPU上高效运行?

问题原因与解决方案

你遇到的GPU利用率低、训练速度慢的问题是JAX + Haiku训练时非常典型的设备数据不对齐、编译优化不全导致的,具体原因和修复方法如下:

核心问题原因

  • 数据加载默认走CPU,每次迭代都要做CPU到GPU的张量拷贝,GPU大部分时间处于等待数据的空闲状态,因此利用率极低,显存占用高是因为JAX初始化时会预先申请大量显存缓存,属于正常现象。
  • loss函数没有加@jax.jit装饰器,打印损失时执行的是未优化的CPU侧计算,还会触发GPU-CPU同步阻塞。
  • 你将mlp作为参数传入loss函数,JAX的JIT编译会把它当成动态参数处理,每次迭代都会触发重复编译,额外占用大量时间。
  • 初始化模型参数时传入的输入张量是CPU侧的,导致模型参数默认创建在CPU上,每次训练迭代都要来回拷贝参数和数据。
  • MNIST模型本身计算量极小,如果batch size设置过小,GPU还没启动计算就已经完成了单步迭代,无法跑满算力。

具体修复步骤

1. 调整数据加载逻辑,提前把数据放到GPU

MNIST数据集体积很小,可以直接把全量训练数据预加载到GPU,完全避免迭代过程中的跨设备拷贝:

# 假设你用的是torchvision的MNIST数据集,提前转成JAX数组并放到GPU
import numpy as np
from torchvision.datasets import MNIST

train_dataset = MNIST(root="./data", train=True, download=True)
train_images = jax.device_put(jnp.array(train_dataset.data.numpy(), dtype=jnp.float32).reshape(-1, 784) / 255.0)
train_labels = jax.device_put(jnp.array(train_dataset.targets.numpy(), dtype=jnp.int32))

# 写纯JAX实现的batch生成器,完全跳过PyTorch DataLoader的CPU流转
def get_batches(rng, images, labels, batch_size=256):
    perm = jax.random.permutation(rng, len(images))
    for i in range(0, len(images), batch_size):
        batch_ids = perm[i:i+batch_size]
        yield images[batch_ids], labels[batch_ids]

2. 优化loss和update函数的JIT编译

移除loss中多余的model参数,给loss加上JIT装饰器,避免重复编译:

@jax.jit
def loss(params, images, labels):
  logits = mlp.apply(params = params, images = images)
  labels = jax.nn.one_hot(labels, num_classes = 10)
  cross_entropy_loss = -jnp.sum(labels*logits)/len(labels)
  return cross_entropy_loss

@jax.jit
def update(params, opt_state, images, labels):
  grads = jax.grad(loss)(params, images, labels)
  updates, opt_state = opt.update(grads, opt_state)
  return optax.apply_updates(params, updates), opt_state

3. 确保模型参数初始化在GPU上

初始化模型参数时,传入GPU侧的张量作为输入:

# 取GPU上的第一个batch用来初始化
dummy_batch = train_images[:256]
params = mlp.init(rng = jax.random.PRNGKey(0), images = dummy_batch)
opt_state = opt.init(params = params)

4. 调整训练循环

用纯JAX的batch生成器替换原来的PyTorch DataLoader:

def train(params, opt_state, epochs, batch_size=256):
  rng = jax.random.PRNGKey(0)
  for epoch in range(epochs):
    rng, perm_rng = jax.random.split(rng)
    for batch_idx, (images, labels) in enumerate(get_batches(perm_rng, train_images, train_labels, batch_size)):
      if batch_idx == 0:
        print(f"Epoch {epoch} : loss = {loss(params,images,labels)}")
      params, opt_state = update(params, opt_state, images,labels)

%time train(params, opt_state, epochs = 10)

5. 可选优化

如果你的batch size小于256,可以适当调大到256/512,进一步提升GPU利用率。

内容的提问来源于stack exchange,提问作者Adrien Lagesse

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.26 18:27:02