从梯度下降到神经网络:一次简单的机器学习实验
机器学习的目标,是通过数据学习模型参数,使预测结果尽可能接近真实值。
本文以线性回归为例,介绍损失函数、梯度下降和 Python 实现。
- 一个简单的问题
给定数据:
$x$| $y$ 1| 3 2| 5 3| 7 4| 9 5| 11
可以看出:
$$ y=2x+1 $$
模型为:
$$ \hat y=wx+b $$
目标是学习出 $w=2$、$b=1$。
1.1 损失函数
使用均方误差衡量预测效果:
$$ L(w,b)= \frac{1}{n} \sum_{i=1}^{n} (wx_i+b-y_i)^2 $$
预测越接近真实值,损失越小。
- 梯度下降
我们希望找到使损失最小的参数:
$$ (w^,b^)=\arg\min_{w,b}L(w,b) $$
梯度下降的更新公式为:
$$
\theta_{t+1}
\theta_t-\eta\nabla L(\theta_t) $$
其中 $\eta$ 是学习率。
梯度为:
$$
\frac{\partial L}{\partial w}
\frac{2}{n} \sum_{i=1}^{n} x_i(wx_i+b-y_i) $$
$$
\frac{\partial L}{\partial b}
\frac{2}{n} \sum_{i=1}^{n} (wx_i+b-y_i) $$
然后更新:
$$ w\leftarrow w-\eta\frac{\partial L}{\partial w} $$
$$ b\leftarrow b-\eta\frac{\partial L}{\partial b} $$
- Python 实现
import numpy as np
x = np.array([1, 2, 3, 4, 5], dtype=float) y = np.array([3, 5, 7, 9, 11], dtype=float)
w, b = 0.0, 0.0 learning_rate = 0.01
for epoch in range(1000): y_pred = w * x + b error = y_pred - y
loss = np.mean(error ** 2)
dw = 2 * np.mean(x * error)
db = 2 * np.mean(error)
w -= learning_rate * dw
b -= learning_rate * db
print("w =", w) print("b =", b)
结果会逐渐接近:
w ≈ 2 b ≈ 1
- PyTorch 实现
import torch
x = torch.tensor([[1.0], [2.0], [3.0], [4.0], [5.0]]) y = torch.tensor([[3.0], [5.0], [7.0], [9.0], [11.0]])
model = torch.nn.Linear(1, 1) loss_fn = torch.nn.MSELoss() optimizer = torch.optim.SGD(model.parameters(), lr=0.01)
for _ in range(1000): prediction = model(x) loss = loss_fn(prediction, y)
optimizer.zero_grad()
loss.backward()
optimizer.step()
"loss.backward()" 会自动计算梯度,这个过程叫自动微分。
- 神经网络
多维线性模型可以写成:
$$ \hat{\mathbf y}=X\mathbf w+b $$
加入非线性函数后,就得到简单的神经网络:
$$ \mathbf h=\sigma(W_1\mathbf x+\mathbf b_1) $$
$$ \mathbf y=W_2\mathbf h+\mathbf b_2 $$
常用激活函数 ReLU 为:
$$ \operatorname{ReLU}(x)=\max(0,x) $$
- Transformer 注意力
Transformer 的注意力公式为:
$$
\operatorname{Attention}(Q,K,V)
\operatorname{softmax} \left( \frac{QK^T}{\sqrt{d_k}} \right)V $$
- 总结
从线性回归到神经网络,训练过程都可以概括为:
- 输入数据;
- 前向传播;
- 计算损失;
- 反向传播;
- 更新参数。
核心思想是:
$$
\text{数据}
+
\text{模型}
+
\text{损失函数}
+
\text{优化算法}
\text{机器学习训练} $$
博客 Markdown 渲染测试完成。