【问题标题】:How to compute the joint probability density function from a joint cumulative density function in Jax?如何从 Jax 中的联合累积密度函数计算联合概率密度函数?
【发布时间】: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


    【解决方案1】:

    JAX 通常将渐变视为相对于单个参数,而不是参数中的元素。在这种情况下,一个与您想要做的类似(但不完全相同)的内置函数是 jax.hessian,它计算二阶导数的 hessian 矩阵;例如:

    import jax
    import jax.numpy as jnp
    
    def f(x):
      return jnp.prod(x ** 2)
    
    x = jnp.arange(1.0, 4.0)
    print(jax.hessian(f)(x))
    # [[72. 72. 48.]
    #  [72. 18. 24.]
    #  [48. 24.  8.]]
    

    对于数组中单个元素的高阶导数,我认为您必须手动嵌套渐变。您可以使用如下所示的辅助函数来执行此操作:

    def grad_all(f):
      def gradfun(x):
        args = tuple(x)
        f_args = lambda *args: f(jnp.array(args))
        for i in range(len(args)):
          f_args = jax.grad(f_args, argnums=i)
        return f_args(*args)
      return gradfun
    
    print(grad_all(f)(x))
    # 48.0
    

    【讨论】:

    • 您的函数似乎运行良好,但我得到了负值,我认为这是不可能的。会不会是你的函数是正确的,但是 jax 设法在不应该的地方得到负值?
    • 根据输入的不同,浮点舍入错误可能会导致 JAX 产生负值,而预期不会产生负值。
    • 我明白了,有没有一种简单的方法可以防止这种情况发生?例如,我可以在循环的每个步骤中截取渐变的值吗?
    • 也许吧?没有更多信息很难说。但是对裁剪函数取梯度可能会遇到边界条件的其他问题。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2017-09-20
    • 2010-11-12
    • 2018-06-25
    • 1970-01-01
    • 2013-07-30
    • 2012-11-21
    • 2020-02-22
    相关资源
    最近更新 更多