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

如何展开训练循环以在GPU/TPU上实现Jax多步训练

Jax/Flax 5步多步训练实现方案

核心思路

多步训练的关键是一次性加载连续5个批次的数据,并用Jax的jax.lax.scan(静态循环,适配JIT编译)替代Python循环,避免编译开销,充分利用GPU/TPU算力。


1. 多批次数据构造

无需改造复杂加载器,只需从数据迭代器中一次性取出5个批次,用jax.tree_map堆叠成带step维度的结构:

  • 单批次结构:(x: (batch_size, ...), y: (batch_size, ...))
  • 5步批次结构:(x: (5, batch_size, ...), y: (5, batch_size, ...))

示例代码:

import jax
import jax.numpy as jnp
from flax import linen as nn
from flax.training import train_state
import optax

# 模拟数据迭代器
def data_iterator(batch_size=32, num_batches=100):
  for _ in range(num_batches):
    x = jnp.random.normal(size=(batch_size, 28, 28))
    y = jnp.random.randint(0, 10, size=(batch_size,))
    yield (x, y)

# 构造5步批次
def get_multistep_batch(iterator, num_steps=5):
  batches = [next(iterator) for _ in range(num_steps)]
  # 堆叠所有批次的x和y,添加step维度
  x_multistep = jax.tree_map(lambda *xs: jnp.stack(xs, axis=0), *[b[0] for b in batches])
  y_multistep = jax.tree_map(lambda *xs: jnp.stack(xs, axis=0), *[b[1] for b in batches])
  return (x_multistep, y_multistep)

2. 多步训练的JIT封装

用jax.lax.scan遍历5个step的批次,累积更新模型参数和优化器状态,同时计算平均损失。相比Python循环,scan是JIT友好的静态循环,能最大化硬件利用率。

先定义单步训练函数(匹配用户提供的单步逻辑)

class SimpleCNN(nn.Module):
  @nn.compact
  def __call__(self, x):
    x = nn.Conv(features=32, kernel_size=(3,3))(x)
    x = nn.relu(x)
    x = nn.avg_pool(x, window_shape=(2,2), strides=(2,2))
    x = x.reshape((x.shape[0], -1))  # 展平
    x = nn.Dense(features=10)(x)
    return x

# 单步训练:更新一次参数
def train_step(state, batch):
  x, y = batch
  def loss_fn(params):
    logits = state.apply_fn({'params': params}, x)
    loss = optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()
    return loss
  grad_fn = jax.value_and_grad(loss_fn)
  loss, grads = grad_fn(state.params)
  new_state = state.apply_gradients(grads=grads)
  return new_state, loss

改成5步多步训练函数(JIT编译)

# 多步训练:用scan遍历5个step
@jax.jit
def multistep_train_step(state, multistep_batch):
  x_steps, y_steps = multistep_batch

  # 定义scan的单步逻辑:输入当前状态+批次,返回新状态+当前损失
  def scan_step(carry, batch):
    current_state, total_loss = carry
    new_state, step_loss = train_step(current_state, batch)
    return (new_state, total_loss + step_loss), step_loss

  # 初始状态:(当前训练状态,累计损失0)
  initial_carry = (state, jnp.array(0.0))
  # 遍历5个step的批次
  (final_state, total_loss), step_losses = jax.lax.scan(
      scan_step, initial_carry, (x_steps, y_steps)
  )
  # 返回更新后的状态和平均损失
  return final_state, total_loss / 5.0

3. 完整训练循环

def main():
  # 初始化模型和训练状态
  rng = jax.random.PRNGKey(42)
  model = SimpleCNN()
  dummy_x = jnp.zeros((1, 28, 28))
  params = model.init(rng, dummy_x)['params']
  tx = optax.adam(learning_rate=1e-3)
  state = train_state.TrainState.create(
      apply_fn=model.apply, params=params, tx=tx
  )

  # 数据迭代器
  iter = data_iterator()

  # 训练10个多步循环(共50步参数更新)
  for epoch in range(10):
    multistep_batch = get_multistep_batch(iter)
    state, avg_loss = multistep_train_step(state, multistep_batch)
    print(f"Epoch {epoch+1}, Avg Loss: {avg_loss:.4f}")

if __name__ == "__main__":
  main()

关键注意点

  • JIT兼容性:jax.lax.scan是静态循环,num_steps必须是编译时已知的常量(这里固定为5);若需动态步数,可结合jax.lax.cond或动态维度,但性能会略有下降。
  • 数据一致性:确保所有批次的形状、预处理逻辑完全一致,否则jnp.stack会报错。
  • 梯度累积等价性:多步训练等价于梯度累积5次再更新参数,可将学习率乘以5,和单步训练效果一致,但硬件利用率更高。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.23 01:02:38