监督学习:学习过程中可以看到输入和输出值,以及计算出的输出值与实际输出值的对应关系

Linear Model:y^=x*w+b     作为基线

求解目标--> w和b的取值

从随机猜测开始-->评价偏移误差 y(i)^-y(i)-->loss=(y^-y)²    Mean Square Error=1/N*loss

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):
    return x*w

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

w_list=[]
mse_list=[]

for w in np.arange(0.0,4.1,0.1):
    print('w=',w)
    l_sum=0
    for x_val,y_val in zip(x_data,y_data):#用zip拼成x,y的形式
        y_pred_val=forward(x_val)
        loss_val=loss(x_val,y_val)
        l_sum+=loss_val
        print('\t',x_val,y_val,y_pred_val,loss_val)
    print('MSE=',l_sum/len(x_data))
    w_list.append(w)
    mse_list.append(l_sum/3)

plt.plot(w_list,mse_list)
plt.show()

此处只探讨了y=x*w的情况,下面讨论线性函数的一般形式,即y=x*w+b

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=[5.0,7.0,9.0]

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

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


w_range=np.arange(0.0,4.1,0.1)
# print(w_range)
b_range=np.arange(0.0,4.1,0.1)
W,B=np.meshgrid(w_range,b_range)#meshgrid 用来将二者所有交叉点坐标生成一个网格
mse_grid=np.zeros_like(W)#二维空数组来存储每个(w,b)对的MSE
# print(W,'\n','B=',B)

for i,w in enumerate(w_range):
    for j,b in enumerate(b_range):
        l_sum=0
        for x_val,y_val in zip(x_data,y_data):
            loss_val=loss(x_val,y_val,w,b)
            l_sum+=loss_val

            mse_grid[j,i]=l_sum/len(x_data)


min_mse=np.min(mse_grid)
min_index_flat = np.argmin(mse_grid)
min_j, min_i = np.unravel_index(min_index_flat, mse_grid.shape)

# 3. 根据索引找到对应的w和b
optimal_w = w_range[min_i]
optimal_b = b_range[min_j]
print(optimal_b,optimal_w)
# 3. 创建3D图形
fig = plt.figure()
# 使用推荐的方式创建3D坐标轴
ax = fig.add_subplot(projection='3d')

# 4. 绘制曲面图
# X轴是W,Y轴是B,Z轴是对应的MSE值(mse_grid)
# W, B, mse_grid 必须是同样形状的二维数组
ax.plot_surface(W, B, mse_grid, cmap='viridis') # 使用cmap参数添加颜色映射,更美观

# 添加坐标轴标签,让图像更易懂
ax.set_xlabel('Weight (w)')
ax.set_ylabel('Bias (b)')
ax.set_zlabel('Mean Squared Error (MSE)')

# 显示图形
plt.show()

Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐