pytorch深度学习笔记6-线性回归
·
import numpy as np
import torch
import matplotlib.pyplot as plt
#1.创建数据:一个numpy数组作为线性回归的数据集合
x = np.array([1,2,0.5,2.5,2.6,3.1], dtype=np.float32).reshape((-1, 1))
y = np.array([3.7,4.6,1.65,5.68,5.98,6.95], dtype=np.float32).reshape(-1, 1)
#2.创建一个模型:该模型对象继承自pytroch的model 超类,为创建pytorch中创建的模型的范式
class LinearRegressionModel(torch.nn.Module):# 注意的python中子类的继承方式
def __init__(self, input_dim, output_dim): #子类的构造函数的,编写方式,输入参数为input_dim以及output_dim为超参数
super(LinearRegressionModel, self).__init__() #在子类构造函数中对于超类的进行初始化
self.linear = torch.nn.Linear(input_dim, output_dim)# 在子类的构造函数中将构造函数的两个参数传递给到pytroch linear,定义线性层,
# 返回一个模型的实例对象
def forward(self, x): # 重写向前传播的函数
out = self.linear(x) #调用linear对象
return out
input_dim = 1#定义输入以及输出维度
output_dim = 1
model = LinearRegressionModel(input_dim, output_dim) #创建线性回归模型
criterion = torch.nn.MSELoss() #定义L2损失函数
learning_rate = 0.01 #定义学习率,即反向传播计算完梯度后,更新的量的大小
optimizer = torch.optim.SGD(model.parameters(), lr=learning_rate) #定义优化函数,用于根据损失更新参数
#使用反向传播的计算梯度更新优化的方式进行线性回归的计算
for epoch in range(100):
epoch += 1
# Convert numpy array to torch Variable #numpy对象转到tensor对象
inputs = torch.from_numpy(x).requires_grad_() #注意这里有说明需要梯度
labels = torch.from_numpy(y)
# Clear gradients w.r.t. parameters #计算梯度前先清除优化器现有的梯度
optimizer.zero_grad()
# Forward to get output #前向传播,计算结果
outputs = model(inputs)
# Calculate Loss #根据结果计算以及标签值计算损失
loss = criterion(outputs, labels)
# Getting gradients w.r.t. parameters #反向传播计算梯度
loss.backward()
# Updating parameters #根据反向传播的梯度以及设定的学习率更新参数
optimizer.step()
print('epoch {}, loss {}'.format(epoch, loss.item()))
# Purely inference
predicted_y = model(torch.from_numpy(x).requires_grad_()).data.numpy() #调用模型输入数据,得到预测值
print("标签Y:", y)
print("预测Y:", predicted_y)
# Clear figure
plt.clf()
# Get predictions
predicted = model(torch.from_numpy(x).requires_grad_()).data.numpy()
# Plot true data
plt.plot(x, y, 'go', label='True data', alpha=0.5)
# Plot predictions
plt.plot(x, predicted_y, '--', label='Predictions', alpha=0.5)
# Legend and plot
plt.legend(loc='best')
plt.show()
#
更多推荐
所有评论(0)