我想讓看網站的人直接用自己的照片試 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。