AI 数学课程 · 第 7 课,共 180

面向 AI 的幂:为什么上下文翻倍会让注意力工作量变成四倍

学习指数律与根式,然后使用对数刻度探索器和 NumPy 观察为什么模型维度或上下文翻倍会产生四倍数量。

分享这篇文章

长度为 2,0482{,}048 个 token 的序列包含 4,194,3044{,}194{,}304 个有序 token 对。 把序列加倍到 4,0964{,}096 个 token 后,配对数变成 16,777,21616{,}777{,}216——是原来的四倍。数字 22 本身没有改变;改变的是对它施加的运算。

这种运算叫作。幂把重复乘法压缩成小小的上标,而上标告诉我们:当输入改变尺度时,一个量会如何响应。本课结束时,你将能够处理正、零、负数和分数指数;把根看作逆幂并注意实数域限制;还能够判断 AI 中的量何时线性增长、何时二次增长。

本课建立在第 2 课把函数看作可复用规则的观点之上。这里的规则可以是 f(n)=n2f(n)=n^2,问题则是当 nn 改变时输出如何变化。官方 Mathematics for Machine Learning 配套网站将数学基础与使用这些基础的机器学习系统分开介绍。它的导论章节把这些基础作为理解模型及其假设的工具;下面的指数代数是本课程原创的先修衔接内容。

一个上标控制增长规则

apa^p 中,aa底数pp指数。指数说明底数如何参与运算。

“把底数自乘 pp 次”是直观理解,但它只定义了正整数指数。形式化的起点是:

零指数和负指数扩展了这个定义,同时保留相同的代数律。对于任意非零 aa

a0=1以及ap=1ap.a^0=1 \qquad\text{以及}\qquad a^{-p}=\frac{1}{a^p}.

条件 a0a\ne0 很重要。负幂会产生倒数,而除以零没有定义。例如,

23=123=18,2^{-3}=\frac{1}{2^3}=\frac{1}{8},

030^{-3} 没有实数值。

零指数规则来自商法则,而不是“完全没有乘法”。若 a0a\ne0,那么

a3a3=1a3a3=a33=a0.\frac{a^3}{a^3}=1 \qquad\text{且}\qquad \frac{a^3}{a^3}=a^{3-3}=a^0.

两个表达式只有在 a0=1a^0=1 时才能相等。由于这个约分用到了 a0a\ne0,它不能决定 000^0;这个边界需要结合具体语境处理。

指数律说明每一个因子

对于正整数 mmnn,同底数幂相乘时,重复因子连接起来:

aman=aam 个因子aan 个因子=am+n.a^m a^n =\underbrace{a\cdots a}_{m\text{ 个因子}} \underbrace{a\cdots a}_{n\text{ 个因子}} =a^{m+n}.

同样的计数思路给出下面的核心规律。负指数要求底数非零,分数指数则要求留意下一节介绍的定义域。

运算规律必要提醒
同底数相乘aman=am+na^m a^n=a^{m+n}两个幂都必须有定义
同底数相除aman=amn\dfrac{a^m}{a^n}=a^{m-n}a0a\ne0
幂的幂(am)n=amn(a^m)^n=a^{mn}对分数幂检查实数域
积的幂(ab)n=anbn(ab)^n=a^n b^n整数 nn 总是安全;实数指数要谨慎

有一个很容易误用的模式并不是指数律:

(a+b)pap+bp一般而言.(a+b)^p\ne a^p+b^p\quad\text{一般而言}.

a=2a=2b=3b=3p=2p=2,左边是 52=255^2=25,而错误提出的右边是 22+32=132^2+3^2=13。括号里的加法必须先完成,或者使用分配律展开后再处理幂。

根反向执行幂,但带有定义域

非负实数 yy主平方根写作 y\sqrt{y},它是平方等于 yy 的非负实数 rr

r=yr0 且 r2=y.r=\sqrt{y} \quad\Longleftrightarrow\quad r\ge0\text{ 且 }r^2=y.

因此 49=7\sqrt{49}=7。虽然 727^2(7)2(-7)^2 都等于 4949,根号 49\sqrt{49} 指的是非负根。解方程 x2=49x^2=49 是另一项任务,它有两个实数解 x=7x=7x=7x=-7

a0a\ge0 时,平方根就是二分之一幂:

a1/2=a.a^{1/2}=\sqrt{a}.

更一般地,对正数 aa 和正整数 qqa1/q=aqa^{1/q}=\sqrt[q]{a}。有理数指数结合了根与幂:

am/q=(aq)m.a^{m/q}=\left(\sqrt[q]{a}\right)^m.

平方根不表示“除以二”。36/236/21818,而 36=6\sqrt{36}=6,因为 62=366^2=36。被二分的是指数:对实数 aaa4=(a4)1/2=a2\sqrt{a^4}=(a^4)^{1/2}=a^2。更一般地,x2=x\sqrt{x^2}=|x|,不总是 xx;当 x=5x=-5 时,两边都等于 55

翻倍测试揭示线性与二次增长

假设一个量遵循幂规则

f(n)=Cnp,f(n)=C n^p,

其中 n>0n>0 是输入规模,pp 是缩放指数,C>0C>0 是比较过程中不变的常数。如果输入乘以尺度因子 k>0k>0,那么

f(kn)=C(kn)p=Ckpnp=kpf(n).f(kn)=C(kn)^p=Ck^pn^p=k^p f(n).

因此输出改变了 kpk^p 倍。翻倍时 k=2k=2

指数 pp规则n2nn\to2n 的效果本课名称
11CnCn21=22^1=2线性增长
22Cn2Cn^222=42^2=4二次增长
33Cn3Cn^323=82^3=8三次增长
1-1Cn1=C/nCn^{-1}=C/n是原来的 21=1/22^{-1}=1/2逆向增长

这个探索器在两条轴上以等距放置 22 的幂。这是对数刻度:相等的步长表示相等的乘法因子,而不是相等的加法。用指针或方向键移动序列长度控件,比较直线实线 nn 与虚线方形曲线 n2n^2

以 128 个 token 为基线比较线性增长与二次增长。图表、读数和可展开数据表显示相同的数值。

选择 128 到 8,192 个 token 之间、取值为 2 的幂的长度。方向键每次改变一次翻倍。

1,024
线性 n二次 n²
对数轴上的线性与二次增长带圆形标记的实线表示线性增长;带方形标记的虚线在每次翻倍时上升两倍远,表示二次增长。竖直点线指示选定的序列长度。1×4×16×64×256×1,024×4,096×1282565121,0242,0484,0968,192序列长度 n(token,log₂)相对于 n = 128 的增长(log₂)
选定长度
1,024
线性增长
8×
二次增长
64×
n² 个得分位置
1,048,576

当 n = 1,024 时,线性增长是基线的 8 倍,二次增长是 64 倍,共有 1,048,576 个得分位置。

将图表数据显示为表格
相对于 128 个 token 的线性与二次增长
序列长度 n线性 n二次 n²得分位置 n²当前状态
1281×1×16,384
2562×4×65,536
5124×16×262,144
1,0248×64×1,048,576← 已选
2,04816×256×4,194,304
4,09632×1,024×16,777,216
8,19264×4,096×67,108,864

服务器渲染的初始状态选中 n=1,024n=1{,}024。相对于 128128,长度是 88 倍,但平方是 82=648^2=64 倍。图表下方的表格让每个绘制值都可用,不依赖图形形状或 JavaScript。

密集注意力把 token 对变成方格

原始 Transformer 论文用 QKTQK^{\mathsf T} 定义缩放点积注意力:每个查询都会在分数缩放、归一化并用于组合值之前与每个键进行比较。在经典的密集自注意力分析中,论文给出每层复杂度项 O(n2d)O(n^2d),其中 nn 是序列长度,dd 是表示宽度(Vaswani 等,Attention Is All You Need)。

即使课程还没有正式介绍矩阵,也可以把矩阵记号读成一个网格。对于 nn 个 token 位置:

  • nn 行查询;
  • 每一行对 nn 个键各有一个得分位置;
  • 因而每个头的完整得分网格包含 n×n=n2n\times n=n^2 个位置。

如果每个查询—键比较使用 dkd_k 个特征,形成这些点积的主要算术项与 n2dkn^2d_k 成正比。在只改变 nn 并固定 dkd_k 时,就得到开头的翻倍计算:

这个计数是完整密集得分网格的数学性质。它不能保证墙钟时间或峰值内存一定恰好增加 44 倍。并行硬件、分块、融合或重新计算的内核、掩码,以及避开完整网格的注意力机制,都可能改变实现存储的内容和运行速度。稳妥的说法更窄:经典密集自注意力有 n2n^2 个查询—键得分位置,其标准算术复杂度包含序列长度的二次项。

参数增长取决于哪些维度发生变化

dind_{\text{in}} 个输入特征映射到 doutd_{\text{out}} 个输出特征的密集层,需要为每一个输入—输出对设置一个权重:

Nweights=dindout.N_{\text{weights}}=d_{\text{in}}d_{\text{out}}.

可选偏置还会增加 doutd_{\text{out}} 个参数,但本例的主要缩放关系是成对权重表。

假设一层从 512512 个输入、2,0482{,}048 个输出变为 1,0241{,}024 个输入、4,0964{,}096 个输出。两个维度都翻倍:

Nold=512×2,048=1,048,576,Nnew=1,024×4,096=4,194,304.\begin{aligned} N_{\text{old}} &= 512\times2{,}048=1{,}048{,}576,\\ N_{\text{new}} &= 1{,}024\times4{,}096=4{,}194{,}304. \end{aligned}

比值使指数显现出来:

NnewNold=(2512)(22,048)5122,048=22=22=4.\frac{N_{\text{new}}}{N_{\text{old}}} =\frac{(2\cdot512)(2\cdot2{,}048)}{512\cdot2{,}048} =2\cdot2 =2^2 =4.

如果只有输出宽度翻倍,数量也只会翻倍。把所有参数增长都称为“二次”会掩盖究竟是哪些维度发生了变化。只有当两个相乘的维度一起缩放时,这里才出现平方律。

NumPy 让公式与其逆运算相互检查

NumPy 的 np.logspace 在对数刻度上以相等间隔构造数值;以 22 为底时,整数端点 771313 生成上下文长度 272^72132^{13}。它的 np.power 将对应数组元素提升到幂, np.sqrt 则逐元素返回非负平方根。

import numpy as np

lengths = np.logspace(7, 13, num=7, base=2, dtype=np.int64)
score_positions = np.power(lengths, 2)
quadratic_growth = score_positions // score_positions[0]
recovered_lengths = np.sqrt(score_positions)

print("lengths:", lengths)
print("score positions:", score_positions)
print("quadratic growth:", quadratic_growth)
print("recovered lengths:", recovered_lengths)
lengths: [ 128  256  512 1024 2048 4096 8192]
score positions: [   16384    65536   262144  1048576  4194304 16777216 67108864]
quadratic growth: [   1    4   16   64  256 1024 4096]
recovered lengths: [ 128.  256.  512. 1024. 2048. 4096. 8192.]

最后一行检查了非负输入上的逆关系:n2=n\sqrt{n^2}=n,因为数组中的每个序列长度都是正数。如果数组可能包含负数 nn,正确的恒等式会是 n2=n\sqrt{n^2}=|n|

边界情况是运算的一部分

指数记号足够紧凑,可能隐藏定义域错误。阅读公式或代码时,应让以下情况保持可见:

  • 000^0 取决于语境。 上面的规则 a0=1a^0=1 假设 a0a\ne0,而正 pp0p=00^p=0。在 000^0 处,这两个扩展相遇,却没有确定一个值。有些组合公式和软件系统为方便而定义它为 11;初等实数指数代数通常把它留作未定义。要说明所采用的约定,不要默默选择。
  • 负实数的偶次根不是实数。 没有实数 rr 满足 r2=9r^2=-9,所以 9\sqrt{-9} 没有实数值。复数可以扩展定义域,但不在本课范围内。
  • 负底数与分数指数需要特别小心。 实立方根 83=2\sqrt[3]{-8}=-2 存在,因此精确的有理表达式 (8)1/3(-8)^{1/3} 可以解释为 2-2。不过像 1 / 3 这样的浮点指数只是近似值,NumPy 的实值 np.power 对负底数和非整数指数会返回 nan。不要假设软件会还原你想表达的分数。
  • 幂一般不能对加法分配。 反例 (2+3)2=2513=22+32(2+3)^2=25\ne13=2^2+3^2 足以否定这个捷径。
  • 根不是把数值除以根的阶数。 36=6\sqrt{36}=6,不是 1818。在相应定义域假设成立时,根会除以指数。

对数刻度探索器已经暗示了下一个问题。如果 2p=82^p=8,哪个运算可以恢复 p=3p=3?第 8 课将介绍对数作为幂的逆运算,并用它把乘法尺度变成加法步长。在它上线前,Math for AI 课程页面会保持经过验证的发布顺序。

检查你的理解

问题 1

展开幂并求值 3⁴,同时指出底数和指数。

显示分步解答

343^4 中,底数是 33,指数是 44。正整数指数表示底数的重复因子:

34=3×3×3×3.3^4=3\times3\times3\times3.

分步相乘:

3×3=9,9×3=27,27×3=81.3\times3=9, \qquad 9\times3=27, \qquad 27\times3=81.

因此 34=813^4=81。把 3×4=123\times4=12 当作答案,是把因子数量和普通乘法混淆了。

问题 2

对于 x ≠ 0,化简 x³x⁻⁵,改写为不含负指数的形式,并在 x = 2 时检查。

显示分步解答

底数相同,所以相乘时指数相加:

x3x5=x3+(5)=x2.x^3x^{-5}=x^{3+(-5)}=x^{-2}.

负指数表示倒数,假设 x0x\ne0 使这个倒数有效:

x2=1x2.x^{-2}=\frac{1}{x^2}.

x=2x=2 时,

2325=8132=14.2^3\cdot2^{-5} =8\cdot\frac{1}{32} =\frac{1}{4}.

化简形式也给出 1/22=1/41/2^2=1/4,因此计算相互吻合。

问题 3

计算 64⁻²ᐟ³,并解释指数中负号、分子 2 和分母 3 各自的作用。

显示分步解答

负号要求取倒数:

642/3=1642/3.64^{-2/3}=\frac{1}{64^{2/3}}.

分母 33 要求取立方根,分子 22 随后要求平方:

642/3=(643)2=42=16.64^{2/3}=\left(\sqrt[3]{64}\right)^2=4^2=16.

合并这些步骤:

642/3=116.64^{-2/3}=\frac{1}{16}.

底数 6464 为正且非零,因此根和倒数在实数中都有效。

问题 4

一个密集层的权重从 300 × 200 增长到 600 × 400。计算两个数量,并解释为什么两个维度翻倍会产生四倍变化。

显示分步解答

原层为每个输入—输出对设置一个权重:

300×200=60,000 个权重.300\times200=60{,}000\text{ 个权重}.

较大的层有

600×400=240,000 个权重.600\times400=240{,}000\text{ 个权重}.

它们的比值是

240,00060,000=4.\frac{240{,}000}{60{,}000}=4.

等价地,每个维度都获得了 22 倍因子,所以乘积获得 2×2=22=42\times2=2^2=4 倍。这个推理只计算权重;可选偏置为每个输出增加一个数,而不是再增加一张完整的成对表。

问题 5

经典密集自注意力在头宽固定时把序列从 512 个 token 增加到 1,024 个。计算前后的查询—键得分位置,然后说明这个比值能证明什么、不能证明什么。

显示分步解答

完整密集查询—键网格的位置数是 n2n^2。增加前,

5122=262,144.512^2=262{,}144.

长度翻倍后,

1,0242=1,048,576.1{,}024^2=1{,}048{,}576.

比值为

1,048,576262,144=4.\frac{1{,}048{,}576}{262{,}144}=4.

这证明完整网格有四倍得分位置,并且在头宽固定时,标准密集注意力算术包含长度的二次项。但它不能证明实测延迟或峰值内存一定恰好变化 44 倍,因为实现和硬件行为也是变量。

问题 6

一名学生声称 (a + b)² = a² + b²。用 a = 2、b = 3 检验该声称,然后给出正确展开式。

显示分步解答

a=2a=2b=3b=3 代入左边:

(2+3)2=52=25.(2+3)^2=5^2=25.

声称的右边为

22+32=4+9=13.2^2+3^2=4+9=13.

由于 251325\ne13,一个反例就足以否定所声称的恒等式。正确展开式来自对两个因子使用分配律:

(a+b)2=(a+b)(a+b)=a2+ab+ba+b2=a2+2ab+b2.\begin{aligned} (a+b)^2 &=(a+b)(a+b)\\ &=a^2+ab+ba+b^2\\ &=a^2+2ab+b^2. \end{aligned}

a=2a=2b=3b=3 时,结果为 4+12+9=254+12+9=25,与左边一致。

问题 7

在实数语境中,判断 0⁻²、√(−9)、(−8)¹ᐟ³、√((−5)²) 和 0⁰,并分别给出理由或约定警告。

显示分步解答

逐一考虑它们所需的定义域:

  1. 02=1/020^{-2}=1/0^2 会除以零,所以未定义。
  2. 9\sqrt{-9} 没有实数值,因为没有实数的平方是 9-9
  3. (8)1/3(-8)^{1/3} 可以读作实立方根 83=2\sqrt[3]{-8}=-2,因为根的阶数为奇数。不过浮点幂函数可能无法保留这个精确的有理数解释。
  4. (5)2=25=5\sqrt{(-5)^2}=\sqrt{25}=5。主平方根为非负,所以它等于 5|-5|,而不是 5-5
  5. 000^0 取决于语境。零指数的推导假设底数非零,因此不能确定这个情况。公式或软件系统必须说明采用像 11 这样的值,还是将其留作未定义。

共同的教训是:记号本身不会抹去定义运算时所用的假设。

资料来源

  1. Mathematics for Machine Learning companion website
  2. Mathematics for Machine Learning book PDF
  3. NumPy documentation: numpy.power
  4. NumPy documentation: numpy.sqrt
  5. NumPy documentation: numpy.logspace
  6. Attention Is All You Need