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

构建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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 18:40:34