理解 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-basebigbird-roberta-large 以及 bigbird-base-trivia-itc

Sources