FlexAttention API — score_mod · mask_mod
FlexAttention API — score_mod · mask_mod
PyTorch 공식 API 문서는 FlexAttention의 두 확장점인 score_mod 와 mask_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_mask로BlockMask를 만들어 전달하면, 비어 있는 블록을 건너뛰어 2배 가까운 속도 향상을 얻을 수 있어요.
조합 팁
mask_mod와score_mod를 동시에 전달 가능하고, score_mod는 마스크되지 않은 위치에만 적용돼요.BlockMask만 바꾸는 것은 재컴파일을 유발하지 않아, 배치마다 문서 경계가 달라지는 상황에 효율적이에요.