理解 BigBird 的區塊稀疏注意力機制
BigBird 解決了標準 Transformer 模型 $O(n^2)$ 的時間與記憶體複雜度問題,使其能夠處理長達 4096 個 token 的序列。透過將全注意力機制(full attention)替換為區塊稀疏注意力機制(block sparse attention),BigBird 在長序列任務(如長文件摘要和長上下文問答)上取得了尖端(state-of-the-art)的成果。
區塊稀疏注意力機制
BigBird 的區塊稀疏注意力是對 BERT 全注意力機制的一種近似,其設計目標是效率而非卓越性。它透過結合三種不同類型的注意力,減少了查詢 token 需要關注的 token 數量:
- 滑動注意力 (Sliding Attention):Token 會關注其鄰近的鄰居。這捕捉了局部依賴關係,因為序列中的單詞通常與鄰近的 token 高度相關。
- 全域 Token (Global Tokens):一小組特定的 token 會關注序列中的每一個其他 token,並且也會被序列中的每一個其他 token 所關注。這使得資訊能比單靠滑動注意力更快速地在序列中傳遞。
- 隨機 Token (Random Tokens):隨機選擇少數 token 來關注其他 token,進一步縮短了資訊在序列圖形中遠距離節點之間傳遞所需的距離。
雖然全注意力機制允許資訊在單個層級中於任意兩個 token 之間傳遞,但區塊稀疏注意力可能需要多個層級才能讓資訊在序列中傳遞。全域與隨機連接的結合,確保了資訊僅需透過少數幾層即可快速傳遞。
技術實作
為了在 GPU 和 TPU 上高效實作,BigBird 使用了基於區塊(block-based)的方法。序列被劃分為大小為 $b$ 的區塊。
依區塊進行注意力計算
- 第一個與最後一個區塊:第一個區塊 ($q_1$) 與最後一個區塊 ($q_n$) 對序列中所有其他 token 執行正常的注意力操作。
- 中間區塊:對於位於區塊 $q_{3:n-2}$ 中的 token,模型會收集全域、滑動與隨機的 key,並僅針對這些選定的 key 計算注意力。
- 邊界區塊:區塊 $q_2$ 與 $q_{n-1}$ 會收集特定子集的 key(包括第一個區塊、最後一個區塊以及鄰近的滑動區塊),以維持稀疏結構。
複雜度比較
隨著序列長度增加,BigBird 顯著降低了計算負擔。與 BERT 的二次方增長相比,BigBird 呈線性增長:
| 注意力類型 | 序列長度 | 時間與記憶體複雜度 |
|---|---|---|
| Original Full (BERT) | 512 | $T$ |
| Original Full (BERT) | 1024 | $4 ✕ T$ |
| Original Full (BERT) | 4096 | $64 ✕ T$ |
| Block Sparse (BigBird) | 1024 | $2 ✕ T$ |
| Block Sparse (BigBird) | 4096 | $8 ✕ T$ |
訓練策略:ITC 與 ETC
BigBird 可以使用兩種不同的策略進行訓練:
- 內部 Transformer 建構 (Internal Transformer Construction, ITC):使用少量的全域 token。它需要的計算量較少,且對於許多任務來說已經足夠。
- 擴展 Transformer 建構 (Extended Transformer Construction, ETC):增加額外的全域 token。這對於問答任務特別有用,因為整個問題可能需要被上下文 token 全域地關注。
與 Hugging Face Transformers 的整合
BigBird 可透過 transformers 函式庫中的 BigBirdModel 使用。關鍵實作細節包括:
- 序列長度要求:序列長度必須是
block_size的倍數。Hugging Face 會自動將序列填充(pad)至該區塊大小最小的倍數。 - 自動回退機制:如果序列長度太短,無法支援所需的全域、隨機與滑動 token,或者模型被用作解碼器 (
BigBirdForCausalLM),函式庫會自動將attention_type切換為original_full。 - 建議用法:作者建議對於短於 1024 個 token 的序列,將
attention_type設定為original_full。
可用的檢查點(checkpoints)包括 bigbird-roberta-base、bigbird-roberta-large 以及 bigbird-base-trivia-itc。