【发布时间】:2021-11-16 10:38:36
【问题描述】:
我在 python 中定义了一个联合累积密度函数作为 jax 数组的函数并返回单个值。 比如:
def cumulative(inputs: array) -> float:
...
要获得梯度,我知道我可以做grad(cumulative),但这只是给我累积相对于输入变量的一阶偏导数。
相反,我想做的是计算这个,假设 F 是我的函数,f 是联合概率密度函数:
偏导的顺序无关紧要。
所以,我有几个问题:
- 如何在 Jax 中高效计算?我想我不能只打电话给 grad n 次
- 一旦计算出结果函数,结果函数的调用复杂度是否会比原始函数更高(是增加了 O(n),还是常数,还是其他什么)?
- 或者,我如何计算仅关于输入数组的一个变量而不是整个数组的单个偏导数? (我将重复此 n 次,每个变量一次)
【问题讨论】:
标签: python probability-density probability-distribution automatic-differentiation jax