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
相关产品推荐
相关产品推荐

