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

Jax局部变量在JIT函数中不更新,普通函数正常的问题解决

JAX JIT函数无法感知闭包内变量更新的简便解决办法

我在闭包内定义了weights变量,通过nonlocal语句修改该变量后,未使用jax.jit装饰的mult函数能正常识别更新后的权重,但用@jax.jit装饰的jitted_mult函数完全无法感知变量变化。不想引入Haiku这类框架改造代码,求更简便的解决方法。

问题代码

from typing import Callable, List

import chex
import jax.numpy as jnp
import jax

Weights = List[jnp.ndarray]


@chex.dataclass(frozen=True)
class Model:
    mult: Callable[
        [jnp.ndarray],
        jnp.ndarray
    ]

    jitted_mult: Callable[
        [jnp.ndarray],
        jnp.ndarray
    ]

    weight_updater: Callable[
        [jnp.ndarray], None
    ]


def create_weight():
    return jnp.ones((2, 5))


def wrapper():
    weights = create_weight()

    def mult(input_var):
        return weights.dot(input_var)

    @jax.jit
    def jitted_mult(input_var):
        return weights.dot(input_var)

    def update_locally_created(new_weights):
        nonlocal weights
        weights = new_weights
        return weights

    return Model(
        mult=mult,
        jitted_mult=jitted_mult,
        weight_updater=update_locally_created
    )


if __name__ == '__main__':
    tester = wrapper()
    to_mult = jnp.ones((5, 2))
    for i in range(5):
        print(jnp.sum(tester.mult(to_mult)))
        print(jnp.sum(tester.jitted_mult(to_mult)))

        if i % 2 == 0:
            tester.weight_updater(jnp.zeros((2, 5)))
        else:
            tester.weight_updater(jnp.ones((2, 5)))

        print("*" * 10)

问题原因

JAX的JIT编译机制会在函数第一次调用时,捕获闭包内变量的当前值并固化到XLA计算图中。后续修改闭包变量不会触发JIT函数的重新编译,因此JIT函数会一直使用最初编译时捕获的weights值,完全感知不到后续的更新。

简便解决方案

方法1:将weights作为显式参数传入JIT函数(推荐)

修改JIT函数让它接收weights作为参数,同时通过包装函数对外保持原有调用接口,这样每次调用都会传递最新的权重值:

def wrapper():
    weights = create_weight()

    def mult(input_var):
        return weights.dot(input_var)

    # 让JIT函数显式接收weights参数
    @jax.jit
    def jitted_mult(weights, input_var):
        return weights.dot(input_var)

    # 包装JIT函数,对外隐藏weights参数
    def wrapped_jitted_mult(input_var):
        return jitted_mult(weights, input_var)

    def update_locally_created(new_weights):
        nonlocal weights
        weights = new_weights
        return weights

    return Model(
        mult=mult,
        jitted_mult=wrapped_jitted_mult,
        weight_updater=update_locally_created
    )

这种方式既符合JAX的设计理念,又能充分利用JIT的性能优化——如果weights的形状和类型不变,JAX只会编译一次计算图,后续调用直接复用。

方法2:用动态更新维护可追踪的权重

将weights包装在一个可通过JAX动态更新的结构中(比如单元素数组),让JIT函数始终能追踪到最新值:

def wrapper():
    # 用单元素数组包装权重,支持动态更新
    weights = jnp.array([create_weight()])

    def mult(input_var):
        return weights[0].dot(input_var)

    @jax.jit
    def jitted_mult(input_var):
        return weights[0].dot(input_var)

    # 用JAX的动态更新操作修改权重
    @jax.jit
    def update_locally_created(new_weights):
        nonlocal weights
        weights = jax.lax.dynamic_update_slice(weights, new_weights[None], (0,))
        return weights

    return Model(
        mult=mult,
        jitted_mult=jitted_mult,
        weight_updater=lambda w: update_locally_created(w)
    )

这种方式适合需要频繁更新权重且不想修改函数参数结构的场景。

方法3:通过Host Callback读取最新权重(不推荐)

如果完全不想修改函数结构,可以用JAX的Host Callback绕开JIT的静态捕获机制,但这种方式会破坏JAX的端到端优化,性能较差:

def wrapper():
    weights = create_weight()

    def mult(input_var):
        return weights.dot(input_var)

    @jax.jit
    def jitted_mult(input_var):
        # 通过Host Callback获取最新的闭包权重
        def get_current_weights(_):
            return weights
        current_weights = jax.experimental.host_callback.id_tap(
            get_current_weights, None, result_shape=weights.shape
        )
        return current_weights.dot(input_var)

    def update_locally_created(new_weights):
        nonlocal weights
        weights = new_weights
        return weights

    return Model(
        mult=mult,
        jitted_mult=jitted_mult,
        weight_updater=update_locally_created
    )

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 09:55:32