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

Jax:如何在jax.lax.scan的扫描函数中传递常量参数

JAX scan 中移除不变carry参数的正确做法

针对你遇到的问题——bar在jax.lax.scan全程不变,却被迫作为carry传递的情况,有两种高效且不会触发重复JIT编译的解决方法:

方法1:利用闭包绑定静态参数

把bar作为外层函数的参数,在内部定义foo时直接引用这个闭包变量,这样foo就不需要再从carry中解包bar,也不会因为partial导致重复编译:

import jax
import jax.numpy as jnp

def create_scan_foo(bar):
    def foo(carry, x):
        # 直接使用外层的bar,无需从carry中提取
        current_value = carry
        # 执行你的计算逻辑,比如用bar处理x和current_value
        new_carry = current_value + x * bar
        output = new_carry * 2
        return new_carry, output
    return foo

# 使用示例
initial_carry = jnp.array(0.0)
xs = jnp.arange(5)
bar = jnp.array(2.0)

# 创建绑定了bar的foo函数
scan_foo = create_scan_foo(bar)
final_carry, outputs = jax.lax.scan(scan_foo, initial_carry, xs)

这种方式下,bar会被JAX识别为静态常量,只要bar的值在编译时确定,foo只会被JIT编译一次,不会每次调用都重新编译。

方法2:使用static_broadcasted_args参数(JAX 0.3.14+)

JAX的scan提供了专门的参数static_broadcasted_args,用来传递扫描过程中完全不变的静态参数,不需要放到carry或xs里:

import jax
import jax.numpy as jnp

def foo(carry, x, bar):
    # bar作为静态参数传入,全程不变
    current_value = carry
    new_carry = current_value + x * bar
    output = new_carry * 2
    return new_carry, output

# 使用示例
initial_carry = jnp.array(0.0)
xs = jnp.arange(5)
bar = jnp.array(2.0)

final_carry, outputs = jax.lax.scan(
    foo,
    initial_carry,
    xs,
    static_broadcasted_args=(bar,)
)

如果bar是运行时可能变化的参数,你可以配合jax.jit的static_argnames来标记它,避免每次bar变化都触发重编译:

@jax.jit(static_argnames=["bar"])
def run_scan(initial_carry, xs, bar):
    return jax.lax.scan(
        foo,
        initial_carry,
        xs,
        static_broadcasted_args=(bar,)
    )

为什么之前的方法不行?

  • 用partial(foo, bar=bar)时,如果bar是动态值(编译时无法确定),JAX会认为每次bar变化都对应一个新的函数,从而触发重复编译,导致速度变慢。
  • 把bar放到xs里报错,是因为xs要求每个输入都带有扫描维度(即第一个维度对应扫描步数),而bar没有这个维度。如果硬要这么做,你需要用jnp.repeat(jnp.expand_dims(bar, 0), xs.shape[0], axis=0)来扩展维度,但这会浪费内存,完全没必要。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.24 03:19:55