如何在Jax中基于联合累积分布函数计算联合概率密度函数
问题1:Jax中高效计算n阶混合偏导的实现方案
你猜测的重复调用grad方法其实可以得到正确结果,但对于维度n较大的场景性能不算最优,推荐两种更高效的实现方式:
- 前向模式自动微分实现:因为你的需求是对每个输入维度各求一次偏导,输入维度n远小于输出维度(每次求导后输出都是标量),前向模式天生比
grad默认的反向模式效率更高。可以通过嵌套jax.jvp(雅可比向量积)实现,也可以用functools.reduce简化写法:from functools import reduce import jax import jax.numpy as jnp def get_mixed_deriv(f, n_dim): def _step(g, i): return lambda x: jax.jvp(g, (x,), (jnp.eye(n_dim)[i],))[1] return reduce(_step, range(n_dim), f) # 调用示例,假设输入维度是3 pdf_fn = get_mixed_deriv(cumulative, n_dim=3) # 输入x是长度为3的Jax数组,返回值就是对应的联合概率密度 pdf_val = pdf_fn(x) - 高阶导数专用API优化:如果n较大,可以用Jax的
jax.experimental.jet模块,它对高阶导数的计算做了算子融合优化,比朴素嵌套求导的内存占用和运行速度都好很多。
问题2:导数函数的调用复杂度
不管用哪种实现方式,你的需求是对每个不同的维度各求1次偏导,不属于对单一变量求n阶导数的场景,所以不会出现指数级的复杂度增长:
- 朴素嵌套反向模式
grad的复杂度是O(n * C),其中C是原cumulative函数的执行复杂度,系数通常是2~4倍的原函数开销,也就是整体复杂度相比原函数提升O(n)量级 - 前向模式或者jet优化的实现,系数会更小,常数开销更低,整体依然是
O(n)的复杂度提升,不会出现常数级或者更高阶的增长。
问题3:单个变量偏导数的计算方案
不需要计算全梯度再取对应索引,直接通过jax.jvp指定单位切线向量即可实现单维度偏导的计算,开销和输入总维度无关,只和原函数复杂度成正比:
def get_single_partial(f, k): # 返回对输入数组第k个分量求偏导的函数 def partial_deriv(x): tangent = jnp.zeros_like(x).at[k].set(1.0) return jax.jvp(f, (x,), (tangent,))[1] return partial_deriv # 调用示例:获取对第2个分量的一阶偏导函数 partial_k2 = get_single_partial(cumulative, k=2)
如果你需要基于单维度偏导循环计算最终的混合偏导,直接嵌套调用上述的单维度偏导函数即可,逻辑和混合偏导的实现一致。
内容的提问来源于stack exchange,提问作者Ben
相关产品推荐
相关产品推荐

