【发布时间】:2019-09-23 10:50:45
【问题描述】:
我使用 mxnet.autograd 实现了一个简单的线性回归梯度下降算法。
一切正常,但性能很糟糕。我使用的是普通梯度下降而不是 SGD,但我怀疑这是问题所在……如果我只是使用梯度的解析表达式,则超过 1000 个 epoch 的整个过程大约需要 2 秒,但使用 autograd 可以达到 147 秒。
这是代码的实现
from mxnet import nd, autograd, gluon
import pandas as pd
def main():
# learning algorithm parameters
nr_epochs = 1000
alpha = 0.01
# read data
data = pd.read_csv("dataset.txt", header=0, index_col=None, sep="\s+")
# ---------------------------------
# -- using gradient descent ---
# ---------------------------------
data.insert(0, "x_0", 1, True) # insert column of "1"s as x_0
m = data.shape[0] # number of samples
n = data.shape[1] - 1 # number of features
X = nd.array(data.iloc[:, 0:n].values) # array with x values
Y = nd.array(data.iloc[:, -1].values) # array with y values
theta = nd.zeros(n) # initial parameters array
theta.attach_grad() # declare gradient with respect to theta is needed
# ----------------------------------------------------------
theta, Loss = GradientDescent(X, Y, theta, alpha, nr_epochs)
# ----------------------------------------------------------
print("Theta by gradient descent:")
print(theta)
#--------------#
# END MAIN #
#--------------#
#-------------------#
# loss function #
#-------------------#
def LossFunction(X, Y, theta):
m = X.shape[0] # number of training samples
loss = 0
for i in range(X.shape[0]):
loss = loss + (1 / (2 * m)) * (H(X[i, :], theta) - Y[i]) ** 2
return loss
#----------------#
# hypothesis #
#----------------#
def H(x, theta):
return nd.dot(x, theta)
#----------------------#
# gradient descent #
#----------------------#
def GradientDescent(X, Y, theta, alpha, nr_epochs):
Loss = nd.zeros(nr_epochs) # array containing values of loss function over iterations
for epoch in range(nr_epochs):
with autograd.record():
loss = LossFunction(X, Y, theta)
loss.backward()
Loss[epoch] = loss
for j in range(len(theta)):
theta[j] = theta[j] - alpha * theta.grad[j]
return theta, Loss
if __name__ == "__main__":
main()
瓶颈是调用
theta, Loss = GradientDescent(X, Y, theta, alpha, nr_epochs)
我做错了吗? 我看过其他一些例子,这些例子比我的工作得快得多,有什么我可以修改以减少运行时间的吗? 谢谢!
【问题讨论】:
标签: python-3.x linear-regression gradient-descent mxnet