Paged 어텐션
Paged 어텐션
이 페이지는 continuous batching에서 사용되는 paged 어텐션 forward 함수를 설명합니다. 이 함수는 두 가지 버전의 flash attention 커널을 감싸서 다양한 배치 구성을 효율적으로 처리합니다.
출처: 문서
본문
이 페이지는 continuous batching에서 사용되는 paged 어텐션 forward 함수를 설명합니다. 이 함수는 두 가지 버전의 flash attention 커널을 감싸서 다양한 배치 구성을 효율적으로 처리합니다.
Varlen 경로
flash_attn_varlen_func 커널은 가변 길이 배치를 처리합니다. 이 경로는 prefill에서 요청 수가 많은 배치에 권장됩니다.
캐시 동작
이 커널은 paged 캐시와 직접 상호작용할 메커니즘이 없으므로, 캐시는 ~PagedAttentionCache.update 메서드를 사용해 수동으로 읽고 쓰입니다. 시퀀스 길이가 길어지면 이것이 병목 지점이 될 수 있습니다.
인덱싱 메커니즘
이 커널은 최대 시퀀스 길이(max_seqlen_q, max_seqlen_k)와 누적 시퀀스 길이(cu_seq_lens_q, cu_seq_lens_k)를 사용해 각 시퀀스의 어텐션을 계산합니다.
예시
쿼리 길이 [10, 3, 1], 키 길이 [0, 1, 7]인 3개 시퀀스의 배치를 생각해 봅시다.
cu_seq_lens_q = [0, 10, 13, 14]
cu_seq_lens_k = [0, 0, 1, 8]
max_seqlen_q = 10
max_seqlen_k = 7
입력 형태:
Q: [1, 10+3+1, num_heads, head_dim] = [1, 14, num_heads, head_dim]
K or V: [1, 0+1+7, num_kv_heads, head_dim] = [1, 8, num_kv_heads, head_dim]
커널은 누적 시퀀스 길이를 사용해 각 쿼리 및 키/값 토큰을 시퀀스에 할당합니다.
Q request index: [r0, r0, r0, r0, r0, r0, r0, r0, r0, r0, r1, r1, r1, r2]
cu_seq_lens_q: 0____________________________________10__________13__14
K request index: [r1, r2, r2, r2, r2, r2, r2, r2] (r0 has 0 K tokens)
cu_seq_lens_k: 0,0_1_______________________8
Decode 경로
flash_attn_with_kvcache 커널은 각 시퀀스에 정확히 하나의 쿼리 토큰이 있는 decode 전용 배치를 처리합니다. 이 경로는 varlen 경로보다 더 효율적이지만, prefill 요청이 있는 배치는 처리할 수 없습니다.
캐시 동작
이 커널은 block_table을 사용해 paged 캐시에 인덱싱하고 제자리(in-place)에서 업데이트합니다. 블록 테이블의 형태는 (batch_size, max_blocks_per_seq)이며, 각 행은 KV 캐시 텐서에서 요청의 캐시 블록이 위치한 물리적 위치를 담고 있습니다.
인덱싱 메커니즘
이 커널은 cache_seqlens를 사용해 각 시퀀스의 캐시 길이를 가져옵니다. 각 쿼리 토큰이 서로 다른 시퀀스에 속한다고(시퀀스당 토큰 하나) 가정합니다.
예시
쿼리 길이 [1, 1, 1], 키 길이 [30, 32, 70]인 3개 시퀀스의 배치를 생각해 봅시다. 캐시 블록 크기는 32이고 시퀀스당 최대 블록 수는 4입니다.
캐시 시퀀스 길이는 단순히 키 길이입니다.
cache_seqlens = [30, 32, 70]
블록 테이블의 형태는 (3, 4)입니다. 예시 주소를 사용하면:
block_table = [[2, -1, -1, -1],
[0, 1, -1, -1],
[3, 5, 6, -1]]
-1 값은 할당되지 않은 블록을 나타냅니다.
- Sequence 0 (캐시된 토큰 30개):
KV_cache[2]에 캐시됨. 새 토큰이 맞습니다(30 + 1 = 31 < 32). - Sequence 1 (캐시된 토큰 32개):
KV_cache[0]과KV_cache[1]에 캐시됨. 32 + 1 > 32이므로 두 번째 블록이 필요합니다. - Sequence 2 (캐시된 토큰 70개):
KV_cache[3],KV_cache[5],KV_cache[6]에 캐시됨. 블록이 반드시 연속적이지는 않다는 점에 주목하세요. 이것이 paged 캐시의 핵심 장점입니다. 새 토큰은 세 번째 블록에 맞습니다.