SDPA 실전 — 고성능 Transformer를 직접 만들기

SDPA 실전 — 고성능 Transformer를 직접 만들기

이 튜토리얼은 scaled_dot_product_attention 을 직접 써서 Transformer를 구현하는 방법을 보여줘요. 이 함수는 이미 nn.MultiheadAttentionnn.TransformerEncoderLayer 에 통합되어 있어요.

왜 SDPA인가

순수 PyTorch 연산으로도 어텐션을 쓸 수 있지만, 융합 구현(fused implementation) 은 커널 내부에서 전체를 처리해 naive 구현 대비 큰 성능 이득을 줘요. 특히 FlashAttention 커널은 시퀀스 길이가 길수록 압도적이에요.

실습 순서

  1. 하이퍼파라미터(배치·시퀀스·차원·헤드 수)를 정의.
  2. 쿼리·키·값을 만들고 F.scaled_dot_product_attention(...) 호출.
  3. sdpa_kernel 컨텍스트로 강제할 커널(FlashAttention 등) 지정.
  4. torch.compile 과 호환되는 CausalSelfAttention 모듈 구성 — NestedTensor까지 지원.

벤치마크로 체감

같은 입력으로 기본 구현 vs 융합 구현의 시간을 프로파일로 찍어 차이를 직접 확인할 수 있어요. 어텐션 바이어스(ALiBi 류)는 torch.compile 과도 함께 동작해요.

더 알아보기