微分


文档摘要

微分 微分刻画的是瞬时的变化率。本文件涵盖极限、导数、求导法则、链式法则(反向传播的数学根基),以及机器学习中常见的导数。 在前面的章节里,我们学会了如何把数据表示成向量、用矩阵去变换它。但现实世界里的很多现象并不是静止的:汽车在加速、股价在波动、神经网络的损失随权重更新而变化。微积分(calculus)就是研究变化的数学。 微积分要回答两个问题:此刻变化有多快?(微分)以及,一段时间里累积了多少?(积分)。本节先来回答「有多快」这个问题。 想象你正在开车,瞥了一眼速度表,读数是 60 km/h。这个数字并不是你整段行程的平均速度,而是你此时此刻的瞬时速度。微分给我们的,正是计算这种瞬时变化率的工具。 不过在那之前,我们先回顾一下直线的方程:$y = mx + b$。

微分

微分刻画的是瞬时的变化率。本文件涵盖极限、导数、求导法则、链式法则(反向传播的数学根基),以及机器学习中常见的导数。

  • 在前面的章节里,我们学会了如何把数据表示成向量、用矩阵去变换它。但现实世界里的很多现象并不是静止的:汽车在加速、股价在波动、神经网络的损失随权重更新而变化。**微积分(calculus)**就是研究变化的数学。

  • 微积分要回答两个问题:此刻变化有多快?(微分)以及,一段时间里累积了多少?(积分)。本节先来回答「有多快」这个问题。

  • 想象你正在开车,瞥了一眼速度表,读数是 60 km/h。这个数字并不是你整段行程的平均速度,而是你此时此刻的瞬时速度。微分给我们的,正是计算这种瞬时变化率的工具。

  • 不过在那之前,我们先回顾一下直线的方程:y = mx + b

  • 这是两个量之间最简单的关系。

    • b纵截距(y-intercept),也就是直线与 y 轴相交的位置(当 x = 0 时的起始值)。
    • m斜率(slope),也就是变化率:x 每增加 1 个单位,y 就改变 m
  • 如果 m = 3,直线陡峭上升;如果 m = 0,直线水平;如果 m = -2,直线下降。

  • 斜率的计算公式是 m = \frac{\Delta y}{\Delta x} = \frac{y_2 - y_1}{x_2 - x_1},即「y 变化了多少」与「x 变化了多少」之比。

直线的方程:b 是纵截距,m 是斜率(纵向变化除以横向变化)

  • 一旦知道了 mb,就可以对任意 x 算出 y

  • 例如,若 m = 2b = 3,那么当 x = 5 时:y = 2(5) + 3 = 13

  • 这两个参数完全确定了这条直线,预测任何输出都只是把值代进去而已。

  • 对于直线而言,斜率处处相同。

  • 这个想法可以推广到直线以外。任何函数都是一个把输入映射到输出的规则,一旦你知道了它的公式(参数和形状),就能对任意输入算出输出,并画出结果。

  • y = x^2 给出一条抛物线,y = \sin(x) 给出一条波浪线,y = e^x 给出指数增长。每一个公式都定义了一条特定的曲线,而习惯于「把函数看成一种形状」是理解后续一切内容的前提。

  • 对于直线而言,斜率处处相同。但大多数有趣的函数都是弯曲的,斜率会随点的位置而变化。微积分给了我们一种方法,去求曲线上任意一点的斜率。

  • 我们还需要**极限(limit)**这个概念。极限描述的是:当函数的输入越来越接近某个目标值时,函数值会趋向于什么——而不一定要真正到达那个值。

\lim_{x \to a} f(x) = L
  • 读作:「当 x 趋近于 a 时,f(x) 趋近于 L。」函数在 x = a 处并不一定要真的等于 L,只要能任意接近即可。

  • 举个例子,取 f(x) = \frac{x^2 - 1}{x - 1}。如果你直接代入 x = 1,会得到 \frac{0}{0},这是没有定义的。

  • 但试试接近 1 的值:f(0.9) = 1.9f(0.99) = 1.99f(1.01) = 2.01。输出明显在奔向 2。

  • 从代数上也能看出为什么:把分子因式分解成 (x-1)(x+1),约去 (x-1) 这一项,就得到对所有 x \neq 1 都成立的 f(x) = x + 1。所以当 x \to 1 时,f(x) \to 2

  • 函数在 x = 1 处有一个「洞」,但极限依然存在。

  • 极限是整个微积分的地基,其他一切都建立在它之上。

  • 函数 f(x) 在点 x = a 处的**导数(derivative)**衡量的是瞬时变化率。从几何上看,它就是曲线在该点处切线的斜率。

导数是曲线上某一点处切线的斜率

  • 为了算出这个斜率,我们先在曲线上取两个点,算出穿过它们的直线(割线(secant line))的斜率。然后让第二个点不断靠近第一个点,观察割线的斜率趋向于什么。这就是差商(difference quotient)
f'(a) = \lim_{h \to 0} \frac{f(a + h) - f(a)}{h}

随着 h 缩小,割线趋近于切线

  • 分子 f(a+h) - f(a) 是输出的变化量。分母 h 是输入的变化量。它们的比值是一段极小区间上的平均变化率。当 h \to 0 时,这个平均值就变成了瞬时变化率。

  • 例如,令 f(x) = x^2。在 x = 3 处:

f'(3) = \lim_{h \to 0} \frac{(3+h)^2 - 9}{h} = \lim_{h \to 0} \frac{9 + 6h + h^2 - 9}{h} = \lim_{h \to 0} (6 + h) = 6
  • 所以在 x = 3 处,函数 x^2 正以「每单位输入 6 单位输出」的速率在增长。

  • 如果这个极限存在,就说函数在该点是可导的(differentiable)。为此,函数必须连续(没有跳跃)、光滑(没有尖角),并且在该点的某个邻域内有定义。

  • 如果你能不抬笔、不打结地画出这条曲线,那么它在那里多半是可导的。

  • 每次都从极限定义去算导数太繁琐了。好在有几条法则,能让我们几乎对任何函数都快速求导。

  • 常数法则(constant rule):常数的导数为零。若 f(x) = 5,则 f'(x) = 0。水平线的斜率为零。

  • 幂法则(power rule):求导的主力。把指数拿下来当系数,再把指数减一:

\frac{d}{dx} x^n = n x^{n-1}
  • 例如:\frac{d}{dx} x^3 = 3x^2。三次函数变成了二次函数。这对任意实数指数都成立,包括负数和分数:\frac{d}{dx} x^{-1} = -x^{-2},以及 \frac{d}{dx} \sqrt{x} = \frac{d}{dx} x^{1/2} = \frac{1}{2}x^{-1/2}

  • 和差法则(sum/difference rule):逐项求导即可。

\frac{d}{dx}[f(x) \pm g(x)] = f'(x) \pm g'(x)
  • 乘积法则(product rule):两个函数相乘时,导数可不是简单地把两个导数相乘。正确的做法是:
\frac{d}{dx}[f(x) \cdot g(x)] = f'(x)g(x) + f(x)g'(x)
  • 可以这样记:「第一个的变化率乘以第二个,加上第一个乘以第二个的变化率。」例如,\frac{d}{dx}[x^2 \sin x] = 2x \sin x + x^2 \cos x

  • 商法则(quotient rule):针对两个函数的比值:

\frac{d}{dx}\left[\frac{f(x)}{g(x)}\right] = \frac{f'(x)g(x) - f(x)g'(x)}{[g(x)]^2}
  • 一个有用的口诀:「下乘上导减上乘下导,除以下面的平方。」

  • 链式法则(chain rule):对机器学习而言最重要的一条法则。当函数被复合起来(一个套在另一个里面)时,导数等于链路上各导数的乘积:

\frac{d}{dx} f(g(x)) = f'(g(x)) \cdot g'(x)
  • 可以想象成剥洋葱:先对最外层函数求导(保持内层函数不动),再乘以内层函数的导数。

链式法则:先对外层求导,再乘以内层的导数

  • 例如,\frac{d}{dx} (3x + 1)^5 = 5(3x+1)^4 \cdot 3 = 15(3x+1)^4。外层函数是 (\cdot)^5,内层是 3x+1

  • 链式法则是神经网络中**反向传播(backpropagation)**的数学根基。一个深度网络本身就是一长串复合函数。为了计算损失相对于每个权重是如何变化的,我们从输出层一路反向应用链式法则,回到输入层,在每一步乘上局部的导数。

  • 下面是你最常遇到的几个导数。每一个都能从极限定义推出来,但熟记它们能省下大量时间:

函数 导数 备注
e^x e^x 唯一一个导数等于自身的函数
a^x a^x \ln a 推广的指数函数
\ln x \frac{1}{x} 自然对数
\log_a x \frac{1}{x \ln a} 一般对数
\sin x \cos x
\cos x -\sin x 注意那个负号
\tan x \sec^2 x
  • 指数函数 e^x 非常特别:它是唯一一个等于自身导数的函数。这就是为什么 e 在机器学习中无处不在——从 softmax 激活到概率分布都能看到它。

  • **洛必达法则(L'Hopital's Rule)**专门处理那些会产生「不定型」的极限,比如 \frac{0}{0}\frac{\infty}{\infty}。当直接代入得到这些形式时,你可以分别对分子和分母求导,然后再试一次极限:

\lim_{x \to a} \frac{f(x)}{g(x)} = \lim_{x \to a} \frac{f'(x)}{g'(x)}
  • 条件:fg 都必须在 a 附近可导,且在 a 附近(a 本身除外)g'(x) \neq 0。同时,原来的极限必须确实是不定型。

  • 例如:\lim_{x \to 0} \frac{\sin x}{x}。直接代入得到 \frac{0}{0}。应用洛必达法则:\lim_{x \to 0} \frac{\cos x}{1} = 1。这个极限非常根本,它在信号处理和傅里叶分析中都会出现。

  • 如果结果仍然是不定型,可以反复应用这条法则。比如 \lim_{x \to 0} \frac{1 - \cos x}{x^2} 得到 \frac{0}{0}。第一次应用:\lim_{x \to 0} \frac{\sin x}{2x},仍然是 \frac{0}{0}。第二次应用:\lim_{x \to 0} \frac{\cos x}{2} = \frac{1}{2}

  • 如果两个函数都可导,那么它们的和、差、积、复合以及商(分母非零处)也都可以求导。正因如此,我们才有底气去对那些由简单零件拼装出来的复杂表达式求导。

编程练习(使用 CoLab 或 notebook)

  1. 可视化常见函数。把 x^2\sin(x)e^x 并排画出来,培养「不同公式产生不同形状」的直觉。试试改变参数(比如 2x^2\sin(2x)),观察曲线如何变化。
import jax.numpy as jnp import matplotlib.pyplot as plt x = jnp.linspace(-3, 3, 300) fig, axes = plt.subplots(1, 3, figsize=(12, 3)) axes[0].plot(x, x**2, color="#e74c3c") axes[0].set_title("x² (parabola)") axes[1].plot(x, jnp.sin(x), color="#3498db") axes[1].set_title("sin(x) (wave)") axes[2].plot(x, jnp.exp(x), color="#27ae60") axes[2].set_title("eˣ (exponential)") for ax in axes: ax.axhline(0, color="gray", linewidth=0.5) ax.axvline(0, color="gray", linewidth=0.5) plt.tight_layout() plt.show()
  1. 用 JAX 的自动微分计算 f(x) = x^3 - 2x + 1 在若干点处的导数,并与解析导数 f'(x) = 3x^2 - 2 对比。
import jax import jax.numpy as jnp f = lambda x: x**3 - 2*x + 1 df = jax.grad(f) for x in [0.0, 1.0, 2.0, -1.0]: print(f"x={x:5.1f} autodiff: {df(x):.4f} analytical: {3*x**2 - 2:.4f}")
  1. 数值验证链式法则。定义 f(x) = \sin(x^2),用 jax.grad 计算它的导数,并与解析结果 2x\cos(x^2) 对比。
import jax import jax.numpy as jnp f = lambda x: jnp.sin(x**2) df = jax.grad(f) for x in [0.5, 1.0, 2.0]: auto = df(x) analytical = 2*x * jnp.cos(x**2) print(f"x={x:.1f} autodiff: {auto:.6f} analytical: {analytical:.6f}")
  1. 可视化导数。把 f(x) = x^3 - 3x 和它的导数 f'(x) = 3x^2 - 3 画在同一张图上。注意 f'(x) = 0 的位置正好对应 f 的峰和谷。
import jax import jax.numpy as jnp import matplotlib.pyplot as plt f = lambda x: x**3 - 3*x # jax.grad 只对标量生效;jax.vmap 把它向量化,一次处理一整组输入 df = jax.vmap(jax.grad(f)) x = jnp.linspace(-2.5, 2.5, 200) plt.plot(x, jax.vmap(f)(x), label="f(x)") plt.plot(x, df(x), label="f'(x)", linestyle="--") plt.axhline(0, color="gray", linewidth=0.5) plt.legend() plt.title("A function and its derivative") plt.show()

发布者: 作者: HenryNdubuaku 转发
评论区 (0)
U