[논문 리뷰] Magic Pyramid: Accelerating Inference with Early Exiting and Token Pruning
매직 피라미드(MP)는 토큰 프루닝과 일찍 종료하는 기법을 유기적으로 융합하여 넓이와 깊이 모두에서 계산을 감소시킴으로써 BERT 추론을 가속화한다. 최대 11.95배의 속도 향상을 이룩했으며, 정확도 저하가 0.5% 미만이어서 최신 기법들보다 최대 2.13배 빠른 성능을 보이며 다양한 입력 길이에서 견고한 성능을 유지한다.
Pre-training and then fine-tuning large language models is commonly used to achieve state-of-the-art performance in natural language processing (NLP) tasks. However, most pre-trained models suffer from low inference speed. Deploying such large models to applications with latency constraints is challenging. In this work, we focus on accelerating the inference via conditional computations. To achieve this, we propose a novel idea, Magic Pyramid (MP), to reduce both width-wise and depth-wise computation via token pruning and early exiting for Transformer-based models, particularly BERT. The former manages to save the computation via removing non-salient tokens, while the latter can fulfill the computation reduction by terminating the inference early before reaching the final layer, if the exiting condition is met. Our empirical studies demonstrate that compared to previous state of arts, MP is not only able to achieve a speed-adjustable inference but also to surpass token pruning and early exiting by reducing up to 70% giga floating point operations (GFLOPs) with less than 0.5% accuracy drop. Token pruning and early exiting express distinctive preferences to sequences with different lengths. However, MP is capable of achieving an average of 8.06x speedup on two popular text classification tasks, regardless of the sizes of the inputs.
연구 동기 및 목표
- BERT와 같은 대규모 미사전학습 모델에서 높은 추론 지연이 발생하는 문제를 해결하기 위해 실시간 생산 환경의 제약 조건을 고려한다.
- 기존 기법의 한계를 극복하기 위해, 긴 시퀀스에서는 효과적인 토큰 프루닝과 짧은 시퀀스에서는 효과적인 일찍 종료 기법이 각각 시퀀스 길이 스펙트럼의 반대편에서 성능이 떨어지는 문제를 해결한다.
- 변동하는 입력 길이에 걸쳐 일관되고 속도 조절이 가능한 추론을 위한 통합 프레임워크를 개발한다.
- 특히 저지연, 고처리량 배포 환경에서 높은 정확도를 유지하면서도 상당한 계산량을 감소시킨다.
제안 방법
- MP는 토큰 프루닝과 일찍 종료 기법을 하나의 계층적 추론 파이프라인에 통합하여, 프루닝은 너비 방향 계산을 감소시키고, 일찍 종료는 깊이 방향 계산을 감소시킨다.
- 각 레이어에서 주의점(attention score)의 불확실성에 기반해 비필수적인 토큰을 식별하고 제거하기 위해 학습 가능한 임계값 메커니즘을 사용한다.
- 각 트랜스포머 블록에 하위 분류기가 부착되어 있으며, 신뢰도가 동적 임계값 τ를 초과하면 예측을 조기에 종료할 수 있다. 이 τ 값은 학습 중에 조정된다.
- 프루닝과 일찍 종료를 공동 최적화할 때 정확도와 속도 간의 상충 관계를 균형 있게 조절하기 위해 온도 스케일링과 손실 가중치(λ)를 활용한다.
- 학습 과정에서는 정확도, GFLOPs, 일찍 종료 확률을 함께 최적화하며, 지식 증착과 불확실성 정규화를 포함한 다중 과제 손실 함수를 사용한다.
- 프레임워크는 BERT 기반 모델에 종단 간(end-to-end)으로 적용되며, 입력 길이에 따라 토큰 프루닝과 일찍 종료 조건을 동적으로 적용하여 추론을 조정한다.
실험 결과
연구 질문
- RQ1다양한 입력 길이에서 토큰 프루닝과 일찍 종료 기법을 효과적으로 융합하여 일관된 추론 가속화를 달성할 수 있는가?
- RQ2너비 방향(프루닝)과 깊이 방향(일찍 종료)의 계산 감소 기법 간의 상호보완적 상호작용이 개별 기법보다 뛰어난 속도 향상을 이끌 수 있는가?
- RQ3통합 기법은 상당한 계산 감소를 이룰 수 있으며, 특히 저지연 생산 환경에서 높은 정확도를 유지할 수 있는가?
- RQ4제안된 방법의 성능는 다양한 시퀀스 길이와 NLP 작업에서 어떻게 변화하는가?
주요 결과
- MP는 AG News와 Yelp 데이터셋에서 최대 11.95배의 속도 향상을 기록했으며, 모든 시퀀스 길이 그룹에서 FastBERT(일찍 종료)와 LTP(토큰 프루닝)를 뛰어넘었다.
- 두 텍스트 분류 작업 평균으로 MP는 입력 길이에 관계없이 평균 8.06배의 속도 향상을 기록하여 다양한 입력에서 일관된 성능을 입증했다.
- BERT 대비 GFLOPs를 최대 70% 감소시켰으며, 정확도 저하가 0.5% 미만이었고, LTP와 FastBERT를 모두 효율-정확도 트레이드오프 측면에서 능가했다.
- 긴 시퀀스(70 토큰 초과)에서는 Yelp에서 8.25배, AG News에서 11.95배의 속도 향상을 기록했으며, 이는 각각 FastBERT의 6.18배와 8.84배를 뛰어넘었다.
- τ = 0.8일 때, MP는 AG News에서 11.95배, Yelp에서 10.10배의 속도 향상을 기록하여 강력한 확장성을 입증했다.
- MP는 BERT와 비교해 정확도를 유지하거나 약간 향상시켰으며(예: AG News에서 94.3%), GFLOPs를 1.8로 감소시켜 4.95배의 속도 향상을 달성했고, 이는 FastBERT의 2.3 GFLOPs(3.97배 속도 향상)를 뛰어넘었다.
더 나은 연구,지금 바로 시작하세요
논문 읽기부터 검토까지, 연구 시간을 획기적으로 줄여보세요.
카드 등록 없음 · 무료 플랜 제공
이 리뷰는 AI가 만들고, 인간 에디터가 검토했습니다.