
MSA는 추론 시점에 적용하는 단순 후처리 최적화가 아닌, 모델 학습 단계부터 반영된 어텐션 구조입니다.
기존 GQA 레이어 위에 학습된 블록 검색기를 결합하여, 전체 KV 캐시를 스캔하는 대신 중요 KV 블록만 선별해 연산합니다.
Query hidden state
├─ 목차 담당: 블록별 중요도 계산 및 Top-k 블록 선택
└─ 실제 읽기 담당: 선택된 블록의 K/V만 로드하여 정확한 어텐션 수행
Index Branch: 가벼운 별도 프로젝션을 통해 Query와 Indexer Key를 생성합니다.
토큰 단위가 아닌 KV 블록 단위의 중요도 점수를 산출하며, GQA 그룹별로 독립적인 Top-k 블록을 선택합니다.
Main Branch: Index Branch가 선택한 블록의 K/V를 가져와 Scaled Dot-Product Attention 및 Softmax 연산을 수행합니다.
선택된 블록 내부 연산은 근사치가 아닌 정확한 계산입니다.
긴 컨텍스트 생성 시 병목은 연산량이 아닌 KV 캐시 메모리 읽기 대역폭에서 발생합니다.
MSA는 전 구간 스캔용 '미리보기 키(토큰당 256B)'와 선택된 본편 K/V 블록만 로드하는 방식으로 대역폭을 절감합니다.
| 컨텍스트 길이 | Dense Attention 로드량 | MSA 로드량 | 대역폭 절감 비율 |
|---|---|---|---|
| 260K | ~60GB | 미리보기 ~4GB + 선택 블록 ~2GB = ~6GB | 1/10 수준 |
| 1M | ~240GB | 미리보기 ~15GB + 선택 블록 ~2GB = ~17GB | 1/14 수준 |
가중치 로드량을 포함한 1M 컨텍스트 전체 읽기량은 Dense 기준 스텝당 263GB(초당 ~11토큰)에서 MSA 기준 40GB 초당 ~70토큰 로 감소합니다.
MSA 구조에서는 컨텍스트 길이가 늘어나더라도 증가하는 연산은 미리보기 스캔뿐이며, 본편 K/V 로드량은 고정됩니다.
추가되는 미리보기 캐시용 용량 메모리는 전체 본편 캐시의 약 6% 수준입니다.
학습 방식: Main Branch의 실제 어텐션 분포를 교사 신호로 활용합니다.
GQA 그룹 단위로 어텐션 분포를 수집한 뒤, Index Branch의 블록 점수 분포와 비교하여 KL Divergence Loss로 학습시킵니다.
Gradient는 두 Branch 간 분리되어 Main Branch의 가중치 오염을 방지합니다.
GPU 커널 최적화: 불규칙한 블록 접근으로 인한 메모리 효율 저하를 막기 위해, Query 기준 연산을 KV 블록 기준으로 역전시켰습니다.
[Query-centric]
Query 1 -> 블록 3, 8, 20
Query 2 -> 블록 3, 9, 20
(블록 3을 중복 로드)
[KV-outer]
블록 3 -> 대상 Query 1, 2 모아서 처리
블록 8 -> 대상 Query 1, 3 처리
(메모리 연속 읽기 및 SRAM 적재 블록 재활용)
kv_unified 제약 사항llama.cpp의 MiniMax M3 구현은 별도의 Indexer KV 캐시를 운용하며 아래 순서로 동작합니다.
kv_unified 분기 조건 및 Dense Fallback 이슈현재 llama.cpp 내부의 MSA 활성화 조건식은 다음과 같습니다.
flash_attention == true && (n_seq_max == 1 || kv_unified == false)
문제 원인: 현재 MSA 구현은 논리적 스트림 내 절대 토큰 위치와 물리적 KV 캐시 셀 위치를 1:1로 매핑합니다.
Fallback 동작: kv_unified=true 상태에서 다중 시퀀스가 입력될 경우, 캐시 셀이 섞여 "시퀀스 상의 블록 위치 = 공유 캐시의 연속 셀 블록" 조건이 깨집니다.
이로 인해 모델은 sparse 연산을 포기하고 Dense 어텐션으로 Fallback합니다.
Dense Fallback의 영향:
수치 차이의 원인: 논문 및 공식 발표 자료 간 성능 배율 차이 Prefill 9~14배, Decode 7.5~15배 는 기준 모델 규모 109B vs 428B/23B MoE , 하드웨어, 컨텍스트 길이 차이에서 비롯됩니다.
멀티 세션 서빙 환경의 해결 과제:
특정 CUDA 커널에 의존하기보다 컴퓨팅 백엔드 CUDA, Vulkan, Metal 등 에 유연하게 대응하는 프레임워크 수준의 그래프 처리가 필요합니다.
Batching 효율성을 위한 kv_unified=true 사용 시 MSA 연산이 유지되도록, "다중 시퀀스 환경에서의 [시퀀스별 Position Cell Block] 명시적 매핑 레이어" 구축이 핵심 개선 항목입니다.