Grouped Query Attention (GQA)
Grouped Query Attention(GQA)가 KV 캐시 메모리를 줄이고, 추론 효율성을 향상시키며, Transformer 모델의 성능을 균형 있게 유지하는 방법을 알아보세요.
Grouped Query Attention (GQA)는 여러 쿼리 헤드가 더 적은 수의 키 및 값 헤드를 공유하는 어텐션 설계 방식입니다. 이는 키와 값에 필요한 메모리 및 데이터 이동량을 줄이면서 다중 헤드 쿼리의 다양한 관점을 유지합니다. GQA는 특히 모델이 한 번에 하나의 토큰을 생성하고 저장된 어텐션 상태를 반복적으로 읽어오는 자가 회귀(autoregressive) 추론 과정에서 매우 유용합니다.
GQA 동작 방식#
attention mechanism 내에서 각 토큰은 세 가지 표현으로 투영됩니다:
- Query: 현재 토큰이 찾고자 하는 정보를 설명합니다.
- Key: 사용 가능한 각 토큰이 무엇을 나타내는지 설명합니다.
- Value: 쿼리가 키와 일치할 때 검색되는 정보를 포함합니다.
기존의 self-attention에서는 모델이 서로 다른 관계를 병렬로 학습할 수 있도록 이러한 투영을 여러 헤드로 나눕니다. GQA는 많은 쿼리 헤드를 유지하되 그 그룹들을 공유되는 키 및 값 헤드에 할당합니다. 예를 들어, 어텐션 레이어는 8개의 쿼리 헤드와 2개의 키-값 헤드를 사용할 수 있습니다. 이 경우 각 키-값 헤드가 4개의 쿼리 헤드를 담당하게 됩니다.
어텐션 연산 자체는 스케일드 닷 프로덕트 어텐션(scaled dot-product attention)으로 유지됩니다. 변경되는 부분은 헤드 배치와 생성, 저장 및 읽기되는 키-값 데이터의 양입니다. 따라서 PyTorch scaled dot-product attention 같은 프레임워크에서는 GQA가 활성화될 때 쿼리 헤드 수가 키-값 헤드 수로 나누어떨어져야 합니다. (docs.pytorch.org)
GQA와 Multi-Head 및 Multi-Query Attention의 비교#
GQA는 멀티 헤드 어텐션(multi-head attention)과 멀티 쿼리 어텐션(multi-query attention) 사이에 위치합니다:
- Multi-head attention: 모든 쿼리 헤드가 자체적인 키 및 값 헤드를 가집니다. 이는 최대의 헤드 독립성을 제공하지만 가장 큰 키-값 메모리 요구 사항을 발생시킵니다.
- Grouped query attention: 여러 쿼리 헤드가 각 키-값 헤드를 공유합니다. 이는 표현 용량과 메모리 효율성의 균형을 맞춰줍니다.
- Multi-query attention: 모든 쿼리 헤드가 하나의 키 헤드와 하나의 값 헤드를 공유합니다. 이는 키-값 스토리지를 최소화하지만 키-값의 다양성을 떨어뜨립니다.
레이어에 32개의 쿼리 헤드가 있는 경우, 멀티 헤드 어텐션은 32개의 키-값 헤드를 사용할 수 있고, GQA는 8개를 사용할 수 있으며, 멀티 쿼리 어텐션은 1개를 사용합니다. GQA는 더 넓은 Transformer architecture를 대체하는 것이 아니며, Transformer 어텐션 레이어 내의 가능한 구성 중 하나입니다. NVIDIA의 multi-head, multi-query, and grouped-query attention documentation에서는 이러한 변형들이 주로 쿼리 헤드를 서비스하는 키-값 헤드의 수에서 차이가 난다고 설명하고 있습니다. (nvidia.github.io)
GQA가 중요한 이유#
자가 회귀 생성 과정에서 모델은 이전 키와 값을 KV cache에 저장합니다. 표준 멀티 헤드 어텐션의 경우 모든 어텐션 헤드가 별도의 캐시된 키와 값을 생성합니다. 따라서 긴 시퀀스, 큰 배치, 그리고 많은 모델 레이어는 상당한 GPU 메모리를 소비할 수 있습니다.
GQA는 더 적은 수의 키-값 헤드를 사용하므로 캐시 크기가 해당 헤드의 감소 비율에 비례하여 줄어듭니다. 이는 다음을 향상시킬 수 있습니다:
- Memory capacity: 더 긴 프롬프트나 더 많은 동시 요청을 가용한 메모리에 수용할 수 있습니다.
- Generation throughput: 생성되는 각 토큰마다 읽어야 하는 키-값 데이터가 줄어듭니다.
- Inference latency: 메모리 트래픽 감소로 인해 토큰 간 처리 시간이 단축될 수 있습니다.
- Context window scalability: 긴 시퀀스를 처리하는 서비스 제공이 더욱 실용적여집니다.
이러한 성능 향상은 하드웨어, 시퀀스 길이, 배치 크기 및 커널 지원 여부에 따라 달라집니다. 프로덕션 런타임에는 여전히 효율적인 캐시 할당이 필요하며, 예를 들어 NVIDIA의 TensorRT-LLM KV cache system은 GQA 지원을 블록 기반 캐싱, 재사용 및 오프로딩과 결합합니다. (nvidia.github.io)
실제 애플리케이션 사례#
Long-context assistants: 코딩 어시스턴트는 답변을 생성하기 전에 수천 개의 소스 코드 토큰을 처리할 수 있습니다. GQA는 해당 기록과 연관된 캐시된 키-값 상태를 줄여주어, 서버가 GPU 메모리를 비례해서 늘리지 않고도 더 긴 파일이나 더 많은 동시 사용자를 처리할 수 있도록 해줍니다.
Multimodal image and video assistants: vision-language model은 이미지 패치나 비디오 프레임을 긴 토큰 시퀀스로 표현할 수 있습니다. GQA는 이러한 시각적 토큰이 언어 디코더에 입력된 후 어텐션 캐시의 압박을 줄여주며, 잠재적으로 추가적인 이미지, 프레임 또는 요청을 위한 메모리를 더 확보할 수 있게 해줍니다. 이는 어텐션이 자가 회귀 KV 캐싱 없이 모든 이미지 패치를 병렬로 처리할 수 있는 표준 Vision Transformer와는 다릅니다.
실용적인 구현 및 트레이드오프#
PyTorch는 enable_gqa=True을 통해 GQA를 지원합니다. 문서화된 이 CUDA 예제는 8개의 쿼리 헤드와 2개의 공유 키-값 헤드를 사용합니다:
import torch
import torch.nn.functional as F
from torch.nn.attention import SDPBackend, sdpa_kernel
assert torch.cuda.is_available(), "A CUDA device is required."
query = torch.randn(1, 8, 16, 32, device="cuda")
key = torch.randn(1, 2, 16, 32, device="cuda")
value = torch.randn(1, 2, 16, 32, device="cuda")
# Four query heads share each key-value head.
with sdpa_kernel(SDPBackend.MATH):
output = F.scaled_dot_product_attention(
query,
key,
value,
enable_gqa=True,
)
print(output.shape)출력은 8개의 쿼리 헤드를 유지하는 반면, 키와 값은 2개의 헤드만 사용합니다. PyTorch attention backend selector를 통해 지원되는 어떤 커널이 연산을 수행할지 제어할 수 있으며, PyTorch Transformer building-block guide는 더 광범위한 구현 맥락을 제공합니다. (docs.pytorch.org)
GQA는 모델을 설계하거나 조정할 때 선택해야 하며, 일반적으로 기존 체크포인트에 대한 추론 시점의 스위치로 작동하지 않습니다. 개발자는 프레임워크 지원, 헤드 가분성, 수치적 거동, 메모리 사용량 및 태스크 품질을 검증해야 합니다. 컴퓨터 비전 배포의 경우 Ultralytics TensorRT integration 및 Ultralytics benchmark mode와 같은 상호 보완적인 워크플로가 아키텍처 및 런타임 최적화가 타겟 하드웨어에서 의미 있는 성능 향상을 가져다주는지 측정하는 데 도움이 됩니다.






