【问题标题】:More efficient way of implement this equation in pytorch (or Numpy)在 pytorch(或 Numpy)中实现这个方程的更有效的方法
【发布时间】:2019-06-27 18:44:33
【问题描述】:

我正在实现这个函数的解析形式

其中 k(x,y) 是 RBF 内核 k(x,y) = exp(-||x-y||^2 / (2h))

我的函数原型是

def A(X, Y, grad_log_px,Kxy):
   pass

XYNxD 矩阵,其中N 是批量大小,D 是一个维度。所以X 是一组x,在上面的等式中N 的大小为grad_log_px 是我使用autograd 计算的一些NxD 矩阵。

KxyNxN 矩阵,其中每个条目(i,j) 是RBF 内核K(X[i],Y[j])

这里的挑战是,在上面的等式中,y 只是一个维度为D 的向量。我有点想传入一批y。 (所以要传递矩阵YNxD 大小)

使用循环遍历批量大小的方程很好,但我无法以更简洁的方式实现

这是我尝试的循环解决方案:

def A(X, Y, grad_log_px,Kxy):
   res = []
   for i in range(Y.shape[0]):
       temp = 0
       for j in range(X.shape[0]):
           # first term of equation
           temp += grad_log_px[j].reshape(D,1)@(Kxy[j,i] * (X[i] - Y[j]) / h).reshape(1,D)
           temp += Kxy[j,i] * np.identity(D) - ((X[i] - Y[j]) / h).reshape(D,1)@(Kxy[j,i] * (X[i] - Y[j]) / h).reshape(1,D) # second term of equation
       temp /= X.shape[0]

        res.append(temp)
    return np.asarray(res) # return NxDxD array 

在等式中:grad_{x}grad_{y} 两个维度 D

【问题讨论】:

  • 你试过用 numpy 来实现方程吗?
  • @taurus05 numpy 有更容易实现这一点的功能吗?
  • Numpy 有一组预定义的函数。它们大多是用 c 编写的。因此,如果您使用它,您将不会遇到任何性能问题。
  • 如果您希望我们 numpy 的人们能够提供帮助,我们真的需要更多线索来了解该方程式到底应该做什么。
  • 你能展示你的循环解决方案吗?

标签: python numpy machine-learning pytorch


【解决方案1】:

鉴于我正确地推断出各种术语的所有维度,这里有一个解决方法。但首先是尺寸的摘要(截图,因为它更容易用数学类型设置来解释;请验证它们是否正确):

还要注意第二项的双导数:

其中下标表示样本,上标表示特征。

所以我们可以使用np.einsum(类似torch.einsum)和array broadcasting来创建这两个词:

grad_y_K = (X[:, None, :] - Y) / h * K[:, :, None]  # Shape: N_x, N_y, D
term_1 = np.einsum('ij,ikl->ikjl', grad_log_px, grad_y_K)  # Shape: N_x, N_y, D_x, D_y
term_2_h = np.einsum('ij,kl->ijkl', K, np.eye(D)) / h  # Shape: N_x, N_y, D_x, D_y
term_2_h2_xy = np.einsum('ijk,ijl->ijkl', grad_y_K, grad_y_K)  # Shape: N_x, N_y, D_x, D_y
term_2_h2 = K[:, :, None, None] * term_2_h2_xy / h**2  # Shape: N_x, N_y, D_x, D_y
term_2 = term_2_h - term_2_h2  # Shape: N_x, N_y, D_x, D_y

那么结果由下式给出:

(term_1 + term_2).sum(axis=0) / N  # Shape: N_y, D_x, D_y

【讨论】:

  • 我在生成的小数据上测试它;结果似乎一致;非常感谢
猜你喜欢
  • 1970-01-01
  • 2013-10-26
  • 1970-01-01
  • 1970-01-01
  • 2019-06-16
  • 2019-11-16
  • 1970-01-01
  • 2023-03-18
  • 2013-06-14
相关资源
最近更新 更多