Lighthouse Attention 기법
선택형 계층적 attention으로 512K 컨텍스트에서 표준 attention보다 약 17배 빨랐다.
긴 컨텍스트 pretraining은 attention의 O(N^2) 비용이 병목이다. Lighthouse Attention은 선택형 계층적 attention으로 Q, K, V를 같은 비율로 대칭 pooling해 L-level pyramid를 만들고, 각 엔트리를 L2 norm으로 점수화한 뒤 top-K만 남긴다. learned scorer head, Gumbel-softmax, STE, auxiliary loss 없이 선택된 토큰만 contiguous dense sub-sequence로 묶어 FlashAttention을 그대로 돌리므로, training과 inference에서 같은 kernel을 쓰고 upstream FlashAttention 개선도 그대로 받는다. 구현은 upstream torchtitan 위에 두 파일과 약 600줄의 수정으로 얹혔다.
핵심은 세 가지다.
- 선택 로직은 kernel 밖에서 처리한다.
torch.gather와torch.sort뒤에 ordinary FlashAttention을 붙여 dense attention 경로와 맞춘다. - 대칭 pooling으로 pooled query와 pooled key가 같은 표현 공간을 공유한다.
- 정렬된 서브시퀀스는 구멍 없는 causal sequence가 되어 표준 lower-triangular mask를 그대로 쓸 수 있다.
학습은 2단계다. 먼저 Lighthouse로 대부분의 스텝을 학습하고, 같은 optimizer state와 dataloader continuation으로 SDPA-resume를 짧게 돌려 dense attention 복원 가능성을 검증한다. 10k+6k, 11k+5k, 12k+4k의 세 split 모두에서 재개 직후 loss가 1.12~1.57 nats 튀지만 1k~1.5k SDPA steps 안에 회복했고, 최종 loss는 dense-from-scratch baseline 0.7237보다 낮은 0.6980~0.7102였다. 50B tokens / 16k steps 기준으로 75~106 B200-hours를 절감했고, sparse training이 dense attention 능력을 망가뜨리지 않는다는 점을 확인했다.
성능 차이도 뚜렷하다. 단일 B200에서 512K context의 forward+backward는 표준 attention보다 약 17× 빨랐고, 98K context pretraining은 end-to-end로 1.4~1.7× 빨라졌다. 530M Llama-3, 16k optimiser steps, 8×B200 single-node 실험에서 stage 1은 84~126k tokens/s/GPU를 냈고, dense SDPA는 약 46k였다. 100K를 넘는 구간에서는 530M 아키텍처가 single B200에서 OOM이므로 context parallelism(CP) 으로 확장했고, 1M-token training은 32 B200s에서 수행했다. scatter-back에는 재현용 deterministic int-atomic 커널도 있지만 1.2~2× 느리며, 기본값은 fp-atomic이다.
이 요약은 원문 이해를 돕기 위한 큐레이션입니다. 저작권은 원저작자에게 있으며, 정확한 내용과 맥락은 원문을 확인하세요.
요약 오류, 출처 표기 문제, 삭제 요청은 문의 · 건의로 알려주세요.