FlexAttention API — score_mod · mask_mod

FlexAttention API — score_mod · mask_mod

PyTorch 공식 API 문서는 FlexAttention의 두 확장점인 score_modmask_mod 를 설명해요. 어느 것을 언제 쓰는지가 성능을 결정하는 핵심이에요.

score_mod

  • 실제 점수값에 의존하는 수정(예: bias 추가, soft-capping)을 적용할 때 사용해요.
  • 시그니처: score_mod(score, b, h, q_idx, kv_idx)
  • 점수에 -inf를 더하면 마스킹처럼 동작할 수도 있지만, 블록 희소성 최적화를 잃을 수 있어요.

mask_mod

  • 위치 정보만으로 참여 여부를 결정(attend/마스크)할 때 사용해요.
  • 시그니처: mask_mod(b, h, q_idx, kv_idx) -> bool
  • create_block_maskBlockMask를 만들어 전달하면, 비어 있는 블록을 건너뛰어 2배 가까운 속도 향상을 얻을 수 있어요.

조합 팁

  • mask_modscore_mod동시에 전달 가능하고, score_mod는 마스크되지 않은 위치에만 적용돼요.
  • BlockMask만 바꾸는 것은 재컴파일을 유발하지 않아, 배치마다 문서 경계가 달라지는 상황에 효율적이에요.

더 알아보기