scaled_dot_product_attention — 융합 커널로 더 빠르게

scaled_dot_product_attention — 융합 커널로 더 빠르게

torch.nn.functional.scaled_dot_product_attention(query, key, value, ...) 은 쿼리·키·값 텐서 사이의 스케일된 점곱 어텐션을 계산하는 함수예요. 이름 그대로 '스케일링된 점곱'을 어텐션 가중치로 쓰죠.

동작 원리

쿼리와 키의 점곱을 차원의 제곱근으로 나눠 스케일링하고, 소프트맥스를 거쳐 값과 곱해요. 인과 마스크(is_causal=True)를 주면 미래 토큰을 가리는 디코더용 마스크가 자동으로 만들어져요.

import torch.nn.functional as F
y = F.scaled_dot_product_attention(query, key, value, is_causal=True)

세 가지 구현

이 함수는 CUDA 상에서 세 구현 중 하나로 분기돼요.

  • FlashAttention-2: IO 인식형 고속·저메모리 정확 어텐션.
  • Memory-Efficient Attention: 메모리 절약형.
  • PyTorch 자체 C++ 구현: 위 수식과 동일.

커널 선택은 sdpa_kernel 컨텍스트 매니저로 강제할 수 있어요.

마스크 의미 주의

이 함수에서는 True 가 '참여할 값'을 뜻해요. MultiheadAttentionkey_padding_mask 와는 반대이므로, 마이그레이션 시 ~maskmask.logical_not() 으로 반전해야 해요.

더 알아보기