섹션별 상세
FlexAttention은 Python으로 score_mod나 mask_mod 함수를 작성하면 컴파일러가 이를 최적화된 커널로 변환해주는 API로, 이번에 FlashAttention-4 백엔드가 통합되었다.
python
import torch
from functools import partial
from torch.nn.attention.flex_attention import flex_attention
flex_flash = torch.compile(
partial(flex_attention, kernel_options={"BACKEND": "FLASH"}),
dynamic=False
)
def local_boost(score, b_idx, h_idx, q_idx, kv_idx):
return torch.where(torch.abs(q_idx - kv_idx) < 128, score * 1.1, score)
# 실행
out = flex_flash(q, k, v, score_mod=local_boost)FlexAttention에서 FlashAttention-4 백엔드를 사용하여 커스텀 score_mod를 적용하는 예시 코드

NVIDIA의 CuTeDSL을 기반으로 하여, 사용자가 작성한 Python 로직을 TensorSSA 표현식으로 재작성하고 이를 FA4의 비동기 파이프라인에 인라인(inline) 방식으로 삽입한다.
Blackwell(GB200) GPU에서 Triton 대비 Forward 패스는 1.6~3.2배, Backward 패스는 1.85~2.3배의 속도 향상을 기록했으며, 일부 케이스에서는 cuDNN의 성능에 근접하거나 능가한다.


블록 희소(Block-sparse) 반복 기능을 확장하여 커널이 마스크된 빈 블록을 건너뛰도록 설계되었으며, Blackwell의 Cluster Launch Control(CLC) 기능을 통해 동적 작업 스케줄링의 이점을 누린다.
현재 블록 크기 제한(Hopper 128x128, Blackwell 256x128)과 동적 스칼라 값 변경 시 재컴파일이 필요한 점, 그리고 학습 가능한 바이어스 텐서의 그래디언트 미지원 등 일부 제약 사항이 존재한다.
python
def tanh_softcap(score, b, h, q_idx, kv_idx):
return soft_cap * tanh(score / soft_cap)동적 스칼라 값(soft_cap)을 사용하는 예시로, 현재 백엔드에서는 값이 바뀔 때마다 재컴파일이 발생하는 제약이 있음
용어 해설
- 플렉스 어텐션(FlexAttention)
- — PyTorch API로, 사용자가 Python 함수(score_mod, mask_mod)만으로 복잡한 어텐션 메커니즘을 정의하면 이를 고성능 커널로 자동 컴파일해주는 기술이다. CUDA 코드를 직접 작성하지 않고도 ALiBi, 슬라이딩 윈도우 등 다양한 변형을 효율적으로 구현할 수 있게 돕는다.
- 플래시 어텐션 4(FlashAttention-4)
- — 최신 GPU 아키텍처에 최적화된 어텐션 알고리즘의 4번째 버전으로, 메모리 접근을 최소화하고 연산 유닛의 활용도를 극대화한다. 특히 NVIDIA Blackwell 아키텍처의 하드웨어 기능을 활용하여 이전 버전보다 높은 처리량을 제공한다.
- 큐트 DSL(CuTeDSL)
- — NVIDIA가 제공하는 Python 기반의 도메인 특화 언어(DSL)로, 복잡한 CUDA 커널을 추상화된 텐서 연산 단위로 작성하고 최적화할 수 있게 한다. FlexAttention의 Python 로직을 고성능 하드웨어 명령어로 변환하는 핵심 가교 역할을 한다.
- 블록 희소(Block-sparse)
- — 전체 행렬 대신 유의미한 값이 있는 특정 블록들만 연산하여 계산량과 메모리 사용량을 획기적으로 줄이는 기법이다. 데이터에 따른 동적 마스킹을 지원하며, 불필요한 연산을 건너뛰어 성능을 최적화한다.
- 적시 인스턴스화(JIT Instantiation)
- — 실행 시점에 필요한 파라미터나 로직에 맞춰 최적화된 기계어 코드를 즉석에서 생성하고 실행하는 방식이다. FlexAttention에서는 사용자의 Python 함수를 기반으로 최적화된 FlashAttention-4 커널을 실시간으로 생성한다.
기술
- PyTorch
- FlashAttention-4
- CuTeDSL
- Triton
- Blackwell GPU
- Hopper GPU
활용 사례
- Custom Attention (ALiBi, Sliding Window)
- Long Context LLM Training
- Sparse Attention Implementation
언급된 리소스
AI 분석 전체 내용 보기
AI 요약 · 북마크 · 개인 피드 설정 — 무료
출처 · 인용 안내
원문 발행 2026. 03. 06.수집 2026. 03. 06.출처 타입 RSS
인용 시 "요약 출처: AI Trends (aitrends.kr)"를 표기하고, 사실 확인은 원문 보기 기준으로 진행해 주세요. 자세한 기준은 운영 정책을 참고해 주세요.
