首先,作为从零开始学习机器学习的第一步,我们不需要一开始就处理复杂的数据和模型。先从一个非常简单的训练集开始:
| x | y |
|---|---|
| $0$ | $0$ |
| $1$ | $2$ |
| $2$ | $4$ |
| $3$ | $6$ |
| $4$ | $8$ |
这个训练集的规律很明显,我们一眼就能看出:
$$ y = 2x $$
但是现在先假装我们不知道这个关系,而是尝试用机器学习的方式,让程序自己从数据中找到它。
为了让机器能够学习这个关系,我们先假设模型的形式为:
$$ \hat{y} = wx $$
其中,$\hat{y}$ 表示模型预测出来的结果,$x$ 是输入,$w$ 是模型需要学习的参数。
也就是说,模型当前并不知道 $x$ 和 $y$ 之间到底是什么关系,它只能先通过某个 $w$ 去计算预测值。我们的目标就是让程序不断调整 $w$,直到模型的预测结果尽可能接近训练数据中的真实结果。
不过在开始训练之前,我们先让 $w$ 取一个随机值,看看此时模型会预测出什么结果。
可以先写成下面这个简单的 C 程序:
1 |
|
运行结果:
1 | w : 6.108053 |
可以看出,由于此时的 $w$ 只是随机生成的,模型的预测结果和真实结果相差很大。
但是,“相差很大”只是一个直观感受。如果我们想让机器自己学习,就需要把这个差距变成一个可以计算的数值。这个数值就叫做损失。
损失越大,说明模型预测得越不准确;损失越小,说明模型越接近训练数据中的规律。
在这里,我们使用均方误差(Mean Squared Error, MSE)作为损失函数:
$$ cost = \frac{1}{n} \sum_{i = 0}^{n - 1} (y_i - \hat{y}_i)^2 $$
其中,$n$ 表示训练数据的数量,$y_i$ 表示第 $i$ 个样本的真实值,$\hat{y}_i$ 表示模型对第 $i$ 个样本预测出来的值。
之所以要对误差进行平方,主要有两个原因:
- 如果直接把误差相加,正误差和负误差可能会互相抵消。例如有些预测偏大,有些预测偏小,最后加起来可能接近 0,但这并不代表模型没有出错。
- 平方会让误差始终变成非负数,并且会放大较大的误差。也就是说,预测得越离谱,损失增长得越明显。
1 | float cost(float w) { |
在 main 中调用:
1 | printf("cost : %f\n", cost(w)); |
运行结果:
1 | w : 6.860047 |
现在我们已经能够用一个具体的数值来衡量预测结果和真实结果之间的差距了。那么接下来的问题就是:如何减少这个差距?
由于当前的损失函数本质上是一个关于 $w$ 的函数,所以我们可以观察:当 $w$ 发生一点点变化时,$cost$ 会如何变化。
如果 $w$ 增大一点后,$cost$ 也变大了,说明我们应该让 $w$ 往反方向调整;如果 $w$ 增大一点后,$cost$ 变小了,说明这个方向是有利的。
这个变化趋势可以用导数来描述。
导数的定义是:
$$ f'(x) = \lim_{h \to 0} \frac{f(x + h) - f(x)}{h} $$
在程序中,我们不需要真的取极限,而是可以取一个很小的数 $\epsilon$,例如 $10^{-3}$,用下面的方式近似计算导数:
$$ f'(x) \approx \frac{f(x + \epsilon) - f(x)}{\epsilon} $$
这就是有限差分。
对于当前问题来说,就是:
$$ dcost = \frac{cost(w + \epsilon) - cost(w)}{\epsilon} $$
有了这个值,我们就知道了 $cost$ 在当前 $w$ 附近的变化方向。
接下来就可以按照下面的方式更新 $w$:
$$ w = w - rate \cdot dcost $$
其中,$rate$ 表示学习率。
学习率决定了每次调整参数时走多大一步。学习率太小,训练会很慢;学习率太大,则可能一步跨过最优位置,导致损失在最小值附近来回振荡,甚至无法收敛。
在这里,我们先把学习率设置为:
$$ rate = 10^{-3} $$
1 | float eps = 1e-3; |
运行结果:
1 | w : 5.947466 |
可以看到,经过一次参数更新后,$cost$ 变小了。
这说明我们的调整方向是有效的。不过只更新一次还远远不够,因此接下来我们可以把这个过程重复多次。例如,重复 500 次:
1 | printf("cost : %f\n", cost(w)); |
运行结果:
1 | w : 6.381645 |
可以看到,经过多次调整后,$cost$ 已经非常接近 0 了。
这说明模型的预测结果已经非常接近训练数据中的真实结果。为了确认模型到底学到了什么,我们再把最终的 $w$ 打印出来:
1 | printf("w : %f\n", w); |
运行结果:
1 | w : 6.037195 |
可以看到,最终得到的 $w$ 已经非常接近 2 了。
这正好符合我们一开始观察到的规律:
$$ y = 2x $$
我们的程序根据预测的误差,一次次调整 $w$,让模型逐渐逼近了这组数据背后的真实关系。
这就是一个最简单的训练过程。
最终代码:
1 |
|