自动微分 入门
loss.backward() 只有一行,却以大约两趟前向的代价,把十亿个参数每一个的梯度都交到你手上。这一页把它拆开 —— 那卷磁带、两个扫描方向、显存,以及在哪三个地方它给你的答案并不是你想要的导数。每一个数字都由旁边那张图算出来。
拿到斜率的三条路
模型是一个带十亿个旋钮的程序,而训练需要知道每个旋钮往哪边拧会让损失变小。
对一个写下来的函数求导,只有三条路:用手推出导函数、把输入推一下看输出动多少、或者对程序本身求导。第三条就是自动微分,前两条则是它存在的理由。
先看东西本身。导数就是斜率:在曲线上放一个点,过它的切线带着一个数 —— f 在那里爬得多快。拖动这个点,看那个数跟着走:
注意,切线是唯一一条既碰到曲线、又贴着它走的直线。在 x = 2 处它读作 28,而 f 读作 49 —— 同一个点上的两件事,训练要的是斜率。
怎么把那个数算出来才是问题。定义给了一条路:在右边 h 处再取一点,画出过这两点的直线,然后让 h 变小。把间距收窄,看两条线合到一起:
看得出来,h = 1 时割线读作 32,切线读作 28,滑块每挪一格差距就窄一点 —— 这里误差正好是 4h。那把 h 取到机器能装下的最小值,答案就精确了?并不会。
下面每个值都是真正的 float64,不是模拟。把 h 缩到 以下,误差就不再往下掉,反而重新爬升。两条轴都是对数轴 —— 一格就是十倍:
注意曲线拐弯的地方。左臂往下,是因为割线还只是弦;右臂往上,是因为 f(x+h) 和 f(x) 的前几位越来越一致,而相减恰好把这些位扔掉了。最好的步长是 1e−8 —— 接近 √ε —— 它换来 16 位里的 8 位正确。到 ,x + h 就是 x,答案是零。
精度砍一半还能忍,代价忍不了。取最朴素的 n 元函数 f = x₁·x₂·…·xₙ,数一数各种方法算出整个梯度要多少次乘法 ——手推公式、有限差分,或反向扫一趟:
因为手推的梯度要为每个分量重搭一次连乘,代价是 n(n−2)。有限差分要算 n+1 次,代价差不多,还额外在第八位就错了。反向扫描代价 3n−2,在 n = 5 之前确实落后。到 它领先 341 倍,而语言模型的 n 以十亿计。
沿一条路径的连乘
全部的活儿由一条规则干完,而这条规则每个人都已经学过。
给程序求导听上去比给公式求导更难,其实更容易:程序早被拆成了足够小的零件,小到每个零件的导数都已经有人写下来了。
这是 y = (2x+3)² 拆成的三个算子。值沿每根线从左往右走;每个方框下面坐着这个算子自己的导数,在值落到的位置上求出来。拖动 x,看两排数一起动:
注意,这三个局部导数是整页仅有的微积分:乘 2 是 2,加常数是 1,平方是 2v。它们谁也不知道对方存在。链式法则说答案就是它们的乘积 —— 2 · 1 · 14 = 28 —— 而这个乘积就是从 x 到 y 的一条路径。
乘积可以按任意顺序乘,后面的一切都出自这一个事实。从左边带一个数往前走:从 1 起步,每跨过一个算子就乘上那个局部导数。按播放,或一个算子一个算子地走:
看得出来,这就是前向模式,它带着的那个数是 ẋ —— x 动时这根线动得多快。它和值同向同行,所以什么都不用存。
现在把同样三次乘法从右往左做。从输出端以 1 起步往回推,被带着的那个数是 x̄ —— 这根线动时 y 动多少。播放它,或:
把两趟并排读:往回是 1、14、14、28,往前是 1、2、2、28。同样三个因子、同样的答案,只是结合顺序相反。对一入一出的情形,两种模式好得一模一样。
那卷磁带
要往回走,前向那一趟就必须留下点东西。
前向模式从不回头。反向扫描是从末端起步的,它要用的每个局部导数都是在去往末端的路上算出来的 —— 总得有谁把它们记住。
于是框架在每个算子执行时写下一行:哪个算子、进去什么、出来什么、以及它的反向规则将来需要留下什么。把程序跑起来,看磁带一行行长出来 —— 它一开始是空的:
注意最后一列。+3 的反向规则是「把梯度原样传过去」,从前向那趟里什么都不要;(·)² 的规则是 2v,所以它留下那一个输入。这一列就是自动微分全部的显存代价。
这一整套就是十二行 Python —— 一个列表、一个边算边往里追加的算子,以及一个倒着走这个列表的循环:
tape = [] # (out, ins, locals)
def mul(a, b):
out = V(a.v * b.v)
tape.append((out, (a, b), (b.v, a.v)))
return out
def backward(y):
y.g = 1.0
for out, ins, loc in reversed(tape):
for x, d in zip(ins, loc):
x.g += out.g * d # += , never =磁带写完,反向就是上面那个倒着跑的 for 循环。每一行把到达它输出端的伴随量乘上自己的局部导数,再把结果推给输入。从下往上读这卷磁带:
不变量,而这就是全部的正确性论证:这些行按执行的逆序回放,所以一个值的所有消费者都在它被读到之前处理完了 —— 也就是说轮到它时,它的伴随量里已经装着 ∂y/∂v 沿每一条从 v 到 y 的路径求和的结果。写成循环断言:v.grad == sum(c.grad * dc_dv for c in consumers(v))。
最后那一列是要付钱的,而且各算子价钱不同。用一个 8 × 1024 × 768 的 fp16 激活来量 —— GPT-2 的形状,12.00 MiB —— 一条规则必须留下的东西跨了一个数量级还多。在四个算子之间切换:
因为规则是个常数,add 是免费的。relu 只需要输入的符号 —— 每元素 1 比特,0.75 MiB。square 留下输入,12.00 MiB;而 两个操作数都留,24.00 MiB。用加法换掉乘法,换掉的就是这张表的一部分。
直链藏了一件事。让 x 被用两次 —— y = x²·(x+1) —— 从 x 到 y 就有两条路径。拖动 x,看值怎么出去、每条路径又带回来什么,以及它们汇合的节点拿这两个数做了什么:
看得出来,那个节点做的是加法。x = 2 时两条路径带回12 和 4,x̄ 出来是 16,也就是 3x² + 2x。正是这个 += 让「沿所有路径求和」在没人枚举路径时自动成立 —— 而在 §06,它也是少写一行就毁掉一次训练的原因。
一趟扫描,一个问题
一趟扫描很便宜。但它不通用。
到这里为止都是一入一出,而这恰好是两种模式打平的那种形状。换一个两输入三输出的程序:它的导数不是一个数,而是一个 3 × 2 的块,每一对「输出与输入」占一格。
前向模式一次只带一个种子。令 ẋ₁ = 1、ẋ₂ = 0,一趟扫描告诉你 x₁ 动时三个输出各动多少 —— 关于 x₂ 一个字也没有。切换种子,看哪几根线亮起来:
注意,每个种子产出三个数,也就是这个块的一列。所以这个 3 × 2 的雅可比矩阵要两趟前向,一个输入一趟,而且没有什么巧妙的种子能一次拿到两列:扫描对种子是线性的,两列独立的答案就要两个独立的种子。
放大来看,这就是全部的代价模型。这是一个 4 × 6 的雅可比矩阵 —— 六输入、四输出 ——每趟前向填一列。推动滑块把它填满:
看得出来,六列要六趟:趟数等于输入的个数,输出的个数完全不进这个公式。四输出的程序和四百输出的程序,对前向模式来说一个价钱。
§02 那趟反向扫描就是同一张图转九十度。给某一个输出种上 1 往回推,学到的是这个输出对所有输入的响应 —— 得到的是一行,不是一列:
四趟而不是六趟,因为趟数现在等于输出的个数。这里没做任何优化:同一条链式法则、同样的局部导数、每趟同样多的乘法。只是结合顺序变了,随之变的是你要为哪个维度付钱。
为什么反向扫描赢
训练是十亿个数进去、一个数出来。
损失是一个标量,所以 loss.backward() 要的是一个 1 × n 的雅可比矩阵 —— 一行、n 列,n 是参数量。 §04 已经把填满它的两种方式各自的价钱算过了。
下面两个块就是这个雅可比矩阵,左边一列一列地填,右边一行一行地填。设定输入和输出各有多少个,读出谁先填完:
注意,八输入一输出时,反向那块一趟就填完,前向那块要八趟。把形状翻过来 —— 对上很多个输出 —— 前向就用同样的论证赢回去。沿块短的那一边扫。
而对损失来说,短边是哪条从无悬念。把两种模式所需的趟数对输入个数画出来,输出固定为一。两条轴都是对数轴,所以一条直的对角线就是严格的正比:
反向是压在 1 上的水平线,前向是那条对角线。在处,这是一趟对一百万趟;一个 70 亿参数的模型,要 70 亿趟前向才能拿到一趟反向已经给出的东西。
一趟扫描不是免费的,但它有上界。Griewank 的「廉价梯度」结论把一次反向梯度封在四次函数求值之内,无论有多少输入;在 Transformer 上实测是前向一份、反向约两份。拖动参数量,看这几根柱子纹丝不动:
看得出来,一步训练是一趟前向的 3 倍,并且被证明不超过 4 倍 —— 1.25 亿参数如此,4050 亿参数也如此。训练每 token 6·N·P FLOPs、推理 2·N·P 就是从这里来的 —— 是算术,不是巧合。梯度不会因为要求导的东西更多而更贵。
它贵在别处。第一行反向被读到之前,磁带必须攥着整趟前向存下的每一个激活,所以显存随深度增长 —— 除非你扔掉一部分、用时再重算。在 96 层的栈上,每第 k 个边界留一个:
不做检查点时峰值是 96 层。在 处做检查点把它压到 20 —— 少了 4.8 倍 —— 因为你存 ⌈96/k⌉ 个边界,重放某一段时再让 k 层活着,而这个和在 √96 附近最小。代价是多一趟前向,每步约多三分之一算力。所以「反向模式只要一趟」的第二层答案是:一趟时间,外加一整卷空间。
精确,但精确的是什么
自动微分精确到最后一位。所以更值得把「它精确的是什么」说清楚。
到目前为止每个数都没有截断误差、也没有抵消。但磁带求导的对象是真正跑过的那个程序,而跑过的程序和你想写的函数并不总是同一个东西。
先从深度学习里最常见的激活函数说起。relu 在零点没有切线 —— 一边斜率 0、一边 1,折角上没有唯一的直线。把点拖到折角上,读一读框架交回来的东西:
注意,它按约定返回 0,而且不警告你。0、1、½ 都是站得住脚的次梯度,PyTorch 选了 0。这几乎从不重要 —— 浮点数正好落在 0 上的概率可以忽略 —— 而在某人用 x - x 手写一条规则、然后纳闷梯度去哪了的那天,它重要得要命。
严重的版本是控制流。一个 Python if 只把真正跑过的那一支写上磁带,别的都不写,所以那个判断从来没变成一行,也就没有导数。把点拖过那道坎:
看着读数纹丝不动。点跨过去时 y 从 0 变成 1,而两边的 dy/dx 都读作 0,因为那次比较从来没被记录下来。这是静默失败:模型照样训练、损失因为别的原因照样下降,而那个离散决策对你算过的每个梯度都不可见。
补救是故意去微分另一个函数。把台阶换成σ(x/τ),它处处有梯度,代价是多出一个温度。把 τ 调低,盯住梯度:
看得出来,τ 一降,曲线就越来越像它替代的那个台阶 —— 而它的梯度塌成一根宽约 τ 的尖峰,窄到在 时几乎每个样本都落在外面,什么也学不到。直通估计器和 Gumbel-softmax 是在管理这笔交易,谁也消不掉它。
最后一个是个坑,而它就是 §03 那个 +=。梯度缓冲区在设计上就是累加的,所以不管你有没有意,它都会累加。先不调那一行跑四个批次,再把它打开:
两个循环都跑得干干净净。没有 zero_grad() 时,.grad 装的是至今所有批次的和,于是第 4 批走了一步大四倍的步子 —— 它不崩溃,只是训得很差,这是最贵的那类 bug。这一页三种失败里有两种都这么静默;第三种是原地改写一个存下的张量,它是响亮失败的那种,因为磁带上的版本计数器会当场抓住它。
完整跑一趟
技术的两半,端到端,归同一个控件管。
§03 那十二行,跑在 §03 那张菱形图上:前向三个算子各往里追加一行,然后反向四次沿边回放把它们读回来。播放它,或一步一步地走:
看得出来,dy/dx = 16,出自一卷三行的磁带。每个算子的规则就是 locals 里那两个数;整页别处再没有微积分,没有步长,也没有抵消。
真实框架多出来的一切,都是搭在这之上的工程:规则用 C++ 而不是 Python、用张量而不是浮点数、程序一旦分支就用图而不是列表,以及 §05 那套检查点,好让磁带还塞得进显存。