[論文レビュー] Fast Transformers with Clustered Attention
本稿では、自己注意機構の線形時間計算量近似としてのクラスタリング注意機構を提案する。この手法はクエリをクラスタにグループ化し、注意計算をクラスタ中心に限定することで、計算コストを顕著に削減する。この方法は自動音声認識で最先端の性能を達成でき、25個のクラスタでフルBERTモデルを近似可能であり、GLUEおよびSQuADベンチマークで性能の低下なしに実現する。
Transformers have been proven a successful model for a variety of tasks in sequence modeling. However, computing the attention matrix, which is their key component, has quadratic complexity with respect to the sequence length, thus making them prohibitively expensive for large sequences. To address this, we propose clustered attention, which instead of computing the attention for every query, groups queries into clusters and computes attention just for the centroids. To further improve this approximation, we use the computed clusters to identify the keys with the highest attention per query and compute the exact key/query dot products. This results in a model with linear complexity with respect to the sequence length for a fixed number of clusters. We evaluate our approach on two automatic speech recognition datasets and show that our model consistently outperforms vanilla transformers for a given computational budget. Finally, we demonstrate that our model can approximate arbitrarily complex attention distributions with a minimal number of clusters by approximating a pretrained BERT model on GLUE and SQuAD benchmarks with only 25 clusters and no loss in performance.
研究の動機と目的
- 自己注意機構のTransformerにおける2次時間計算量の問題を解決し、長時間系列データへの適用を制限する要因を軽減すること。
- 推論およびトレーニング時間を著しく短縮しながらも、高いモデル性能を維持する手法を開発すること。
- 微小な精度低下で事前学習済みモデル(例:BERT)を効率的に近似可能にする手法を実現すること。
- 実世界のNLPタスクで見られる複雑でスパースな注意パターンを、この手法が適切に処理できることを示すこと。
- 計算量の増加に伴い、トレーニング時間とCO2排出量を最大50%まで削減できるスケーラビリティを示すこと。
提案手法
- K-meansと局所性に敏感なハッシュ法(LSH)を用いてクエリをクラスタにグループ化し、注意計算の回数を削減する。
- 各クラスタの中心点に対してのみ注意計算を実行することで、ドット積演算の回数を著しく削減する。
- 精度向上のため、各クエリクラスタごとに注意スコアが最も高いキーとの正確なドット積を特定し、計算する。
- 固定されたクラスタ数に対して、シーケンス長に対して線形の計算量を維持する。
- トレーニングおよび推論の両フェーズに適用可能であり、事前学習済みモデルから完全な注意機構を蒸留(distill)するのにも利用可能。
- 完全な自己注意機構との近似誤差を分析するための理論的境界を導出する。
実験結果
リサーチクエスチョン
- RQ1クラスタリング注意機構は、シーケンスモデリングタスクで性能を維持しつつ、線形時間計算量を達成できるか?
- RQ2クラスタリング注意機構は、事前学習済みBERTモデルに見られる複雑な注意分布をどれほど正確に近似できるか?
- RQ3標準的なTransformerと同等の計算予算下で、性能が維持されるか?
- RQ4顕著な精度低下なしに、トレーニング時間とエネルギー消費量を削減できるか?
- RQ5下流タスクで完全な注意機構をほぼ完全に近似するのに必要な最小クラスタ数はどれくらいか?
主な発見
- Switchboard ASRデータセットにおいて、i-clustered注意機構は1回のフォワードパスあたり約50秒の計算予算で、完全な注意機構よりも2ポイント以上低いWER(誤り率)を達成した。
- 12層のモデルで、i-clusteredは1エポックあたりのトレーニング時間を48%短縮(1.91時間 vs. 3.84時間)し、合計収束時間も44%短縮(132.13時間 vs. 228.05時間)した。
- GLUEおよびSQuADベンチマークにおいて、25クラスタでのi-clusteredは、SQuADを除き全タスクで完全なBERTの性能を再現した。SQuADではわずかに劣る結果(F1スコア:0.876 vs. 0.904)であった。
- 25クラスタのクラスタリング注意機構は、GLUEタスクではほぼ性能の低下なしに、SQuADおよびRTEタスクではわずかな低下を示した。これらのタスクは複雑な注意パターンを要する。
- 長時間系列データの処理において、GPUトレーニング時間を50%削減可能であり、直接的にCO2排出量とエネルギー消費量の削減に寄与した。
- 理論的分析により、近似誤差が有界であり、各クラスタごとに高い注意スコアを持つキーを選択することで、誤差を最小化できることが確認された。
より良い研究を、今すぐ始めましょう
論文の読解から最終レビューまで、研究時間を劇的に削減しましょう。
クレジットカード登録不要
このレビューはAIが作成し、人間の編集者が確認しました。