Relaxed-System-Lab/Flash-Sparse-Attention

🚀🚀 Efficient implementations of Native Sparse Attention

Flash‑Sparse‑Attention (FSA)

소개Native Sparse Attention (NSA)의 Triton 기반 오픈소스 구현으로, 커널 루프를 재구성하여 메모리 트래픽과 연산 오버헤드를 획기적으로 줄입니다. 대규모 언어 모델(LLM)의 어텐션 레이어를 대체하는 드롭인 모듈로 제공되며, 최신 NVIDIA GPU(Ampere, Hopper 등)에서 작동합니다.

중요성 – 희소 어텐션(Sparse attention)은 매우 긴 시퀀스(수만 개의 토큰)에서 전체 어텐션의 이차 비용을 관리 가능한 수준으로 유지하는 방법입니다. 기존 NSA 커널은 group‑query‑attention (GQA) 헤드 그룹 크기가 작을 때(오늘날 LLM에서 흔한 경우) 패딩 및 원자적 가산 오버헤드로 인해 성능 저하가 발생합니다. FSA는 외부/내부 루프를 교체하고 작업을 3개의 특수 커널로 분할하여 이러한 병목 현상을 방지하며, 동일한 API를 사용하면서 학습(prefill) 및 추론 모두에서 최대 2‑3배의 속도 향상을 제공합니다.


주요 기능

기능 상세 내용
최적화된 Triton 커널 3개의 커널(메인 배치 query‑to‑KV, 리덕션, 온라인 softmax)을 통해 패딩 데이터 작업을 최소화하고 원자적 연산을 제거합니다.
GQA 인식 폴백 GQA 그룹 크기가 ≥ 8인 경우, 해당 영역에서 더 빠른 원래의 NSA 구현으로 자동 폴백합니다.
광범위한 하드웨어 지원 NVIDIA A100, H100, H200 (SXM 및 PCIe)에서 fp16 및 bf16으로 테스트되었습니다.
드롭인 API FlashSparseAttention은 표준 torch.nn.Module 인터페이스를 미러링하므로 cu_seqlens와 입력 텐서만 제공하면 됩니다.
학습 및 추론 전방 전용 프리필과 전체 역전파(loss‑backward) 모두 지원합니다.
호환성 PyTorch ≥ 2.4, Triton ≥ 3.0, HuggingFace transformers ≥ 4.45 및 공식 Flash‑Attention 2.6.3 커널 기반으로 구축되었습니다.

빠른 시작

# 1️⃣ 의존성 설치 (Python 3.10+ 권장)
pip install -r requirements.txt   # torch, triton, transformers, datasets, accelerate, flash‑attn 설치

# 2️⃣ 모듈 임포트 및 인스턴스화
import torch
from fsa.module.fsa import FlashSparseAttention, RopeConfig

fsa = FlashSparseAttention(
    hidden_size=4096,
    num_q_heads=4,
    num_kv_heads=4,
    head_dim=128,
    kernel_size=32,
    kernel_stride=16,
    block_size=64,
    topk=16,
    init_blocks=1,
    local_blocks=2,
    window_size=512,
    rope_config=RopeConfig(
        max_position_embeddings=131072,
        head_dim=128,
        rope_theta=500000,
        rope_scaling={
            "factor": 8.0,
            "high_freq_factor": 4.0,
            "low_freq_factor": 1.0,
            "original_max_position_embeddings": 8192,
            "rope_type": "llama3",
        },
    ),
).cuda().to(torch.bfloat16)

# 3️⃣ cu_seqlens (누적 길이) 및 입력 준비
seqlens = torch.LongTensor([65536, 32768]).int().cuda()
cu_seqlens = torch.cat([torch.zeros(1, dtype=torch.int32, device="cuda"),
                        torch.cumsum(seqlens, dim=0)], dim=0)

x = torch.randn(cu_seqlens[-1], 4096, device="cuda", dtype=torch.bfloat16)

# 4️⃣ 전방 + 후방 (학습) 예제
y = fsa(x, cu_seqlens)
loss = (y * torch.randn_like(y)).sum(-1).mean()
loss.backward()

일반 transformer와 비교했을 때 유일한 추가 단계는 가변 길이 배치를 인코딩하는 cu_seqlens를 구성하는 것입니다.


벤치마크

저장소에는 두 개의 스크립트가 포함되어 있습니다:

  • scripts/run_unit_test.sh – 전방/후방 정확성을 확인하고 원시 커널 지연 시간을 측정합니다.
  • scripts/run_unit_test_sel_attn.shselected‑attention 부분(주요 병목 구간)을 벤치마크합니다.

README에는 두 개의 성능 표가 있습니다:

  • 커널 수준 – FSA의 지연 시간을 1로 정규화했을 때, NSA 및 전체 Flash‑Attention은 블록 크기/top‑k에 따라 1.7–2.4배 더 느립니다.
  • 엔드 투 엔드 – LLaMA‑2‑70B와 같은 LLM의 경우, 학습 단계 지연 시간이 약 1.9초(NSA)에서 약 1.2초(FSA)로 감소하며 프리필 지연 시간도 비슷하게 개선됩니다.

사용 시기

  • 시퀀스 길이가 ≥ 32k 토큰인 LLM을 학습하거나 서비스하는 경우.
  • 모델이 GQA를 사용하며 KV 그룹당 헤드 수가 ≤ 8인 경우(가장 일반적인 구성).
  • NVIDIA Ampere/Hopper GPU를 보유하고 있고 Triton을 설치할 수 있는 경우.
  • 이미 밀집 어텐션에 Flash‑Attention을 사용 중이며 모델 코드를 다시 작성하지 않고 희소 어텐션 대안을 원하는 경우.

제한 사항 / 향후 계획

  • 현재 NVIDIA GPU만 지원하며 CPU 또는 AMD 경로는 없습니다.
  • 구현은 Q, K, V의 헤드 차원이 동일(≤ 256)하다고 가정합니다.
  • NSA와 FSA 사이를 동적으로 전환할 수 있는 "온라인 프로파일링" 모듈이 향후 릴리스(2025년 9월)에 예정되어 있습니다.

인용

논문에서 FSA를 사용하는 경우, 동봉된 arXiv 프리프린트를 인용하십시오:

@misc{yan2026fsaalternativeefficientimplementation,
  title={{FSA}: An Alternative Efficient Implementation of Native Sparse Attention Kernel},
  author={Ran Yan and Youhe Jiang and Zhuoming Chen and Haohui Mai and Beidi Chen and Binhang Yuan},
  year={2026},
  eprint={2508.18224},
  archivePrefix={arXiv},
  primaryClass={cs.DC},
  url={https://arxiv.org/abs/2508.18224},
}

요약

Flash‑Sparse‑Attention은 긴 컨텍스트를 가진 현대 LLM을 위해 native sparse attention을 실용적으로 만드는 고성능 Triton 커널 라이브러리입니다. 기존 PyTorch/transformers 코드에 드롭인 방식으로 적용 가능하며, Ampere/Hopper GPU에서 실행되고, 동일한 API와 수치적 정확성을 유지하면서 원래의 NSA 구현보다 2‑3배 빠른 속도를 제공합니다.

관련

  • 프로젝트
  • 프로젝트
  • 프로젝트
  • 프로젝트