import numpy as np
import matplotlib.pyplot as plt
from matplotlib.widgets import Slider, Button
%matplotlib widget
plt.rcParams["font.family"] = ["sans-serif", "SimHei"]
plt.rcParams['axes.unicode_minus'] = False
def loss_function(w, b):
return (w - 2)**2 + (b + 1)** 2 + 0.1 * np.sin(5*w) + 0.1 * np.sin(5*b)
def compute_gradient(w, b):
dw = 2 * (w - 2) + 0.5 * np.cos(5*w)
db = 2 * (b + 1) + 0.5 * np.cos(5*b)
return dw, db
def gradient_descent(w_init, b_init, learning_rate, num_iterations):
w_history = [w_init]
b_history = [b_init]
loss_history = [loss_function(w_init, b_init)]
w_current, b_current = w_init, b_init
for i in range(num_iterations):
dw, db = compute_gradient(w_current, b_current)
w_current -= learning_rate * dw
b_current -= learning_rate * db
w_history.append(w_current)
b_history.append(b_current)
loss_history.append(loss_function(w_current, b_current))
if i > 0 and abs(loss_history[-1] - loss_history[-2]) < 1e-6:
break
return w_history, b_history, loss_history
fig = plt.figure(figsize=(12, 9))
plt.subplots_adjust(bottom=0.3)
ax1 = fig.add_subplot(121, projection='3d')
ax2 = fig.add_subplot(122)
ax_init_b = plt.axes([0.25, 0.22, 0.65, 0.03])
ax_init_w = plt.axes([0.25, 0.17, 0.65, 0.03])
ax_iterations = plt.axes([0.25, 0.12, 0.65, 0.03])
ax_learning_rate = plt.axes([0.25, 0.07, 0.65, 0.03])
slider_lr = Slider(ax_learning_rate, '学习率', 0.01, 0.5, valinit=0.1)
slider_iter = Slider(ax_iterations, '迭代次数', 10, 500, valinit=100, valstep=10)
slider_init_w = Slider(ax_init_w, '初始w值', -3, 5, valinit=0)
slider_init_b = Slider(ax_init_b, '初始b值', -3, 5, valinit=0)
reset_ax = plt.axes([0.05, 0.07, 0.1, 0.04])
button = Button(reset_ax, '重置', hovercolor='0.975')
w_grid = np.linspace(-3, 5, 100)
b_grid = np.linspace(-3, 5, 100)
W, B = np.meshgrid(w_grid, b_grid)
L = loss_function(W, B)
surf = ax1.plot_surface(W, B, L, cmap='viridis', alpha=0.7, edgecolor='none')
fig.colorbar(surf, ax=ax1, shrink=0.5, aspect=5)
trajectory_line, = ax1.plot([], [], [], 'r-', linewidth=2, label='梯度下降轨迹')
current_point, = ax1.plot([], [], [], 'bo', markersize=8, label='当前位置')
optimal_point, = ax1.plot([2], [-1], [loss_function(2, -1)], 'go', markersize=10, label='最优解')
ax1.set_xlabel('w参数')
ax1.set_ylabel('b参数')
ax1.set_zlabel('损失值')
ax1.set_title('损失函数曲面与梯度下降轨迹')
ax1.legend()
loss_line, = ax2.plot([], [], 'b-', linewidth=2)
ax2.set_xlabel('迭代次数')
ax2.set_ylabel('损失值')
ax2.set_title('损失值随迭代变化')
ax2.grid(True)
def update(val):
"""
滑块参数变化时触发的更新函数,用于重新执行梯度下降并刷新可视化结果
参数 val: 滑块的当前值(由滑块控件自动传入,此处未直接使用但需保留参数位)
"""
learning_rate = slider_lr.val
num_iterations = int(slider_iter.val)
init_w = slider_init_w.val
init_b = slider_init_b.val
w_history, b_history, loss_history = gradient_descent(
init_w, init_b, learning_rate, num_iterations
)
trajectory_line.set_data_3d(w_history, b_history, loss_history)
current_point.set_data_3d([w_history[-1]], [b_history[-1]], [loss_history[-1]])
loss_line.set_data(range(len(loss_history)), loss_history)
ax2.relim()
ax2.autoscale_view()
fig.canvas.draw_idle()
def reset(event):
slider_lr.reset()
slider_iter.reset()
slider_init_w.reset()
slider_init_b.reset()
update(None)
slider_lr.on_changed(update)
slider_iter.on_changed(update)
slider_init_w.on_changed(update)
slider_init_b.on_changed(update)
button.on_clicked(reset)
update(None)
plt.show()