线性模型

使用y = w*x 来拟合数据集。

从0.0到4.0之间,按照0.1的步长依次取得不同的权重值w,计算每个权重w对应的损失函数值,最终,取最小损失函数对应的权重w。这里损失函数使用的是均方误差。

代码:

import numpy as np
import matplotlib.pyplot as plt

x_data = [1.0,2.0,3.0]
y_data = [2.0,4.0,6.0]

#定义模型
def forward(x,w):
    return x * w

#定义损失函数
def loss(x,y,w):
    y_pred = forward(x,w)
    return (y - y_pred)*(y - y_pred)

w_list = []
mse_list = []

for w in np.arange(0.0,4.1,0.1):
    print("w=:",w)
    l_sum = 0
    for x_value,y_value in zip(x_data,y_data):
        y_pred_val = forward(x_value,w)
        loss_val = loss(x_value,y_value,w)
        l_sum += loss_val
        print('\t',x_value,y_value,y_pred_val,loss_val)
    print('MSE=',l_sum / 3)
    w_list.append(w)
    mse_list.append(l_sum / 3)

plt.plot(w_list,mse_list)
plt.xlabel('w')
plt.ylabel('Loss')
plt.show()

运行结果:

根据(w,mse)图可以看出,权重w增加,损失函数先减小后增加。当w = 2.0 时,损失函数为0,达到最小。该线性模型y=w*x的权重w应该取2.0,此时模型完美拟合数据集。

总结:使用y = w*x 模型,不断改变参数w,选择一个最符合数据集的模型。

作业

线性模型:y = w * x + b

代码来自:https://blog.csdn.net/qq_39804263/article/details/139685123?fromshare=blogdetail&sharetype=blogdetail&sharerId=139685123&sharerefer=PC&sharesource=m0_63829662&sharefrom=from_link 如下:

import numpy as np
import matplotlib.pyplot as plt
from mpl_toolkits.mplot3d import Axes3D

x_data = [1.0, 2.0, 3.0]
y_data = [3.0, 4.0, 6.0]

def forward(x, w, b):
    return x * w + b

def loss(x, y, w, b):
    y_pred = forward(x, w, b)
    loss = (y_pred - y) ** 2
    return loss

w_list = np.arange(0.0, 4.1, 0.1)
b_list = np.arange(-2.0, 2.1, 0.1)

# mse_matrix用于存储不同 w,b 组合下的均方误差损失
mse_matrix = np.zeros((len(w_list), len(b_list)))

for i, w in enumerate(w_list):
    for j, b in enumerate(b_list):
        l_sum = 0
        for x_val, y_val in zip(x_data, y_data):
            l_sum += loss(x_val, y_val, w, b)
        mse_matrix[i, j]= l_sum/len(x_data)
W, B = np.meshgrid(w_list, b_list)
fig = plt.figure('Linear Model Cost Value')
ax = fig.add_subplot(111, projection='3d')
ax.plot_surface(W, B, mse_matrix.T, cmap='viridis')
ax.set_xlabel('w')
ax.set_ylabel('b')
ax.set_zlabel('loss')
plt.show()

三个for循环,遍历不同的权重 (w) 和偏置 (b) 组合,计算每一种组合对应的均方误差 (MSE),最终生成一个 MSE 矩阵,同时为后续可视化(如绘制等高线 / 3D 图)生成网格状的 w 和 b 数组

更多推荐