我想让看网站的人直接用自己的照片试 LumiGrid,而且照片不能离开他们的设备。所以整个模型要在浏览器里跑。这篇记录怎么拆、怎么验证,以及一个花了我最多时间的错。
怎么拆
LumiGrid 有三个部分,我分别用最适合的方式实作:
| 部分 | 浏览器里怎么跑 |
|---|---|
| 网格 CNN(读 256² 缩略图,输出 16×16×8 网格) | ONNX,WebAssembly |
| 亮度导引 + 三线性切片 + 8 轮曲线 + 色彩矩阵 | 自己写的 WebGL2 着色器,fp32 |
| NAFNet 细修 | ONNX,WebGPU(没有就退回 WebAssembly) |
中间那段本来是 PyTorch 的 grid_sample,五维的三线性内插在 ONNX 里支持度不一,而且它本来就是「每个像素做一样的事」,最适合写成着色器。每个像素在着色器里:跑一个 3→16→1 的小网络算亮度座标(GELU 用 erf 近似,误差约 1e-7),到网格里取 8 个邻居内插出 36 个系数,套 8 轮曲线,再乘上 3×4 色彩矩阵。缩略图的面积平均也用 JavaScript 照 OpenCV 的 INTER_AREA 重写一遍。
怎么验证
我先用 PyTorch 在 6 张保留测试图上产生参考输出,再用无头 Chrome 跑网页版,逐像素比较。
WebAssembly 版:63.87 dB,最大误差只有 1 个灰阶,等于完全一致。每一段中间结果也对过:切片结果和 PyTorch 差距小于 1e-6。
WebGPU 版:28.75 dB。图看起来大致正常,数字却明显不对。
猜错了三次
我一开始的三个猜测都错了:
- PixelShuffle 导出成 DepthToSpace(CRD 模式)。 改成明确的 reshape + transpose 重新导出:还是 28.75 dB。
- LayerNorm 里的
pow(x, 2)。 WebGPU 的 WGSL 里,pow对负数底数没有定义,而算变异数时差值本来就会是负的。改成d * d:还是 28.75 dB。 - ONNX Runtime Web 的版本。 从 1.23.2 换到 1.30.0:还是 28.75 dB。
三次都一模一样,代表这不是浮点误差,而是某个运算在 WebGPU 上的行为根本不同。
二分
接着改成有系统地找:
- 单独测每一种运算(LayerNorm、信道注意力、闸控、depthwise 卷积、下采样、上采样、一般卷积):WebGPU 和 WebAssembly 全部一致。
- 把真正的模型每一层都拉出来当输出,逐层比较:差异从第一层就出现了。
- 第一层是「把输入和上一步结果串接成 6 个信道,再做 3×3 卷积」。把这两件事拆开各测一次:
| 测试 | WebGPU 与 WebAssembly 最大差距 |
|---|---|
| 只做串接 | 0 |
| 6 个输入信道的 3×3 卷积 | 0.69 |
| 同一个卷积,补零到 8 个信道 | 6 × 10⁻⁷ |
问题就在这里:在这台 Windows 笔记本电脑的 Chrome 上,ONNX Runtime Web 的 WebGPU 后端遇到 6 个输入信道的卷积会算错,补成 8 个就正常。
修正
把细修网络的输入改成 8 个信道:[输入, 上一步结果, 0, 0],第一层卷积多出来的两个信道权重设为 0。数学上完全等价(PyTorch 里比对差距是 0),但避开了那条出错的路径。
修完之后两个后端都是 63.87 dB、最大误差 1 个灰阶。在 1016 × 680 的图上,我的笔记本电脑大约是:网格 CNN 50–250 ms、着色器 35–250 ms、细修 WebGPU 约 1.8–2.1 秒、WebAssembly 约 2.4 秒。
学到的事
- 一定要有数值上的参考答案。 「看起来差不多」会让这种错活很久;28.75 dB 的图乍看完全正常。
- 把中间层拉出来比,比猜测快很多。 我猜了三次,二分只做两轮就找到。
- 我还没有把它缩成最小重现回报给 ONNX Runtime,这是下一件要做的事。
自己试试看:在浏览器运行的 LumiGrid。