Loading
在浏览器里跑 Transformer,显而易见的做法是 ONNX Runtime Web。这个项目最初用的就是它,而且能跑。
代价是:打包里多一个 WebAssembly 运行时、一条允许 WASM 编译的 Content-Security-Policy 例外,以及一个我无法掌控发布节奏的依赖——而这一切,只是为了在最多二十四个词元的句子上执行四个 Transformer 块。
所以我把它删掉,手写了前向传播。
关键结论: 四层模型不需要一个通用推理引擎。它需要的是大约七百行在
Float32Array上的算术——然后引擎带来的所有约束都消失了。
如今整个运行时都是普通 TypeScript:嵌入、layer norm、每层三次线性投影、缩放点积、逐行 softmax、value 加权求和、带精确 GELU 的两层前馈网络,以及残差连接。连同 WordPiece 分词器在内,约 770 行。
对十二个词的句子,这需要几十毫秒。它运行在 module 类型的 Web Worker 里,主线程从不阻塞,模型运行时动画依然保持帧率。
它换来了:
'wasm-unsafe-eval'。 CSP 收紧为 script-src 'self',而浏览器测试现在断言这个响应头不包含它过去所必需的那条例外。写前向传播花了一个下午。让它与 PyTorch 一致花的时间要长得多——而那才是唯一值得写的部分。
权重以量化形式发布:int8,每个输出通道一个缩放系数,而不是每个张量一个——因为单一系数在前馈层上损失太大。按通道的系数每行多花四个字节,而每一个字节都值。
门限是:相对 fp32 PyTorch 参考,最坏情况注意力误差小于 1e-3,这样显示的第三位小数才可信。界面写 0.612,模型自己的答案就必须四舍五入到 0.612。
朴素的全 int8 测得 2.97e-2,是门限的三十倍。注意力肉眼可见地错了。
相对 fp32 PyTorch 的最坏情况注意力误差,在量化设计的每一步上测得。门限是 1e-3。数字来自 engine/bert.parity.test.ts 的对齐门。
修好它的不是"统一提高精度",而是弄清楚误差在哪里会变得可见,并只在那里花字节。
位置嵌入与类型嵌入用 fp32。 它们是很小的表,而它们的误差会在任何一个块运行之前,直接进入第 0 层的 logits。用 int8 时,仅它们就把对齐卡在 9.4e-3——网络还什么都没做,就已经很粗糙了。
Q 与 K 用 int16。 它们的误差会被 softmax 的指数放大。logit 上的微小扰动,会变成概率上的巨大扰动。
V 与注意力输出用 int16。 它们会写入残差流,误差因此传播到之后的每一层。
除最后一层外,前馈网络都用 int16。 最后一层的前馈不再喂给任何注意力——下游没有任何东西会读它——所以它可以免费地保持 int8。
最终测量:1.218e-4,稳稳落在门限之内,冷启动 6.32 MiB。
上面每一级都是被测量过的假设,而不是经验法则。"全部量化成 int8"和"全部保持 fp32"都很省事,也都错了:前者毁掉输出,后者把几兆字节浪费在根本不在乎精度的表上。
有用的问题从来不是"用什么精度",而是"误差会从哪些张量里逃出来,又会在哪里被放大"。softmax 会放大。残差会传播。一个走进死胡同的前馈网络两者都不会。
这只能靠逐层测量才能知道。对齐测试之所以分层报告误差,正是为此:如果第 0 层就已经粗糙,问题出在嵌入路径;如果误差随深度增长,说明块内算术在漂移。两个不同的 bug、两种不同的修法——而一个汇总数字会把两者都藏起来。
1.2e-4 的对齐说明:在测试过的句子上,这个实现与 PyTorch 一致。它没有说这套算术在一般意义上是正确的——那是关于所有输入的断言,而任何有限的测试都给不出这个结论。
这也正是测试套件不止步于参考句的原因,以及为什么还存在第二个与它毫无共享代码的实现。第 8 篇讲的就是它,以及它找到的那个 bug。
下一篇:从三十一兆字节的表里只下载二十行。