【发布时间】:2021-06-18 13:43:57
【问题描述】:
当我学习用 Python 编写神经网络时,我刚刚编写了以下线性关联网络,它接收 K 输入向量 x_1, ..., x_K 各自长度 L 和 K 各自长度的输出向量 @ 987654326@ 并使用梯度下降找到最佳权重。
由于在调整K、L 和N 时计算时间会迅速爆炸,我正在寻找如何加快计算速度。我发现了 cupy,但在这种情况下,cupy 比 numpy 慢得多。 为什么会这样?将代码更改为 cupy 变体时,我只将每个 np 替换为 cp,因为我将 cupy 导入为 cp。
我也使用过f = njit()(ManyAssociations.fit),但后来我不得不使用return W,而不是写ManyAssociations.weights = W。 有什么方法可以在课堂内使用 njit,或者除此之外还有更好的方法来使用 numba/cuda?在使用第一个函数调用“热身”之后,结果证明它要快得多,但它仍然在K = L = N = 9 附近的上述形状的向量处达到其极限。
还有哪些其他可以加快代码速度的好方法?我能写得更有效率吗?我可以更好地使用 GPU 吗?
import numpy as np
class ManyAssociations:
def fit(x_train, y_train, learning_rate, tol):
L_L = x_train.shape[1]
L_N = y_train.shape[1]
W = np.zeros((L_N, L_L))
for n in range(L_N):
learning = True
w = np.random.rand(L_L)
while learning:
delta = (x_train @ w - y_train[:,n])
grad_E = delta @ x_train
w = w - learning_rate * grad_E
if (grad_E @ grad_E) < tol:
W[n] = w
learning = False
ManyAssociations.weights = W
def predict(x_pred, W):
preds = []
for k in range(x_pred.shape[0]):
preds.append(W @ x_pred[k])
return np.array(preds)
【问题讨论】:
-
您可以尝试使用JAX。它有一个类似于 numpy 的 API,具有自动微分和 GPU 支持。
-
请提供一个可重现的最小示例,否则人们将无法为您提供帮助。
-
@jakub 谢谢你的建议。不幸的是,原来 JAX 还不支持 Windows。
标签: python performance numpy neural-network gpu