【问题标题】:numerically stable way to multiply log probability matrices in numpy在numpy中乘以对数概率矩阵的数值稳定方法
【发布时间】:2014-05-13 11:40:22
【问题描述】:

我需要取两个包含对数概率的 NumPy 矩阵(或其他二维数组)的矩阵乘积。天真的方式np.log(np.dot(np.exp(a), np.exp(b))) 显然不是首选。

使用

from scipy.misc import logsumexp
res = np.zeros((a.shape[0], b.shape[1]))
for n in range(b.shape[1]):
    # broadcast b[:,n] over rows of a, sum columns
    res[:, n] = logsumexp(a + b[:, n].T, axis=1) 

有效,但运行速度比 np.log(np.dot(np.exp(a), np.exp(b))) 慢约 100 倍

使用

logsumexp((tile(a, (b.shape[1],1)) + repeat(b.T, a.shape[0], axis=0)).reshape(b.shape[1],a.shape[0],a.shape[1]), 2).T

或其他 tile 和 reshape 组合也可以工作,但运行速度甚至比上面的循环还要慢,因为实际大小的输入矩阵需要大量的内存。

我目前正在考虑用 C 语言编写一个 NumPy 扩展来计算它,但我当然宁愿避免这种情况。有没有一种既定的方法可以做到这一点,或者有人知道执行这种计算的内存密集度较低的方法吗?

编辑: 感谢 larsmans 提供的解决方案(推导见下文):

def logdot(a, b):
    max_a, max_b = np.max(a), np.max(b)
    exp_a, exp_b = a - max_a, b - max_b
    np.exp(exp_a, out=exp_a)
    np.exp(exp_b, out=exp_b)
    c = np.dot(exp_a, exp_b)
    np.log(c, out=c)
    c += max_a + max_b
    return c

使用 iPython 的神奇 %timeit 函数将此方法与上面发布的方法 (logdot_old) 进行快速比较,得出以下结果:

In  [1] a = np.log(np.random.rand(1000,2000))

In  [2] b = np.log(np.random.rand(2000,1500))

In  [3] x = logdot(a, b)

In  [4] y = logdot_old(a, b) # this takes a while

In  [5] np.any(np.abs(x-y) > 1e-14)
Out [5] False

In  [6] %timeit logdot_old(a, b)
1 loops, best of 3: 1min 18s per loop

In  [6] %timeit logdot(a, b)
1 loops, best of 3: 264 ms per loop

显然 larsmans 的方法抹杀了我的方法!

【问题讨论】:

  • 如果你已经了解 C,你可以使用 scipy.weave.blitz 在你的 Python 代码中加入几行 C
  • 唉,scipy.weave 不适用于 python3
  • 在您的示例中,我不认为scipy.misc.logsumexp 正在做您认为的事情-根据the docs b= 参数实际上是exp(a) 的缩放因子,即@ 987654333@.
  • @mart:你为什么将权重解释为概率?
  • Weave 处于弃用周期。任何新代码都应该使用 Cython。

标签: python numpy matrix matrix-multiplication logarithm


【解决方案1】:

logsumexp 通过计算等式的右侧来工作

log(∑ exp[a]) = max(a) + log(∑ exp[a - max(a)])

即,它在开始求和之前提取最大值,以防止exp 中的溢出。在做矢量点积之前也可以应用同样的方法:

log(exp[a] ⋅ exp[b])
 = log(∑ exp[a] × exp[b])
 = log(∑ exp[a + b])
 = max(a + b) + log(∑ exp[a + b - max(a + b)])     { this is logsumexp(a + b) }

但是通过在推导中采取不同的转变,我们得到

log(∑ exp[a] × exp[b])
 = max(a) + max(b) + log(∑ exp[a - max(a)] × exp[b - max(b)])
 = max(a) + max(b) + log(exp[a - max(a)] ⋅ exp[b - max(b)])

最终形式的内部有一个矢量点积。它也很容易扩展到矩阵乘法,所以我们得到了算法

def logdotexp(A, B):
    max_A = np.max(A)
    max_B = np.max(B)
    C = np.dot(np.exp(A - max_A), np.exp(B - max_B))
    np.log(C, out=C)
    C += max_A + max_B
    return C

这会创建两个 A 大小的临时对象和两个 B 大小的临时对象,但每个都可以通过

exp_A = A - max_A
np.exp(exp_A, out=exp_A)

B 也是如此。 (如果输入矩阵可以被函数修改,则可以消除所有的临时矩阵。)

【讨论】:

  • 谢谢!如果这能达到我所希望的性能,我会尝试。
  • 这比原来较慢的解决方案更不稳定。考虑 logdotexp([[0,0],[0,0]], [[-1000,0], [-1000,0]])。
  • @identity-m 你是对的。此方法不能控制所有元素的稳定性。请参阅我的回答,它也处理您提供的反例。
  • logdotexp(np.array([[0., -1000.]]), np.array([[-1000.], [0.]]))) 失败。当它应该在 [[-999.3]] 附近时产生结果 [[-inf]]
【解决方案2】:

假设A.shape==(n,r)B.shape==(r,m)。在计算矩阵乘积C=A*B 时,实际上有n*m 求和。为了在日志空间中工作时获得稳定的结果,您需要在每个求和中使用 logsumexp 技巧。幸运的是,使用 numpy 广播很容易分别控制 A 和 B 的行和列的稳定性。

代码如下:

def logdotexp(A, B):
    max_A = np.max(A,1,keepdims=True)
    max_B = np.max(B,0,keepdims=True)
    C = np.dot(np.exp(A - max_A), np.exp(B - max_B))
    np.log(C, out=C)
    C += max_A + max_B
    return C

注意:

这背后的原因类似于 FredFoo 的回答,但他为每个矩阵使用了一个最大值。由于他没有考虑每个n*m 求和,最终矩阵的某些元素可能仍然不稳定,如其中一个 cmets 中所述。

使用@identity-m 反例与当前接受的答案进行比较:

def logdotexp_less_stable(A, B):
    max_A = np.max(A)
    max_B = np.max(B)
    C = np.dot(np.exp(A - max_A), np.exp(B - max_B))
    np.log(C, out=C)
    C += max_A + max_B
    return C

print('old method:')
print(logdotexp_less_stable([[0,0],[0,0]], [[-1000,0], [-1000,0]]))
print('new method:')
print(logdotexp([[0,0],[0,0]], [[-1000,0], [-1000,0]]))

打印出来的

old method:
[[      -inf 0.69314718]
 [      -inf 0.69314718]]
new method:
[[-9.99306853e+02  6.93147181e-01]
 [-9.99306853e+02  6.93147181e-01]]

【讨论】:

  • 你的方法也不理想。例如采取a = np.array([[-500., 900.]], dtype=np.float64), b = np.array([[900., -500.]], dtype=np.float64)。您的logdotexp 返回-inf,而scipy.special.logsumexp(a_np[0] + b_np[:, 0]) 正确返回400.69314718055995。
【解决方案3】:

Fred Foo 目前接受的答案以及 Hassan 的答案在数值上不稳定(Hassan 的答案更好)。稍后将提供 Hassan 的回答失败的输入示例。我的实现如下:

import numpy as np
from scipy.special import logsumexp

def logmatmulexp(log_A: np.ndarray, log_B: np.ndarray) -> np.ndarray:
    """Given matrix log_A of shape ϴ×R and matrix log_B of shape R×I, calculates                                                                                                                                                             
    (log_A.exp() @ log_B.exp()).log() in a numerically stable way.                                                                                                                                                                           
    Has O(ϴRI) time complexity and space complexity."""
    ϴ, R = log_A.shape
    I = log_B.shape[1]
    assert log_B.shape == (R, I)
    log_A_expanded = np.broadcast_to(np.expand_dims(log_A, 2), (ϴ, R, I))
    log_B_expanded = np.broadcast_to(np.expand_dims(log_B, 0), (ϴ, R, I))
    log_pairwise_products = log_A_expanded + log_B_expanded  # shape: (ϴ, R, I)                                                                                                                                                              
    return logsumexp(log_pairwise_products, axis=1)

就像 Hassan 的回答和 Fred Foo 的回答一样,我的回答的时间复杂度为 O(ϴRI)。他们的答案有空间复杂度 O(ϴR+RI) (我实际上不确定),而不幸的是我的空间复杂度 O(ϴRI) - 这是因为 numpy 可以将 ϴ×R 矩阵乘以 R×I 矩阵而无需分配一个大小为 ϴ×R×I 的附加数组。具有 O(ϴRI) 空间复杂度不是我的方法的固有属性 - 我认为如果你使用循环写出来,你可以避免这种空间复杂度,但不幸的是我不认为你可以使用股票 numpy 函数来做到这一点。

我检查了我的代码实际运行了多少时间,它比常规矩阵乘法慢 20 倍。

您可以通过以下方式知道我的答案在数值上是稳定的:

  1. 显然,除返回线之外的所有线在数值上都是稳定的。
  2. 已知logsumexp 函数在数值上是稳定的。
  3. 因此,我的 logmatmulexp 函数在数值上是稳定的。

我的实现还有另一个不错的属性。如果不使用 numpy,而是在 pytorch 中编写相同的代码或使用另一个具有自动微分功能的库,您将自动获得数值稳定的反向传递。以下是我们如何知道反向传播将在数值上稳定:

  1. 我的代码中的所有函数在任何地方都是可区分的(不像np.max
  2. 显然,通过除返回线之外的所有线的反向传播在数值上是稳定的,因为那里绝对没有发生任何奇怪的事情。
  3. 通常 pytorch 的开发人员知道他们在做什么。因此,相信他们以数值稳定的方式实现了 logsumexp 的反向传递就足够了。
  4. 其实logsumexp的梯度就是softmax函数(参考google“softmax是logsumexp的梯度”或者见https://arxiv.org/abs/1704.00805命题1)。众所周知,softmax 可以以数值稳定的方式计算。所以 pytorch 开发人员可能只是在那里使用 softmax(我实际上没有检查过)。

下面是 pytorch 中的相同代码(以防您需要反向传播)。由于 pytorch 反向传播的工作原理,在正向传递期间,它将保存 log_pairwise_products 张量以用于反向传递。这个张量很大,您可能不希望它被保存 - 您可以在反向传递期间再次重新计算它。在这种情况下,我建议您使用检查点 - 这真的很简单 - 请参阅下面的第二个功能。

import torch
from torch.utils.checkpoint import checkpoint

def logmatmulexp(log_A: torch.Tensor, log_B: torch.Tensor) -> torch.Tensor:
    """Given matrix log_A of shape ϴ×R and matrix log_B of shape R×I, calculates                                                                                                                                                             
    (log_A.exp() @ log_B.exp()).log() and its backward in a numerically stable way."""
    ϴ, R = log_A.shape
    I = log_B.shape[1]
    assert log_B.shape == (R, I)
    log_A_expanded = log_A.unsqueeze(2).expand((ϴ, R, I))
    log_B_expanded = log_B.unsqueeze(0).expand((ϴ, R, I))
    log_pairwise_products = log_A_expanded + log_B_expanded  # shape: (ϴ, R, I)                                                                                                                                                              
    return torch.logsumexp(log_pairwise_products, dim=1)


def logmatmulexp_lowmem(log_A: torch.Tensor, log_B: torch.Tensor) -> torch.Tensor:
    """Same as logmatmulexp, but doesn't save a (ϴ, R, I)-shaped tensor for backward pass.                                                                                                                                                   

    Given matrix log_A of shape ϴ×R and matrix log_B of shape R×I, calculates                                                                                                                                                                
    (log_A.exp() @ log_B.exp()).log() and its backward in a numerically stable way."""
    return checkpoint(logmatmulexp, log_A, log_B)

这是 Hassan 的实现失败但我的实现给出正确输出的输入:

def logmatmulexp_hassan(A, B):
    max_A = np.max(A,1,keepdims=True)
    max_B = np.max(B,0,keepdims=True)
    C = np.dot(np.exp(A - max_A), np.exp(B - max_B))
    np.log(C, out=C)
    C += max_A + max_B
    return C

log_A = np.array([[-500., 900.]], dtype=np.float64)
log_B = np.array([[900.], [-500.]], dtype=np.float64)
print(logmatmulexp_hassan(log_A, log_B)) # prints -inf, while the correct answer is approximately 400.69.

【讨论】:

    【解决方案4】:

    您正在访问resb 的列,它们的locality of reference 很差。尝试的一件事是将这些存储在column-major order

    【讨论】:

    • 我也注意到了这一点,但是对于较大的数组(大小 > 1000),logsumexp 操作占主导地位。
    猜你喜欢
    • 2021-11-14
    • 2021-05-11
    • 1970-01-01
    • 2011-11-12
    • 1970-01-01
    • 1970-01-01
    • 2021-06-02
    • 2016-10-10
    • 1970-01-01
    相关资源
    最近更新 更多