面向 AI 的幂:为什么上下文翻倍会让注意力工作量变成四倍
学习指数律与根式,然后使用对数刻度探索器和 NumPy 观察为什么模型维度或上下文翻倍会产生四倍数量。
长度为 个 token 的序列包含 个有序 token 对。 把序列加倍到 个 token 后,配对数变成 ——是原来的四倍。数字 本身没有改变;改变的是对它施加的运算。
这种运算叫作幂。幂把重复乘法压缩成小小的上标,而上标告诉我们:当输入改变尺度时,一个量会如何响应。本课结束时,你将能够处理正、零、负数和分数指数;把根看作逆幂并注意实数域限制;还能够判断 AI 中的量何时线性增长、何时二次增长。
本课建立在第 2 课把函数看作可复用规则的观点之上。这里的规则可以是 ,问题则是当 改变时输出如何变化。官方 Mathematics for Machine Learning 配套网站将数学基础与使用这些基础的机器学习系统分开介绍。它的导论章节把这些基础作为理解模型及其假设的工具;下面的指数代数是本课程原创的先修衔接内容。
一个上标控制增长规则
在 中, 是底数, 是指数。指数说明底数如何参与运算。
“把底数自乘 次”是直观理解,但它只定义了正整数指数。形式化的起点是:
零指数和负指数扩展了这个定义,同时保留相同的代数律。对于任意非零 ,
条件 很重要。负幂会产生倒数,而除以零没有定义。例如,
而 没有实数值。
零指数规则来自商法则,而不是“完全没有乘法”。若 ,那么
两个表达式只有在 时才能相等。由于这个约分用到了 ,它不能决定 ;这个边界需要结合具体语境处理。
指数律说明每一个因子
对于正整数 和 ,同底数幂相乘时,重复因子连接起来:
同样的计数思路给出下面的核心规律。负指数要求底数非零,分数指数则要求留意下一节介绍的定义域。
| 运算 | 规律 | 必要提醒 |
|---|---|---|
| 同底数相乘 | 两个幂都必须有定义 | |
| 同底数相除 | ||
| 幂的幂 | 对分数幂检查实数域 | |
| 积的幂 | 整数 总是安全;实数指数要谨慎 |
有一个很容易误用的模式并不是指数律:
取 、、,左边是 ,而错误提出的右边是 。括号里的加法必须先完成,或者使用分配律展开后再处理幂。
根反向执行幂,但带有定义域
非负实数 的主平方根写作 ,它是平方等于 的非负实数 :
因此 。虽然 和 都等于 ,根号 指的是非负根。解方程 是另一项任务,它有两个实数解 和 。
当 时,平方根就是二分之一幂:
更一般地,对正数 和正整数 ,。有理数指数结合了根与幂:
平方根不表示“除以二”。 是 ,而 ,因为 。被二分的是指数:对实数 ,。更一般地,,不总是 ;当 时,两边都等于 。
翻倍测试揭示线性与二次增长
假设一个量遵循幂规则
其中 是输入规模, 是缩放指数, 是比较过程中不变的常数。如果输入乘以尺度因子 ,那么
因此输出改变了 倍。翻倍时 :
| 指数 | 规则 | 的效果 | 本课名称 |
|---|---|---|---|
| 倍 | 线性增长 | ||
| 倍 | 二次增长 | ||
| 倍 | 三次增长 | ||
| 是原来的 | 逆向增长 |
这个探索器在两条轴上以等距放置 的幂。这是对数刻度:相等的步长表示相等的乘法因子,而不是相等的加法。用指针或方向键移动序列长度控件,比较直线实线 与虚线方形曲线 。
选择 128 到 8,192 个 token 之间、取值为 2 的幂的长度。方向键每次改变一次翻倍。
- 选定长度
- 1,024
- 线性增长
- 8×
- 二次增长
- 64×
- n² 个得分位置
- 1,048,576
当 n = 1,024 时,线性增长是基线的 8 倍,二次增长是 64 倍,共有 1,048,576 个得分位置。
将图表数据显示为表格
| 序列长度 n | 线性 n | 二次 n² | 得分位置 n² | 当前状态 |
|---|---|---|---|---|
| 128 | 1× | 1× | 16,384 | |
| 256 | 2× | 4× | 65,536 | |
| 512 | 4× | 16× | 262,144 | |
| 1,024 | 8× | 64× | 1,048,576 | ← 已选 |
| 2,048 | 16× | 256× | 4,194,304 | |
| 4,096 | 32× | 1,024× | 16,777,216 | |
| 8,192 | 64× | 4,096× | 67,108,864 |
服务器渲染的初始状态选中 。相对于 ,长度是 倍,但平方是 倍。图表下方的表格让每个绘制值都可用,不依赖图形形状或 JavaScript。
密集注意力把 token 对变成方格
原始 Transformer 论文用 定义缩放点积注意力:每个查询都会在分数缩放、归一化并用于组合值之前与每个键进行比较。在经典的密集自注意力分析中,论文给出每层复杂度项 ,其中 是序列长度, 是表示宽度(Vaswani 等,Attention Is All You Need)。
即使课程还没有正式介绍矩阵,也可以把矩阵记号读成一个网格。对于 个 token 位置:
- 有 行查询;
- 每一行对 个键各有一个得分位置;
- 因而每个头的完整得分网格包含 个位置。
如果每个查询—键比较使用 个特征,形成这些点积的主要算术项与 成正比。在只改变 并固定 时,就得到开头的翻倍计算:
这个计数是完整密集得分网格的数学性质。它不能保证墙钟时间或峰值内存一定恰好增加 倍。并行硬件、分块、融合或重新计算的内核、掩码,以及避开完整网格的注意力机制,都可能改变实现存储的内容和运行速度。稳妥的说法更窄:经典密集自注意力有 个查询—键得分位置,其标准算术复杂度包含序列长度的二次项。
参数增长取决于哪些维度发生变化
把 个输入特征映射到 个输出特征的密集层,需要为每一个输入—输出对设置一个权重:
可选偏置还会增加 个参数,但本例的主要缩放关系是成对权重表。
假设一层从 个输入、 个输出变为 个输入、 个输出。两个维度都翻倍:
比值使指数显现出来:
如果只有输出宽度翻倍,数量也只会翻倍。把所有参数增长都称为“二次”会掩盖究竟是哪些维度发生了变化。只有当两个相乘的维度一起缩放时,这里才出现平方律。
NumPy 让公式与其逆运算相互检查
NumPy 的 np.logspace 在对数刻度上以相等间隔构造数值;以 为底时,整数端点 到 生成上下文长度 到 。它的 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.]
最后一行检查了非负输入上的逆关系:,因为数组中的每个序列长度都是正数。如果数组可能包含负数 ,正确的恒等式会是 。
边界情况是运算的一部分
指数记号足够紧凑,可能隐藏定义域错误。阅读公式或代码时,应让以下情况保持可见:
- 取决于语境。 上面的规则 假设 ,而正 时 。在 处,这两个扩展相遇,却没有确定一个值。有些组合公式和软件系统为方便而定义它为 ;初等实数指数代数通常把它留作未定义。要说明所采用的约定,不要默默选择。
- 负实数的偶次根不是实数。 没有实数 满足 ,所以 没有实数值。复数可以扩展定义域,但不在本课范围内。
- 负底数与分数指数需要特别小心。 实立方根 存在,因此精确的有理表达式 可以解释为 。不过像
1 / 3这样的浮点指数只是近似值,NumPy 的实值np.power对负底数和非整数指数会返回nan。不要假设软件会还原你想表达的分数。 - 幂一般不能对加法分配。 反例 足以否定这个捷径。
- 根不是把数值除以根的阶数。 ,不是 。在相应定义域假设成立时,根会除以指数。
对数刻度探索器已经暗示了下一个问题。如果 ,哪个运算可以恢复 ?第 8 课将介绍对数作为幂的逆运算,并用它把乘法尺度变成加法步长。在它上线前,Math for AI 课程页面会保持经过验证的发布顺序。
检查你的理解
问题 1
展开幂并求值 3⁴,同时指出底数和指数。
显示分步解答
在 中,底数是 ,指数是 。正整数指数表示底数的重复因子:
分步相乘:
因此 。把 当作答案,是把因子数量和普通乘法混淆了。
问题 2
对于 x ≠ 0,化简 x³x⁻⁵,改写为不含负指数的形式,并在 x = 2 时检查。
显示分步解答
底数相同,所以相乘时指数相加:
负指数表示倒数,假设 使这个倒数有效:
在 时,
化简形式也给出 ,因此计算相互吻合。
问题 3
计算 64⁻²ᐟ³,并解释指数中负号、分子 2 和分母 3 各自的作用。
显示分步解答
负号要求取倒数:
分母 要求取立方根,分子 随后要求平方:
合并这些步骤:
底数 为正且非零,因此根和倒数在实数中都有效。
问题 4
一个密集层的权重从 300 × 200 增长到 600 × 400。计算两个数量,并解释为什么两个维度翻倍会产生四倍变化。
显示分步解答
原层为每个输入—输出对设置一个权重:
较大的层有
它们的比值是
等价地,每个维度都获得了 倍因子,所以乘积获得 倍。这个推理只计算权重;可选偏置为每个输出增加一个数,而不是再增加一张完整的成对表。
问题 5
经典密集自注意力在头宽固定时把序列从 512 个 token 增加到 1,024 个。计算前后的查询—键得分位置,然后说明这个比值能证明什么、不能证明什么。
显示分步解答
完整密集查询—键网格的位置数是 。增加前,
长度翻倍后,
比值为
这证明完整网格有四倍得分位置,并且在头宽固定时,标准密集注意力算术包含长度的二次项。但它不能证明实测延迟或峰值内存一定恰好变化 倍,因为实现和硬件行为也是变量。
问题 6
一名学生声称 (a + b)² = a² + b²。用 a = 2、b = 3 检验该声称,然后给出正确展开式。
显示分步解答
将 、 代入左边:
声称的右边为
由于 ,一个反例就足以否定所声称的恒等式。正确展开式来自对两个因子使用分配律:
当 、 时,结果为 ,与左边一致。
问题 7
在实数语境中,判断 0⁻²、√(−9)、(−8)¹ᐟ³、√((−5)²) 和 0⁰,并分别给出理由或约定警告。
显示分步解答
逐一考虑它们所需的定义域:
- 会除以零,所以未定义。
- 没有实数值,因为没有实数的平方是 。
- 可以读作实立方根 ,因为根的阶数为奇数。不过浮点幂函数可能无法保留这个精确的有理数解释。
- 。主平方根为非负,所以它等于 ,而不是 。
- 取决于语境。零指数的推导假设底数非零,因此不能确定这个情况。公式或软件系统必须说明采用像 这样的值,还是将其留作未定义。
共同的教训是:记号本身不会抹去定义运算时所用的假设。