3 min read

LiteRT 探索 - attention.tflite 分析

Table of Contents

前陣子在讀 LiteRT runtime 的原始碼,想搞懂它怎麼決定模型哪些部分丟給 NPU、哪些退回 CPU。讀到一半卡住,才發現我跳過了一步:要談「切開」一個模型之前,得先回答一個更基本的問題,模型作為一個檔案,實體上到底是什麼東西?

這問題聽起來很蠢,但我認真想了一下發現自己答不出來。

短答案是:模型就是一堆數字,加上一份「拿這些數字做什麼」的說明書。數字是權重,說明書是一張圖。檔案裡幾乎全都是權重,但幾乎所有有趣的工程問題都住在那張圖上。

這邊我們就把一個真實的 .tflite 剖開來看。

「圖」是自然長出來的,不是設計出來的

神經網路本質上就是函數的合成,就這樣:

y = f₄(f₃(f₂(f₁(x))))

每個 f 是一層運算,卷積、矩陣相乘、正規化、activation。輸入流進去,一路變形,輸出跑出來。

這種「一個接一個」的結構是一條鏈。再加上跳接(ResNet 的 skip connection)或多輸入運算(attention 需要 Q、K、V 同時到齊),鏈就變成一張有向無環圖(DAG)。節點是運算,邊是流動的 tensor。

所以「圖」不是誰為了顯得高深而發明的抽象,它是運算依賴關係的直接寫照,那些依賴本來就在那裡。

那為什麼不直接送程式碼過去?

模型是你用 PyTorch 寫的,那是 Python。為什麼不直接把 Python 送到裝置上?

兩個原因,第二個才是真正的答案。

第一,裝置上沒有 Python 直譯器。好吧,這個可以編譯解決。

第二:NPU 根本不執行「程式」。 它是一塊把矩陣乘法做得很快的電路。你沒辦法遞給它一個 for 迴圈,你只能遞給它「這裡有一個卷積、這是它的權重、輸入放在這個位址」。它認得的單位是「運算」,不是指令流。

所以中間需要一個可攜、可分析、可轉換的表示法。這跟編譯器做的事情一模一樣:

C 原始碼  ->  IR  ->  x86 / ARM / RISC-V 機器碼

端側 ML 的對應版本:

PyTorch (Python)  ->  torch.export -> FX Graph  ->  LiteRT FlatBuffer  ->  CPU / GPU / NPU
    創作形式                中間表示 (IR)              可攜的圖              實際執行

Google 的 litert-torch 就是中間那支箭頭。它以前叫 ai-edge-torch,改名還太新,你搜尋大部分會落到舊名字上。

op 是一份 ABI

這是讓我把所有事情串起來的那個切入點。

一個 op 是兩群互不往來的人之間的契約。

用什麼思考
模型作者「我要一層 conv,再接一個 attention」
晶片廠「我為 CONV_2D 寫了一個 kernel,也為 BATCH_MATMUL 寫了一個」

op set 實際上就是一份 ABI。模型作者組合 op,晶片廠實作 op,兩邊不需要知道對方任何其他事。

也因此,「這顆 NPU 支援哪些 op」才會是決定後面一切的問題。先記著這句,那是第二篇的全部主題。

檔案裡實際裝的東西

.tflite 是一個 FlatBuffer。撇開序列化細節,它裝四樣東西:

區段內容
tensors圖裡每一個值:shape、dtype、量化參數、名稱
buffers權重的實際位元組
operators每個 op:opcode、輸入 tensor 索引、輸出 tensor 索引、選項
subgraphs圖本身(一個模型可以有多張)

operator 是用索引指向 tensor,那些索引就是「邊」。這個檔案就是一張序列化的圖,字面意義上的。

打開一個真的來看

我用 tflite 的 Python binding 寫了一支小 dumper,大概 60 行。

第一個拿來拆的是 mnist_quantized.tflite,結果拆出三個只有一個輸入的 ADD,輸出還叫 *_dequant。單輸入的 ADD 根本講不通,我一度以為自己 parser 寫錯,去查了 raw opcode 才確認檔案裡真的寫 0(就是 ADD)。後來想通了:那個檔案住在 tests/models/ 底下,是量化工具的合成測試 fixture,不是能跑的模型。拿它當教材只會把人教歪。

換成 attention.tflite,LiteRT-LM repo 裡的一個測試模型,從 JAX 匯出的單一 attention block。這個就漂亮多了。

file            : attention.tflite  (10,399,276 bytes on disk)
schema version  : 3
subgraphs       : 1
operator codes  : RESHAPE, FULLY_CONNECTED, CAST, DIV, SIN, COS, SLICE, MUL,
                  SUB, ADD, CONCATENATION, TRANSPOSE, BATCH_MATMUL,
                  SELECT_V2, SOFTMAX
weight bytes    : 10,387,852  (99.9% of the file)

99.9% 的檔案是權重。 78 個 tensor、49 個 operator,整份「要算什麼」的描述大約只佔 11 KB。以體積論,圖是誤差等級的東西,但它也是唯一決定這個模型能不能跑在加速器上的部分。

三個輸入直接告訴你 attention block 需要什麼:

tensor 0   BOOL     8x1x100      serving_default_args_3:0    <- attention mask
tensor 1   FLOAT32  8x100x128    serving_default_args_0:0    <- hidden states
tensor 2   INT32    8x100        serving_default_args_1:0    <- token 位置

batch 8、序列長度 100、模型維度 128。

演算法可以直接從 op 列表讀出來

這是我覺得最有意思的部分。以下是 operator 列表,我只加了註解,其他都沒動。先講清楚:這些不是我讀原始碼推出來的,就是檔案裡直接寫的東西。

op[1]  FULLY_CONNECTED   in=[29, 7]   -> 30    ┐
op[3]  FULLY_CONNECTED   in=[29, 6]   -> 32    ├─ Q、K、V 投影
op[5]  FULLY_CONNECTED   in=[29, 5]   -> 34    ┘  (同一個輸入 29,三份權重)

op[8]  CAST              in=[36]      -> 37    ┐
op[9]  DIV               in=[37, 22]  -> 38    ├─ RoPE:從 token 位置
op[11] SIN               in=[39]      -> 40    │   建出角度表
op[12] COS               in=[39]      -> 41    ┘

op[13] SLICE   ┐
op[14] SLICE   │
op[15] MUL     │
op[16] MUL     ├─ rotate-half 套用到 Q
op[17] SUB     │   (x₁cos - x₂sin, x₂cos + x₁sin)
op[18] MUL     │
op[19] MUL     │
op[20] ADD     │
op[21] CONCAT  ┘
op[22]..op[30]  ── 一模一樣的九個 op,套用到 K

op[31] MUL               in=[50, 28]  -> 60    <- 乘上 1/√d 縮放
op[38] BATCH_MATMUL      in=[66, 63]  -> 67    <- Q·Kᵀ,注意力分數
op[41] SELECT_V2         in=[69,68,27]-> 70    <- causal mask(就是那個 BOOL 輸入)
op[42] SOFTMAX           in=[70]      -> 71
op[44] BATCH_MATMUL      in=[72, 65]  -> 73    <- attn · V
op[48] FULLY_CONNECTED   in=[76, 4]   -> 77    <- 輸出投影

帶 RoPE 的 scaled dot-product attention,用 49 個基本運算攤開來寫。這裡沒有 attention op,也沒有 RoPE op,只有切片、相乘、正弦和餘弦。

權重的 shape 洩漏了模型架構

看那三個投影的權重:

tensor 7   FLOAT32  128x128   65,536 bytes   <- Q 投影
tensor 6   FLOAT32   16x128    8,192 bytes   <- K 投影
tensor 5   FLOAT32   16x128    8,192 bytes   <- V 投影

Q 是 128 -> 128,K 和 V 卻是 128 -> 16。這個不對稱不是筆誤。

再往後看,tensor 74 的 shape 是 8x32x100x4:batch 8、32 個 head、序列 100、head 維度 4。所以 query 側是 32 × 4 = 128。key/value 側則是 16 ÷ 4 = 4 個 head

32 個 query head 共用 4 個 key/value head,這是 grouped-query attention(GQA),group size 8。KV head 數量少代表 KV cache 小,那正是 GQA 存在的目的。

這裡要老實說一句:模型檔案本身沒有任何 metadata 告訴你「我用了 GQA」,上面是我從 shape 反推的。但這個推法很硬,128 對 16、head 維度 4,兜不出別的解釋。重點是你不需要任何文件,光看 tensor 的形狀就讀得出模型的架構決策。

為什麼這件事比看起來重要

三個實際後果,它們各自是後續文章的主題。

光有權重不算模型。 同樣這 10 MB 的浮點數,你可以當成卷積 kernel,也可以當成普通矩陣,結果天差地別。圖是說明書,權重是零件,所以你沒辦法「把權重載進 NPU」就了事。

op 列表就是相容性的接觸面。 再看一次那份列表:SINCOSSELECT_V2BATCH_MATMUL。每一個對 NPU 來說都是問號。一顆卷積做得很漂亮的晶片,可能根本沒有 SIN 這顆 kernel,這時候圖就得被切開,一部分跑 NPU、一部分跑 CPU,而每切一刀你就付一次錢。

融合(fusion)發生在這一層。 正因為 RoPE 是攤成九個基本 op 而不是一個,編譯器才有機會把這九個 pattern match 起來,換成單一個硬體原語,前提是晶片廠有提供。這就是為什麼 LiteRT repo 裡會有一份 PATTERN_MATCHING.md


第二篇會看一顆晶片只跑得動這份 op 列表的一部分時會發生什麼事:圖怎麼被切、誰決定切點、什麼時候決定,以及為什麼切得不好反而會比完全不用加速器還慢。

Dump 工具大約 60 行 Python,依賴 tflite 套件。模型是 google-ai-edge/LiteRT-LMschema/testdata/attention.tflite,讀取自 commit effe245