构建VJP时遇到不可哈希静态参数问题
解决JAX中VJP处理静态参数的报错问题
问题根源
你遇到的错误是因为:当把预先指定了static_argnums的jit函数传给jax.vjp时,vjp的反向追踪过程会尝试处理静态参数(A、B),导致它们变成不可哈希的JVPTracer类型,违反了静态参数必须可哈希的要求。
解决方案
不需要先对整个func做jit再传给vjp,而是提前固定静态参数,让vjp只处理动态参数即可,有两种常用方式:
方式1:用functools.partial绑定静态参数
from functools import partial import jax # 固定静态参数A和B,生成仅接收动态参数的函数 func_with_static = partial(func, A=A, B=B) # 对绑定后的函数计算VJP,此时仅需传入动态参数variational_params和e f_eval, vjp_func, aux_output = jax.vjp(func_with_static, variational_params, e, has_aux=True) # 传入余切(注意匹配函数输出结构:第一个输出model_params不需要梯度用None,第二个用dlogp) cotangents = (None, dlogp) vjp_result = vjp_func(cotangents)
方式2:用闭包封装静态参数
import jax # 定义仅接收动态参数的函数,静态参数通过闭包传入 def dynamic_func(variational_params, e): return func(variational_params, e, A, B) # 计算VJP f_eval, vjp_func, aux_output = jax.vjp(dynamic_func, variational_params, e, has_aux=True) # 计算向量-雅可比乘积 cotangents = (None, dlogp) vjp_result = vjp_func(cotangents)
额外说明
如果需要对VJP过程做jit优化,可以直接对vjp_func进行jit:
vjp_func_jitted = jax.jit(vjp_func) vjp_result = vjp_func_jitted(cotangents)
内容的提问来源于stack exchange,提问作者hasco641
相关产品推荐
相关产品推荐

