Introduction
Transformer 모델은 자연어 처리, 컴퓨터 비전, 화학 분자구조 학습과 같은 다양한 분야에서 우수한 성능을 보이며 sequential한 데이터 학습에 자주 사용되었습니다. Transformer의 Multi-Head Attention 모듈은 그림 1, 2와 같이 context 상 한 token이 다른 token에 미치는 중요도를 포착하는 이른바 attention score matrix를 계산해 높은 표현 능력을 자랑하지만, sequence 길이의 제곱에 비례하는 계산 비용이 필요합니다. 이로 인해 Transformer를 문서 요약, 고해상도 이미지 처리, 단백질 구조 모델링과 같이 input sequence의 길이가 긴 작업에 직접 사용하는 데에는 어려움이 있습니다.
그림 1. 기존 Transformer 아키텍처[1,2]
그림 2. 기존 Transformer의 attention 계산식
Sparse한 attention 패턴이나 low-rank approximation을 활용해 Transformer를 가속하는 방법은 이전에도 많이 제시되었으나, 이는 모든 attention layer와 head에 동일한 변경 사항을 적용하는 지나치게 강한 inductive bias가 있어 downstream task를 푸는 과정에서 최적의 결과가 아닌 계산 비용 대비 예측 성능 간 trade-off를 보이기도 합니다. BERT나 GPT-3와 같이 많은 state-of-the-art 시스템들이 dense attention과 sparse attention을 혼합해 활용한다는 점을 생각해보았을 때, 주어진 task를 기반으로 dense 또는 sparse attention 사이를 자체적으로 유연하게 조절할 수 있는 attention 모듈을 개발하는 것이 더 나은 성능을 구현하는 데 도움이 될 것이라고 생각했습니다. 이러한 이유로 각 attention head에 mixed-membership Stochastic Block Model (SBM)을 부여해 attention sparsity뿐만 아니라 주어진 데이터에 맞춰 계산 비용까지 선택할 수 있는 Transformer 변형 모델인 SBM-Transformer를 제안했습니다.
Our Method: SBM-Transformer
SBM-Transformer의 forward step에서 각 attention head는 input으로 주어진 token을 node의 집합으로 보며, query와 key를 연결하는 bipartite graph M을 샘플링합니다. 그 이후에는 샘플링된 그래프에 edge 가 존재하는 경우에만 그에 해당하는 attention score 를 계산할 수 있습니다.
그림 3. SBM-Transformer의 forward step[3]
그래프 샘플링 단계에서는 주어진 input에 따라 parameterize되는 SBM을 사용합니다. 한 SBM이 정의되기 위해선 Query-source node의 cluster membership , key-destination node의 cluster membership , cluster 간 connection probabilities 까지 총 세 개의 non-negative한 parameter가 필요합니다. 의 경우, 각 cluster를 표현하는 학습 가능한 embedding C를 자기 자신과 inner product를 취해 cluster 간 연결성을 계산합니다. 각 node의 cluster membership을 의미하는 와 의 경우, query와 key representation을 MLP에 통과시켜 token representation을 node representation space로 옮긴 후, cluster embedding C와 함께 inner product를 계산하고, cluster membership을 얻습니다. 그림 4에서 input query/key와 cluster embedding을 통해 SBM을 parameterize하는 상세한 과정을 확인할 수 있습니다.
그림 4. Query-key와 cluster embedding로부터 각 SBM을 parameterize하는 상세 과정
SBM에 필요한 세 가지 parameter가 모두 준비되었다면 그림 5와 같이 fastRG[4] 알고리즘을 실행하여 그래프를 샘플링할 수 있습니다. 전반적인 단계는 다음과 같습니다. 1) 그림 5의 1~3과 같이 node membership이 node의 집합 위에 정의된 확률 분포가 되게끔 normalization을 거친 후, 2) 4와 같이 Poisson 분포에서 생성할 edge의 개수를 샘플링하고, 3) 6~11과 같이 각 edge에 대해 cluster pair를 샘플링, 해당하는 node 확률 분포에서 source와 destination node를 샘플링합니다. 전체 프로세스를 수행하는 데 드는 비용은 edge 개수에 비례하며, 후반부 for-loop의 경우 각 edge가 독립적으로 샘플링되기 때문에 병렬적으로 실행해 프로세스를 한층 가속할 수 있습니다.
그림 5. SBM으로부터 샘플링을 할 때 사용되는 fastRG 알고리즘[3]
그래프 샘플링 단계는 단순하고 효율적이지만 본래 discrete합니다. 그 때문에 통상적인 backpropagation은 SBM의 적절한 parameterization을 학습할 수 없습니다. 이러한 non-differentiability에 대처하기 위해 Straight-Through Estimator(STE)를 활용해 gradient가 그래프 샘플링 단계를 지나 edge probability로 직접 전달되게 합니다. 이런 방법으로 샘플링됐던 edge가 예측이 도움이 되었는지 안되었는지에 대한 정보를 제공해 end-to-end로 샘플링 과정 학습이 가능하게 해줍니다. 결과적으로, SBM-Transformer의 backward step 또한 모델이 주어진 데이터에 맞춰 선택한 edge의 개수에 linear한 비용을 사용합니다. 그림 6에서 전체 파이프라인을 확인할 수 있습니다.
그림 6. SBM-Transformer의 attention 모듈[3]
그림 7. SBM으로 표현할 수 있는 attention 패턴의 예시[3]
놀랍게도 그림 7에서 볼 수 있듯 SBM이 latent feature space에서 node와 cluster representation을 기반으로 해 매우 다양한 attention 마스크를 표현할 수 있다는 점을 발견했습니다. Attention 마스크의 전반적인 density는 embedding이 얼마나 밀집해 있는가에 따라 no attention부터 full attention까지 넓은 범위를 표현할 수 있습니다. 이러한 SBM의 유연함 덕분에 SBM-Transformer는 latent graph space에서 low-rank 구조를 가정하고 있음에도 불구하고 일반 Transformer와 동일한 universal approximability를 유지한다는 것을 증명할 수 있었습니다.
Experiments
정량 평가를 위해 우선 LRA 벤치마크[5]에서 본 논문에서 제안한 SBM-Transformer 대비 이전에 발표되었던 Transformer의 변형 모델과 기존 트랜스포머를 비교하는 실험을 진행했습니다. 표 1의 결과를 보면 SBM-Transformer 모델이 다른 효율적 Transformer 변형 모델에 비해 좋은 성능을 보이며, 기존 Transformer와 비교해서도 훨씬 적은 attention을 사용했음에도 나은 성능을 보여주고 있습니다. 표 2에서는 SBM-Transformer가 inference 중 FLOP 수와 최대 메모리 사용량을 대폭 감소시킨다는 것을 볼 수 있습니다.
표 1. LRA 벤치마크에서의 정확도 결과. SBM-T 결과에서 λ는 샘플링 된 각 edge에 페널티를 부여하는 attention density regularizer를 나타내며 소괄호 안의 백분율은 test time 동안의 attention density를 나타냅니다[3]
표 2. LRA test time 동안의 example 당 평균 FLOP 수와 최대 메모리 사용량의 비교[3]
Downstream NLP 상황에서의 모델 성능을 평가하기 위해 GLUE (General Language Understanding Evaluation) 벤치마크[6]에서도 실험을 진행했습니다. 표 3은 각 task에서의 classification 정확도를 보여줍니다. 이를 통해 본 논문이 제안한 모델이 기존에 개발되었던 baseline 대비 경쟁적인 성능을 보이는 것을 알 수 있습니다.
표 3. GLUE 벤치마크에서의 정확도 결과[3]
정성 분석을 위해 LRA 내 image와 pathfinder task에서의 input example 별 attention heatmap을 시각화하여 어떠한 example이 올바른 예측을 하기 위해 dense 혹은 sparse한 attention을 필요로 하는지 확인했습니다. 두 흰색 점이 점선으로 연결되어 있는지 확인하는 task인 LRA Pathfinder에서 흥미로운 점을 발견했는데, SBM-Transformer가 attention을 분포하는 방식이 사람과 동일하게 작동한다는 것이었습니다: 육안으로 확인이 어려운 경우 dense attention을 사용하고, 사람이 보기에도 상대적으로 쉬운 경우에는 훨씬 sparse한 attention을 사용하는 것을 볼 수 있었습니다. CIFAR-10과 같은 LRA Image task에서도 동일한 분석을 진행했는데, 모델이 주로 attention을 분배할 때, 올바른 분류를 위해 이미지 내 개체와 배경 사이 대비의 큰 변화를 인식하는데 주로 attention을 사용한다는 점을 발견했습니다.
그림 8. LRA Image(좌)와 LRA Pathfinder(우)에서의 attention density heatmap 예시[3]
Conclusion
본 논문에서는 널리 사용되는 Transformer 아키텍쳐의 효율성을 개선한 SBM-Transformer를 제시했습니다. 이 모델은 attention score matrix 전체를 계산하지 않고도 주어진 데이터에 맞춰 sparse 또는 dense attention 사이를 선택해 불필요한 비용을 아낄 수 있습니다. 본 연구진은 SBM의 유연성을 통해 low-rank latent graph structure를 가정하고, SBM-Transformer가 기존 Transformer의 표현력을 유지할 수 있음을 증명했으며, LRA와 GLUE 벤치마크에 대한 실험으로 기존 Transformer뿐만 아니라 다른 baseline 대비 경쟁력 있는 성능을 보였다는 것을 입증했습니다. 추후 특정 구조를 가정하지 않는 fine-grained한 sparsity 아래 GPU 친화적인 tensor operation을 통합하고, inference 중 SBM-Transformer에 의해 유도된 query-key 상호작용의 기하학적 구조의 연구를 진행하고자 합니다.
▶Transformers meet Stochastic Blockmodels: Attention with Data-Adaptive Sparsity and Cost (Link)