NumPyro

Stars

建立在 JAX 之上的概率编程框架,覆盖 HMC、NUTS、SVI 等贝叶斯推断算法,并利用 JIT、自动微分和多设备并行提升推断效率。

PROBABILISTIC PROGRAMMING · JAXPR

从私有追踪假设回到公共接口

闭包常量参与计算,却不应占据显式参数的 provenance 位置。

PRIVATE TRACE · JAX 0.11.1
invarsclosed constdynamic x
provenancex
[2] ≠ [1]
PUBLIC jax.make_jaxpr
constsclosed constinvarsdynamic x
ClosedJaxpr 对齐
版本变化

闭包常量进入了输入序列

JAX 0.11.1 将闭包常量提升进 jaxpr 后,provenance 动态输入与变量发生错位。

错位本质

私有 API 隐藏了布局假设

问题来自私有 tracing API 对常量布局的隐式假设;版本升级后常量和动态参数共用输入序列,旧映射关系不再成立。

公共接口

ClosedJaxpr 显式拆分常量

改用公共 jax.make_jaxpr 获取 ClosedJaxpr,显式分离常量与动态输入,重建 provenance 输入映射并解除对私有 tracing API 的耦合。

长期兼容

Provenance 再次准确对齐

恢复新版本 JAX 下的 provenance 正确性,并把实现建立在稳定公共接口上,减少后续 JAX 内部改动造成的兼容性风险。