如何在Jax lax scan中使用迭代不变但每次调用可变的输入?
解决Jax lax.scan中迭代不变但调用时可变的非数组输入问题
当你需要在Jax lax.scan中传入迭代过程保持不变、但每次调用scan时会更新的非数组参数(比如Flax的train_state)时,完全不需要生成冗余的重复序列,以下是两种最优实现方式:
方法一:闭包捕获外部参数
将需要固定的参数(如train_state)定义在scan迭代函数的外部,让函数通过闭包直接访问该参数。每次调用scan前更新这个外部参数,迭代函数就会自动使用最新的值,无需将其放入输入序列。
示例代码:
import jax import jax.numpy as jnp from flax.training import TrainState # 初始化示例train_state params = jnp.array([1.0, 2.0]) train_state = TrainState.create( apply_fn=lambda p, x: p @ x, params=params, tx=None ) # 迭代函数通过闭包捕获外部的train_state def scan_step(carry, x): output = train_state.apply_fn(train_state.params, x) new_carry = carry + output return new_carry, output # 第一次调用scan inputs = jnp.array([[1.0, 0.0], [0.0, 1.0]]) init_carry = jnp.array(0.0) final_carry, outputs = jax.lax.scan(scan_step, init_carry, inputs) # 更新train_state后,第二次调用scan new_params = jnp.array([3.0, 4.0]) train_state = train_state.replace(params=new_params) final_carry_new, outputs_new = jax.lax.scan(scan_step, init_carry, inputs)
方法二:利用scan的多输入与in_axes参数
jax.lax.scan支持传入由多个输入组成的树结构,通过in_axes参数可以指定每个输入是否参与迭代(即是否按序列遍历)。对于迭代中不变的参数,将其对应的in_axes设为None,这样scan会在每次迭代中传入同一个参数值,无需重复生成序列。
这种方式更显式,参数依赖关系清晰,适合复杂场景:
import jax import jax.numpy as jnp from flax.training import TrainState # 初始化train_state params = jnp.array([1.0, 2.0]) train_state = TrainState.create( apply_fn=lambda p, x: p @ x, params=params, tx=None ) # 迭代函数将train_state作为额外输入 def scan_step(carry, (x, static_train_state)): output = static_train_state.apply_fn(static_train_state.params, x) new_carry = carry + output return new_carry, output # 输入序列为(inputs, train_state),in_axes指定第一个输入遍历轴0,第二个输入保持不变 inputs = jnp.array([[1.0, 0.0], [0.0, 1.0]]) init_carry = jnp.array(0.0) final_carry, outputs = jax.lax.scan( scan_step, init_carry, (inputs, train_state), in_axes=(0, None) ) # 更新train_state后再次调用 new_params = jnp.array([3.0, 4.0]) train_state = train_state.replace(params=new_params) final_carry_new, outputs_new = jax.lax.scan( scan_step, init_carry, (inputs, train_state), in_axes=(0, None) )
两种方法对比
- 闭包方式更简洁,适合参数较少的场景;
in_axes方式更显式,便于调试和维护,尤其适合需要对scan做jax.jit编译的场景(闭包变量若需编译后更新,需要额外处理静态参数,而in_axes方式更兼容)。
内容的提问来源于stack exchange,提问作者Jabby
相关产品推荐
相关产品推荐

