目录

1.线性模型基础概念

2.完整模型构建步骤

3.可视化参数


1.线性模型基础概念

  • 线性模型的数学定义与公式表达:
    y^ = Wx + b
    其中 W 为权重矩阵,b为偏置项。
  • 模型学习过程:
  • 学习过程详细讲解:random guess:w1;训练线性模型y^ = W1x时得到一组y^;不考虑负值带来的影响,只看y与y^差距,选取(y^-y)^2;不断地数据训练,最终找到最优的w使得y与y^差距缩到0,真正拟合,就把(y^-y)^2定义成损失;∑(y^-y)^2/n定义成均方误差(MSE)。
  • 应用场景:回归任务、二分类/多分类任务。

2.完整模型构建步骤

  • 使用numpy数组并定义 loss和 forward 方法:
    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):  # 定义线性模型:y_pred = x * w
        return x * w
    
    def loss(x, y):  # 修正参数分隔符(.→,)
        y_pred = forward(x)
        return (y_pred - y) ** 2  # 修正变量名(y_predy→y_pred)
    
    w_list = []  # 存储不同w值
    mse_list = []  # 存储对应w的均方误差
    
    # 遍历w从0.0到4.0(步长0.1)
    for w in np.arange(0.0, 4.1, 0.1):
        print('w=', w)
        l_sum = 0  # 累计总损失,为了求均方误差=loss/n
        for x_val, y_val in zip(x_data, y_data):
            y_pred_val = forward(x_val)  # 计算预测值
            loss_val = loss(x_val, y_val)  # 计算单个样本损失
            l_sum += loss_val  # 累加总损失
            # 打印当前样本的详细信息(修正xval→x_val)
            print('\t', x_val, y_val, y_pred_val, loss_val)
        # 计算并打印均方误差(修正1_sum→l_sum)
        print('MSE=', l_sum / 3)
        w_list.append(w)  # 修正空格问题
        mse_list.append(l_sum / 3)
    
    # 可选:绘制w与MSE的关系图(直观展示损失曲线)
    plt.plot(w_list, mse_list)
    plt.xlabel('w')
    plt.ylabel('MSE')
    plt.title('Loss Curve')
    plt.show()

3.可视化参数

  • 从w和loss关系看模型训练的整个过程,并不是最后训练的w就是最优的
  • 示例代码:
plt.plot(w_list, mse_list)
plt. ylabel('Loss’)
plt. xlabel('w')
plt. show()

扩展感兴趣的可以用np.meshgrid() 3d图像

更多推荐