๐Ÿค– LLM

GQA (Grouped-Query Attention) ๋ถ„์„

GQA โ€” MHA์™€ MQA ์‚ฌ์ด์˜ KV ํ—ค๋“œ ์„ค๊ณ„

Grouped-Query Attention(GQA)์€ Query head๋Š” ์—ฌ๋Ÿฌ ๊ฐœ ์œ ์ง€ํ•˜๋ฉด์„œ Key/Value head๋งŒ ์ค„์ด๊ณ , ์—ฌ๋Ÿฌ Query head๊ฐ€ ํ•˜๋‚˜์˜ KV head๋ฅผ ๊ณต์œ ํ•˜๋Š” Attention ๊ตฌ์กฐ์ž…๋‹ˆ๋‹ค. Multi-Head Attention(MHA)์˜ ํ‘œํ˜„๋ ฅ๊ณผ Multi-Query Attention(MQA)์˜ ์ž‘์€ KV cache ์‚ฌ์ด์— ์ค‘๊ฐ„ ์ง€์ ์„ ๋‘ก๋‹ˆ๋‹ค.

GQA์˜ ์ง์ ‘์ ์ธ ๋ชฉ์ ์€ ๋ชจ๋ธ ์ „์ฒด ํŒŒ๋ผ๋ฏธํ„ฐ๋ฅผ ํฌ๊ฒŒ ์ค„์ด๋Š” ๊ฒƒ์ด ์•„๋‹ˆ๋ผ autoregressive decoding์—์„œ ๋ฐ˜๋ณตํ•ด์„œ ์ฝ๋Š” KV cache์˜ ์šฉ๋Ÿ‰๊ณผ ๋ฉ”๋ชจ๋ฆฌ ๋Œ€์—ญํญ์„ ๋‚ฎ์ถ”๋Š” ๊ฒƒ์ž…๋‹ˆ๋‹ค. Ainslie et al.์€ ๊ธฐ์กด MHA checkpoint๋ฅผ ๊ทธ๋ฃนํ™”๋œ ๊ตฌ์กฐ๋กœ ๋ฐ”๊พผ ๋’ค ์›๋ž˜ ์‚ฌ์ „ํ•™์Šต ๊ณ„์‚ฐ๋Ÿ‰์˜ ์•ฝ 5%๋ฅผ ์‚ฌ์šฉํ•˜๋Š” uptraining recipe๋ฅผ ์ œ์‹œํ–ˆ๊ณ , GQA๊ฐ€ MQA์— ๊ฐ€๊นŒ์šด ์†๋„์™€ MHA์— ๊ฐ€๊นŒ์šด ํ’ˆ์งˆ์„ ์ œ๊ณตํ•  ์ˆ˜ ์žˆ์Œ์„ ๋ณด์˜€์Šต๋‹ˆ๋‹ค.

MHA, GQA, MQA์˜ head ๊ตฌ์„ฑ ๋น„๊ต

๊ทธ๋ฆผ 1. Query head ์ˆ˜๋Š” ์œ ์ง€ํ•˜๊ณ  KV head ์ˆ˜๋งŒ H์—์„œ G๋กœ ์ค„์ด๋Š” ๊ตฌ์กฐ

ํ•ต์‹ฌ ๊ฐœ๋…

Query head์™€ KV head์˜ ๋ถ„๋ฆฌ

Query head ์ˆ˜๋ฅผ H, Key/Value head ์ˆ˜๋ฅผ G๋ผ๊ณ  ํ•˜๋ฉด ์„ธ ๊ตฌ์กฐ๋Š” ๋‹ค์Œ์ฒ˜๋Ÿผ ํ‘œํ˜„ํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.

๊ตฌ์กฐ Query head ์ˆ˜ KV head ์ˆ˜ ๊ทธ๋ฃน๋‹น Query head ์ˆ˜
MHA H H 1
GQA H G, 1 < G < H H/G
MQA H 1 H

Hugging Face Transformers์˜ Llama ๊ณ„์—ด ์„ค์ •์—์„œ๋Š” num_attention_heads=H, num_key_value_heads=G๋กœ ์ด ๊ตฌ๋ถ„์„ ๋‚˜ํƒ€๋ƒ…๋‹ˆ๋‹ค. G=H์ด๋ฉด MHA, G=1์ด๋ฉด MQA, ๊ทธ ์‚ฌ์ด๋ฉด GQA์ž…๋‹ˆ๋‹ค. ์ผ๋ฐ˜์ ์ธ ๊ตฌํ˜„์€ H๊ฐ€ G๋กœ ๋‚˜๋ˆ„์–ด๋–จ์–ด์ง€๋„๋ก ๋ชจ๋ธ์„ ๊ตฌ์„ฑํ•˜์—ฌ ๊ฐ KV head๊ฐ€ ๊ฐ™์€ ์ˆ˜์˜ Query head๋ฅผ ๋‹ด๋‹นํ•˜๊ฒŒ ํ•ฉ๋‹ˆ๋‹ค.

KV cache ์šฉ๋Ÿ‰

๋ ˆ์ด์–ด ์ˆ˜๋ฅผ L, ๋ฐฐ์น˜๋ฅผ B, ๋ฌธ๋งฅ ๊ธธ์ด๋ฅผ T, head dimension์„ D, KV ์›์†Œ ํ•˜๋‚˜์˜ ๋ฐ”์ดํŠธ ์ˆ˜๋ฅผ s๋ผ๊ณ  ํ•˜๋ฉด KV cache์˜ ๊ตฌ์กฐ์  ์šฉ๋Ÿ‰์€ ๋‹ค์Œ๊ณผ ๊ฐ™์ด ๊ทผ์‚ฌํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.

KV_bytes โ‰ˆ 2 ร— L ร— B ร— T ร— G ร— D ร— s

์•ž์˜ 2๋Š” Key์™€ Value๋ฅผ ํ•จ๊ป˜ ์ €์žฅํ•œ๋‹ค๋Š” ๋œป์ž…๋‹ˆ๋‹ค. ๊ฐ™์€ L, B, T, D, ์ž๋ฃŒํ˜•์„ ์œ ์ง€ํ•˜๋ฉด GQA์˜ cache๋Š” MHA ๋Œ€๋น„ G/H ๋น„์œจ์ž…๋‹ˆ๋‹ค. ์˜ˆ๋ฅผ ๋“ค์–ด H=32, G=8์ธ GQA-8์€ head ์ฐจ์›๋งŒ์œผ๋กœ KV cache๋ฅผ 1/4๋กœ ์ค„์ž…๋‹ˆ๋‹ค. ์‹ค์ œ ์‚ฌ์šฉ๋Ÿ‰์—๋Š” block allocator, padding, scale metadata, tensor-parallel shard๊ฐ€ ์ถ”๊ฐ€๋˜๋ฏ€๋กœ ์ด ์‹์€ ์„ค๊ณ„ ๋น„๊ต์šฉ์œผ๋กœ ์‚ฌ์šฉํ•ด์•ผ ํ•ฉ๋‹ˆ๋‹ค.

Grouped broadcast

์ž…๋ ฅ X์—์„œ Query๋Š” H๊ฐœ head๋กœ, Key์™€ Value๋Š” G๊ฐœ head๋กœ ํˆฌ์˜ํ•ฉ๋‹ˆ๋‹ค.

Q = X ยท Wq  -> [B, H, T, D]
K = X ยท Wk  -> [B, G, T, D]
V = X ยท Wv  -> [B, G, T, D]

๊ฐ Query head h๋Š” ๋‹ค์Œ ๊ทธ๋ฃน์˜ KV๋ฅผ ์‚ฌ์šฉํ•ฉ๋‹ˆ๋‹ค.

g(h) = floor(h / (H / G))
score[h] = softmax(Q[h] ยท K[g(h)]แต€ / sqrt(D))
O[h]     = score[h] ยท V[g(h)]

๋…ผ๋ฆฌ์ ์œผ๋กœ๋Š” K/V๋ฅผ H/G๋ฒˆ ๋ฐ˜๋ณตํ•œ ๊ฒƒ๊ณผ ๊ฐ™์ง€๋งŒ, ํšจ์œจ์ ์ธ attention kernel์€ ๋™์ผํ•œ K/V๋ฅผ ์‹ค์ œ๋กœ ๋ณต์ œํ•˜์ง€ ์•Š๊ณ  group index๋กœ ์ฐธ์กฐํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค. ๋”ฐ๋ผ์„œ GQA๋Š” ํ•˜๋‚˜์˜ attention ๊ฒฐ๊ณผ๋ฅผ ๋ชจ๋“  head์— ๋ณต์‚ฌํ•˜๋Š” ๋ฐฉ์‹์ด ์•„๋‹™๋‹ˆ๋‹ค. Query๋งˆ๋‹ค ๋‹ค๋ฅธ score์™€ output์„ ๊ณ„์‚ฐํ•˜๋˜, ์ฝ๋Š” K/V sequence๋งŒ ๊ทธ๋ฃน ๋‚ด๋ถ€์—์„œ ๊ณต์œ ํ•ฉ๋‹ˆ๋‹ค.

GQA attention์˜ ๊ทธ๋ฃน ๋งคํ•‘

๊ทธ๋ฆผ 2. ๊ฐ Query group์ด ํ•˜๋‚˜์˜ KV sequence๋ฅผ ๊ณต์œ ํ•˜๋Š” attention ๊ฒฝ๋กœ

๋น„๊ต/๋ถ„์„

MHA, GQA, MQA ๋น„๊ต

ํ•ญ๋ชฉ MHA GQA MQA
KV head ์ˆ˜ H G, 1<G<H 1
KV cache ๋น„์œจ(MHA=1) 1 G/H 1/H
K/V ํ‘œํ˜„์˜ ๋…๋ฆฝ์„ฑ ๊ฐ€์žฅ ๋†’์Œ ๊ทธ๋ฃน ๋‚ด๋ถ€ ๊ณต์œ  ๋ชจ๋“  Query๊ฐ€ ๊ณต์œ 
Decode KV read volume ๊ฐ€์žฅ ํผ ์ค‘๊ฐ„ ๊ฐ€์žฅ ์ž‘์Œ
๊ธฐ์กด MHA ์ „ํ™˜ ๋‚œ์ด๋„ ํ•ด๋‹น ์—†์Œ ํ‰๊ท  ํ’€๋ง + uptraining ํ‰๊ท  ํ’€๋ง + uptraining
ํ’ˆ์งˆยท์†๋„ ์กฐ์ • ํญ ํ’ˆ์งˆ ๊ธฐ์ค€์„  G๋กœ ์กฐ์ ˆ ์†๋„ยท์šฉ๋Ÿ‰ ์ตœ์ ํ™”, ํ’ˆ์งˆ ์œ„ํ—˜ ํผ

GQA ๊ทธ๋ฃน ์ˆ˜ ์„ ํƒ

G๋ฅผ ์ž‘๊ฒŒ ํ•˜๋ฉด cache์™€ KV read volume์€ ์ค„์ง€๋งŒ ํ•œ KV ํ‘œํ˜„์„ ๊ณต์œ ํ•˜๋Š” Query head ์ˆ˜๊ฐ€ ๋Š˜์–ด ํ’ˆ์งˆ ์ €ํ•˜ ์œ„ํ—˜์ด ์ปค์งˆ ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค. ๋ฐ˜๋Œ€๋กœ G๋ฅผ ํฌ๊ฒŒ ํ•˜๋ฉด MHA์— ๊ฐ€๊นŒ์šด ํ‘œํ˜„๋ ฅ์„ ์œ ์ง€ํ•˜์ง€๋งŒ ๋ฉ”๋ชจ๋ฆฌ ์ ˆ๊ฐ ํญ๊ณผ decode ๋Œ€์—ญํญ ์ด๋“์ด ์ค„์–ด๋“ญ๋‹ˆ๋‹ค. ๋”ฐ๋ผ์„œ G๋Š” ๋‹จ์ˆœํžˆ ์ตœ๋Œ€ ์ ˆ๊ฐ๊ฐ’์œผ๋กœ ๊ณ ๋ฅด๋Š” ๊ฒƒ์ด ์•„๋‹ˆ๋ผ ๋ชจ๋ธ์˜ ํ’ˆ์งˆ ํ‰๊ฐ€์™€ GPU/tensor parallel ๋ฐฐ์น˜ ์ œ์•ฝ์„ ํ•จ๊ป˜ ๋ณด๊ณ  ์ •ํ•ด์•ผ ํ•ฉ๋‹ˆ๋‹ค.

์˜ˆ๋ฅผ ๋“ค์–ด H=32, D=128, L=32, B=1, T=4096, BF16(s=2)์ด๋ฉด ๋‹ค์Œ๊ณผ ๊ฐ™์Šต๋‹ˆ๋‹ค.

๊ตฌ์กฐ G ์ „์ฒด KV cache ๊ทผ์‚ฌ๊ฐ’
MHA 32 2 GiB
GQA-8 8 512 MiB
GQA-4 4 256 MiB
MQA 1 64 MiB

์ด ๊ฐ’์€ ํ•œ ์š”์ฒญ์˜ ๋ชจ๋“  ๋ ˆ์ด์–ด์™€ 4096๊ฐœ ํ† ํฐ์„ ํ•ฉ์นœ ์ด๋ก ๊ฐ’์ž…๋‹ˆ๋‹ค. allocator metadata, batch, padding, ์–‘์žํ™” scale์€ ํฌํ•จํ•˜์ง€ ์•Š์•˜์Šต๋‹ˆ๋‹ค.

Prefill๊ณผ Decode์˜ ์ฐจ์ด

Prefill์—์„œ๋Š” ์—ฌ๋Ÿฌ ์ž…๋ ฅ ํ† ํฐ์„ ๋ณ‘๋ ฌ๋กœ ์ฒ˜๋ฆฌํ•˜๋ฏ€๋กœ GEMM ์ฒ˜๋ฆฌ๋Ÿ‰๊ณผ ์—ฐ์‚ฐ ํšจ์œจ์ด ์ค‘์š”ํ•ฉ๋‹ˆ๋‹ค. Decode์—์„œ๋Š” ์ƒˆ Query๊ฐ€ ๋ณดํ†ต ํ•œ ํ† ํฐ์”ฉ ๋“ค์–ด์˜ค๊ณ , ๊ฐ ๋ ˆ์ด์–ด๊ฐ€ ๊ณผ๊ฑฐ T๊ฐœ ํ† ํฐ์˜ K/V๋ฅผ ๋‹ค์‹œ ์ฝ์Šต๋‹ˆ๋‹ค. ์ด๋•Œ GQA๋Š” ๊ณผ๊ฑฐ sequence์˜ KV head ์ˆ˜ ์ž์ฒด๋ฅผ G๊ฐœ๋กœ ์ค„์—ฌ HBM read pressure๋ฅผ ๋‚ฎ์ถฅ๋‹ˆ๋‹ค.

GQA๊ฐ€ attention์˜ ๋ชจ๋“  ์—ฐ์‚ฐ์„ G/H ๋น„์œจ๋กœ ์ค„์ด๋Š” ๊ฒƒ์€ ์•„๋‹™๋‹ˆ๋‹ค. Query projection, Query head๋ณ„ score ๊ณ„์‚ฐ, output projection์€ ์—ฌ์ „ํžˆ H๊ฐœ Query head๋ฅผ ๊ธฐ์ค€์œผ๋กœ ์ˆ˜ํ–‰๋ฉ๋‹ˆ๋‹ค. ์‹ค์ œ end-to-end ์ด๋“์€ context length, batch size, attention kernel, memory bandwidth, prefill/decode ๋น„์œจ์— ๋”ฐ๋ผ ๋‹ฌ๋ผ์ง‘๋‹ˆ๋‹ค.

๋™์ž‘ ์›๋ฆฌ

1. MHA checkpoint์—์„œ GQA๋กœ ๋ณ€ํ™˜

๊ธฐ์กด MHA์˜ K/V projection์—๋Š” H๊ฐœ์˜ head๊ฐ€ ์žˆ์Šต๋‹ˆ๋‹ค. ๋ชฉํ‘œ GQA๊ฐ€ G๊ฐœ์˜ KV head๋ฅผ ์‚ฌ์šฉํ•˜๊ณ  H/G๊ฐ€ ์ •์ˆ˜๋ผ๋ฉด, ์›๋ž˜ K/V head๋ฅผ ๊ทธ๋ฃน๋ณ„๋กœ ๋ฌถ๊ณ  ๊ฐ ๊ทธ๋ฃน์˜ ๊ฐ€์ค‘์น˜๋ฅผ ํ‰๊ท  ํ’€๋งํ•˜์—ฌ ์ƒˆ K/V projection์„ ์ดˆ๊ธฐํ™”ํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค. Query projection๊ณผ output projection์€ ์œ ์ง€ํ•ฉ๋‹ˆ๋‹ค.

MHA checkpoint๋ฅผ GQA๋กœ ๋ฐ”๊พธ๋Š” ํ๋ฆ„

๊ทธ๋ฆผ 3. ๊ทธ๋ฃน๋ณ„ K/V ํ‰๊ท  ํ’€๋ง๊ณผ uptraining์„ ํ†ตํ•œ checkpoint ๋ณ€ํ™˜

for g in range(G):
    start = g * (H // G)
    end = (g + 1) * (H // G)
    Wk_g = mean(Wk_heads[start:end], axis=head)
    Wv_g = mean(Wv_heads[start:end], axis=head)

ํ‰๊ท  ํ’€๋ง์€ ๋ณ€ํ™˜ ์งํ›„์˜ ์ดˆ๊ธฐํ™” ๋ฐฉ๋ฒ•์ด๋ฉฐ, ํ’ˆ์งˆ ํšŒ๋ณต์„ ๋ณด์žฅํ•˜๋Š” ๊ฒƒ ์ž์ฒด๋Š” ์•„๋‹™๋‹ˆ๋‹ค. GQA ๋…ผ๋ฌธ์€ ๋ณ€ํ™˜ํ•œ checkpoint๋ฅผ ์ œํ•œ๋œ ์ถ”๊ฐ€ ํ•™์Šต์œผ๋กœ ๋ณด์ •ํ•˜๋Š” uptraining์„ ์‚ฌ์šฉํ–ˆ์Šต๋‹ˆ๋‹ค. ์‹ค์ œ ์ „ํ™˜์—์„œ๋Š” tokenizer, RoPE ์„ค์ •, optimizer ์ƒํƒœ, calibration ๋ฐ์ดํ„ฐ์™€ ํ‰๊ฐ€ workload๋ฅผ ๋ชจ๋ธ๋ณ„๋กœ ๋งž์ถฐ์•ผ ํ•ฉ๋‹ˆ๋‹ค.

2. Projection๊ณผ RoPE

๊ฐ decoder layer๋Š” hidden state์—์„œ Q, K, V๋ฅผ ์ƒ์„ฑํ•ฉ๋‹ˆ๋‹ค. RoPE๋ฅผ ์‚ฌ์šฉํ•˜๋Š” ๋ชจ๋ธ์ด๋ผ๋ฉด Query์™€ Key์— ์œ„์น˜ ํšŒ์ „์„ ์ ์šฉํ•œ ๋’ค Key๋ฅผ KV cache์— ์ถ”๊ฐ€ํ•ฉ๋‹ˆ๋‹ค. GQA๋Š” Key/Value head ์ˆ˜๋ฅผ ์ค„์ด๋Š” ๊ตฌ์กฐ์ด๋ฏ€๋กœ Query head๋ณ„ ์œ„์น˜ ์ฒ˜๋ฆฌ์™€ score ๊ณ„์‚ฐ์€ ๋…๋ฆฝ์ ์œผ๋กœ ๋‚จ์Šต๋‹ˆ๋‹ค.

Q_t = RoPE(project_Q(x_t))       # [B, H, 1, D]
K_t = RoPE(project_K(x_t))       # [B, G, 1, D]
V_t = project_V(x_t)              # [B, G, 1, D]
append(K_cache, K_t)
append(V_cache, V_t)

3. Incremental decode

์ƒˆ ํ† ํฐ t๋ฅผ ์ฒ˜๋ฆฌํ•  ๋•Œ ํ˜„์žฌ Query Q_t[h]๋Š” ์ž์‹ ์ด ์†ํ•œ ๊ทธ๋ฃน g(h)์˜ ์ „์ฒด cache์™€ ๋‚ด์ ํ•ฉ๋‹ˆ๋‹ค. ๊ทธ๋ฃน์ด ๊ฐ™์•„๋„ Query๊ฐ€ ๋‹ค๋ฅด๋ฏ€๋กœ attention probability์™€ output์€ ์„œ๋กœ ๋‹ค๋ฆ…๋‹ˆ๋‹ค.

for each layer:
    q_t = project_query(x_t)       # H heads
    k_t = project_key(x_t)         # G heads
    v_t = project_value(x_t)       # G heads
    append(K_cache, k_t)
    append(V_cache, v_t)

    for h in range(H):
        g = h // (H // G)
        a_h = softmax(q_t[h] @ K_cache[g].T / sqrt(D))
        o_h = a_h @ V_cache[g]

    y_t = concat(o_0, ..., o_{H-1}) @ Wo

์ด ๊ฒฝ๋กœ์—์„œ cache์˜ ์‹œ๊ฐ„์ถ• ๊ธธ์ด T๋Š” ๋™์ผํ•˜์ง€๋งŒ KV head ์ถ•์ด H์—์„œ G๋กœ ์ค„์–ด๋“ญ๋‹ˆ๋‹ค. ์ปค๋„์ด K/V๋ฅผ materializeํ•˜์ง€ ์•Š๊ณ  group mapping์„ ์ฒ˜๋ฆฌํ•˜๋ฉด ์ค‘๊ฐ„ broadcast buffer๋„ ์ค„์ผ ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.

4. Tensor parallelism๊ณผ kernel ์ œ์•ฝ

Tensor parallelism์—์„œ๋Š” Query head์™€ KV head๊ฐ€ GPU shard์— ์–ด๋–ป๊ฒŒ ๋ฐฐ์น˜๋˜๋Š”์ง€๊ฐ€ ์ค‘์š”ํ•ฉ๋‹ˆ๋‹ค. G๊ฐ€ GPU ์ˆ˜ ๋˜๋Š” tensor-parallel degree์™€ ์ž˜ ๋งž์ง€ ์•Š์œผ๋ฉด KV head๋ฅผ ๊ท ๋“ฑํ•˜๊ฒŒ ๋‚˜๋ˆ„๊ธฐ ์–ด๋ ต๊ณ , attention ์ „์— KV๋ฅผ ๋ณต์ œํ•˜๊ฑฐ๋‚˜ ํ†ต์‹ ํ•˜๋Š” ๊ฒฝ๋กœ๊ฐ€ ์ƒ๊ธธ ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค. ๋”ฐ๋ผ์„œ ์ด๋ก ์ ์ธ G/H cache ๊ฐ์†Œ๊ฐ€ ๊ฐ™์€ ๋น„์œจ์˜ latency ๊ฐ์†Œ๋กœ ์ด์–ด์ง„๋‹ค๊ณ  ๊ฐ€์ •ํ•˜๋ฉด ์•ˆ ๋ฉ๋‹ˆ๋‹ค.

๊ตฌํ˜„ ์‹œ ๋‹ค์Œ์„ ํ™•์ธํ•ด์•ผ ํ•ฉ๋‹ˆ๋‹ค.

  • num_attention_heads, num_key_value_heads, head_dim์˜ ์ •ํ•ฉ์„ฑ๊ณผ H % G == 0 ์กฐ๊ฑด
  • KV cache layout์ด backend๊ฐ€ ์š”๊ตฌํ•˜๋Š” [batch, kv_heads, sequence, head_dim] ๊ณ„์—ด์ธ์ง€ ์—ฌ๋ถ€
  • repeat_kv๊ฐ€ ์‹ค์ œ ๋ฉ”๋ชจ๋ฆฌ ๋ณต์ œ์ธ์ง€ fused kernel ๋‚ด๋ถ€์˜ ๋…ผ๋ฆฌ์  broadcast์ธ์ง€ ์—ฌ๋ถ€
  • tensor parallel shard๋ณ„ KV head ๋ฐฐ์น˜์™€ all-gather/broadcast ํ†ต์‹ ๋Ÿ‰
  • GQA์™€ FlashAttention, PagedAttention, KV cache quantization์˜ shapeยทscale ํ˜ธํ™˜์„ฑ

์žฅ๋‹จ์ 

์žฅ์ 

  • MHA ๋Œ€๋น„ KV cache ์šฉ๋Ÿ‰๊ณผ decode ๋‹จ๊ณ„์˜ KV memory read๋ฅผ G/H ์ˆ˜์ค€์œผ๋กœ ์ค„์ผ ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.
  • Query head ์ˆ˜๋ฅผ ์œ ์ง€ํ•˜๋ฏ€๋กœ MQA๋ณด๋‹ค ํ‘œํ˜„ ๊ณต๊ฐ„๊ณผ ํ’ˆ์งˆ์„ ๋” ๋ณด์กดํ•  ์—ฌ์ง€๊ฐ€ ์žˆ์Šต๋‹ˆ๋‹ค.
  • ๊ธด context, ํฐ batch, beam search์ฒ˜๋Ÿผ KV cache๊ฐ€ ์ปค์ง€๋Š” serving workload์— ์œ ๋ฆฌํ•ฉ๋‹ˆ๋‹ค.
  • G๋ฅผ ์กฐ์ ˆํ•ด ํ’ˆ์งˆยท๋ฉ”๋ชจ๋ฆฌยท๋Œ€์—ญํญ ์‚ฌ์ด์˜ ์ ˆ์ถฉ์ ์„ ๋ชจ๋ธ๋ณ„๋กœ ์„ ํƒํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.
  • PagedAttention์€ cache ๋ฐฐ์น˜ ๋ฐฉ์‹์„, KV quantization์€ ์›์†Œ ์ •๋ฐ€๋„๋ฅผ ๋ฐ”๊พธ๋ฏ€๋กœ GQA์™€ ์ง๊ต์ ์œผ๋กœ ๊ฒฐํ•ฉํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.

๋‹จ์ 

  • K/V ํ‘œํ˜„์„ ๊ณต์œ ํ•˜๋ฏ€๋กœ MHA์˜ head๋ณ„ ๋…๋ฆฝ์„ฑ๋ณด๋‹ค ํ‘œํ˜„๋ ฅ์ด ์ค„๊ณ  ํ’ˆ์งˆ ์ €ํ•˜๊ฐ€ ๋ฐœ์ƒํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.
  • ๊ธฐ์กด MHA checkpoint๋ฅผ ๋‹จ์ˆœํžˆ head๋ฅผ ์‚ญ์ œํ•˜๋Š” ๋ฐฉ์‹์œผ๋กœ ๋ณ€ํ™˜ํ•˜๋ฉด ๋ถ„ํฌ๊ฐ€ ๋ฐ”๋€Œ๋ฏ€๋กœ ํ‰๊ท  ํ’€๋ง๊ณผ ์ถ”๊ฐ€ ํ•™์Šต์ด ํ•„์š”ํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.
  • ์ž‘์€ G๋Š” tensor parallel shard์— KV head๋ฅผ ๋ฐฐ์น˜ํ•˜๊ธฐ ์–ด๋ ต๊ฒŒ ๋งŒ๋“ค์–ด broadcast ๋˜๋Š” ํ†ต์‹  ๋น„์šฉ์„ ์œ ๋ฐœํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.
  • Prefill์ด compute-bound์ด๋ฉด decode์—์„œ์˜ KV read ์ ˆ๊ฐ์ด ์ „์ฒด ์š”์ฒญ latency์— ์ œํ•œ์ ์œผ๋กœ๋งŒ ๋ฐ˜์˜๋ฉ๋‹ˆ๋‹ค.
  • ๋ชจ๋ธ ์„ค์ •์„ ๋ฐ”๊พธ๋Š” ์•„ํ‚คํ…์ฒ˜ ๊ธฐ๋ฒ•์ด๋ฏ€๋กœ ๊ธฐ์กด weight์™€ inference backend๊ฐ€ GQA shape์„ ์ง€์›ํ•ด์•ผ ํ•ฉ๋‹ˆ๋‹ค.

๊ด€๋ จ ๊ธฐ์ˆ  ๋ฐ ์ฐธ๊ณ  ๋ฌธํ—Œ

๋ฌธ์„œ/์—ฐ๊ตฌ ์—ฐ๊ฒฐ์ 
LLM Basics Transformer decoder, prefill/decode, KV cache ๊ธฐ๋ณธ ๊ฐœ๋…
MQA Analysis GQA์˜ ์–‘ ๋์ ์ธ MHA์™€ MQA์˜ ๊ตฌ์กฐยทcache ๋น„๊ต
PagedAttention Analysis KV cache๋ฅผ ๋ธ”๋ก์œผ๋กœ ๋ฐฐ์น˜ํ•˜๊ณ  ๊ด€๋ฆฌํ•˜๋Š” serving ๊ธฐ๋ฒ•
KV Cache Quantization Analysis GQA์™€ ๊ฒฐํ•ฉํ•  ์ˆ˜ ์žˆ๋Š” KV ์›์†Œ ์ •๋ฐ€๋„ ์ถ•์†Œ
MLA Analysis KV head ๊ณต์œ ๊ฐ€ ์•„๋‹Œ latent ์••์ถ•์„ ์‚ฌ์šฉํ•˜๋Š” ๋Œ€์•ˆ
Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints GQA ์ •์˜, MHA checkpoint ๋ณ€ํ™˜, 5% uptraining recipe
Shazeer, Fast Transformer Decoding: One Write-Head is All You Need MQA์™€ incremental decoding์˜ memory-bandwidth ๋ฌธ์ œ
Hugging Face LlamaConfig ๋ฌธ์„œ num_key_value_heads๋กœ MHA/GQA/MQA๋ฅผ ์„ค์ •ํ•˜๋Š” ๊ทœ์น™

ํ•ต์‹ฌ ์ •๋ฆฌ

GQA๋Š” H๊ฐœ์˜ Query head๋ฅผ ์œ ์ง€ํ•˜๋ฉด์„œ G๊ฐœ์˜ KV head๋งŒ ๋‘๊ณ , ๊ฐ KV head๋ฅผ H/G๊ฐœ์˜ Query head๊ฐ€ ๊ณต์œ ํ•˜๋Š” Attention ๊ตฌ์กฐ์ž…๋‹ˆ๋‹ค. KV cache ์šฉ๋Ÿ‰์€ KV head ์ˆ˜์— ๋น„๋ก€ํ•˜๋ฏ€๋กœ MHA ๋Œ€๋น„ G/H ์ˆ˜์ค€์œผ๋กœ ์ค„์ง€๋งŒ, Query ๊ณ„์‚ฐ๊ณผ ๋ชจ๋“  head์˜ score ๊ณ„์‚ฐ์ด ๊ฐ™์€ ๋น„์œจ๋กœ ์ค„์–ด๋“œ๋Š” ๊ฒƒ์€ ์•„๋‹™๋‹ˆ๋‹ค. ๊ธฐ์กด MHA checkpoint๋Š” ๊ทธ๋ฃน๋ณ„ K/V ํ‰๊ท  ํ’€๋ง ํ›„ ์ œํ•œ๋œ uptraining์œผ๋กœ GQA์— ๋งž์ถœ ์ˆ˜ ์žˆ์œผ๋ฉฐ, G ์„ ํƒ์€ ํ’ˆ์งˆ๋ฟ ์•„๋‹ˆ๋ผ tensor parallel ๋ฐฐ์น˜์™€ kernel ์ง€์›๊นŒ์ง€ ๊ณ ๋ คํ•ด์•ผ ํ•ฉ๋‹ˆ๋‹ค. PagedAttention๊ณผ KV cache quantization์€ ์„œ๋กœ ๋‹ค๋ฅธ ์ถ•์„ ์ตœ์ ํ™”ํ•˜๋ฏ€๋กœ GQA์™€ ํ•จ๊ป˜ ์ ์šฉํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.