为什么我的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
相关产品推荐
相关产品推荐

