본문으로 건너뛰기
PyTorch조회 1

PyTorch FlexAttention, FlashAttention-4 백엔드 도입으로 성능 최대 3.2배 향상

PyTorch의 FlexAttention API가 FlashAttention-4 백엔드를 지원하며 Hopper 및 Blackwell GPU에서 기존 Triton 대비 최대 3.2배의 성능 향상을 달성했습니다.

섹션별 상세

01
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를 적용하는 예시 코드

FlexAttention 출시 이후 월별 및 누적 프로젝트 채택 현황 그래프
Chart2024년 8월 출시 이후 FlexAttention을 사용하는 리포지토리가 꾸준히 증가하여 2025년 12월 기준 누적 1874개에 달함을 보여준다. 이는 커스텀 어텐션 구현에 대한 커뮤니티의 높은 수요를 입증한다.
02
NVIDIA의 CuTeDSL을 기반으로 하여, 사용자가 작성한 Python 로직을 TensorSSA 표현식으로 재작성하고 이를 FA4의 비동기 파이프라인에 인라인(inline) 방식으로 삽입한다.
Blackwell 아키텍처에서의 핑퐁 파이프라인 구조 다이어그램
DiagramBlackwell GPU에서 연산과 메모리 로드를 겹쳐서 처리하는 파이프라인 구조를 설명한다. FlexAttention의 FA4 백엔드가 이러한 하드웨어 특성을 어떻게 활용하여 성능을 극대화하는지 시각화한다.
03
Blackwell(GB200) GPU에서 Triton 대비 Forward 패스는 1.6~3.2배, Backward 패스는 1.85~2.3배의 속도 향상을 기록했으며, 일부 케이스에서는 cuDNN의 성능에 근접하거나 능가한다.
H100 GPU에서 FA3와 FlexAttention의 시퀀스 길이에 따른 성능 비교 차트
ChartForward 및 Backward 패스 모두에서 시퀀스 길이가 길어질수록 성능 차이가 발생하며, FlexAttention이 FA3의 성능을 어느 정도 추격하고 있음을 보여준다.
GB200 GPU에서 cuDNN과 FlexAttention Triton의 성능 비교 차트
ChartForward 패스에서 cuDNN이 FlexAttention Triton보다 약 2.07x에서 2.85x 더 빠른 성능을 보임을 나타낸다. 이는 새로운 FA4 백엔드 도입의 필요성을 뒷받침하는 벤치마크 데이터이다.
04
블록 희소(Block-sparse) 반복 기능을 확장하여 커널이 마스크된 빈 블록을 건너뛰도록 설계되었으며, Blackwell의 Cluster Launch Control(CLC) 기능을 통해 동적 작업 스케줄링의 이점을 누린다.
05
현재 블록 크기 제한(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)"를 표기하고, 사실 확인은 원문 보기 기준으로 진행해 주세요. 자세한 기준은 운영 정책을 참고해 주세요.