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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.27 19:43:01