๋…ผ๋ฌธ ๋งํฌ

๊ฐ•์ ๊ณผ ๋…์ฐฝ์„ฑ, ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜, ๊ทธ๋ฆฌ๊ณ  ํ•œ๊ณ„

๋ฌผ์–ด๋ณธ ๊ฒƒ: ์ด ๋…ผ๋ฌธ์„ ์•„์ฃผ ์ž์„ธํ•˜๊ฒŒ ์ฝ๊ณ  ๋…ผ๋ฌธ์˜ ๊ฐ•์ ๊ณผ ๋…์ฐฝ์ ์ธ ์ง€์ ์„ ์„ค๋ช…ํ•ด์ฃผ๊ณ  ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜์„ ์˜ˆ์‹œ ์ž…๋ ฅ์„ ๋“ค์–ด์„œ ์ „์ฒด์ ์ธ ๊ณผ์ •์„ ์„ค๋ช…ํ•ด์ค˜ ์ถ”๊ฐ€์ ์œผ๋กœ ๋…ผ๋ฌธ์˜ ํ•œ๊ณ„์ ์— ๋Œ€ํ•ด์„œ๋„ ์•Œ๋ ค์ค˜

๐Ÿ“Œ ๊ฒฐ๋ก  ์š”์•ฝ

๋…ผ๋ฌธ โ€˜Seesaw: High-throughput LLM Inference via Model Re-shardingโ€™์€ LLM์˜ ๋‘ ์ฃผ์š” ๋‹จ๊ณ„์ธ prefill๊ณผ decode์—์„œ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋™์ ์œผ๋กœ ๋ณ€๊ฒฝํ•˜๋Š” โ€˜model re-shardingโ€™ ๊ธฐ๋ฒ•์„ ์ œ์•ˆํ•˜์—ฌ, ํ‰๊ท  1.36๋ฐฐ, ์ตœ๋Œ€ 1.78๋ฐฐ์˜ ์ถ”๋ก  ์ฒ˜๋ฆฌ๋Ÿ‰ ๊ฐœ์„ ์„ ๋‹ฌ์„ฑํ•ฉ๋‹ˆ๋‹ค. ์ด ๋ฐฉ์‹์€ ๊ธฐ์กด vLLM์ฒ˜๋Ÿผ ๊ณ ์ •๋œ ๋ณ‘๋ ฌํ™” ์ „๋žต์— ๋น„ํ•ด throughput ์ตœ์ ํ™”๋ฅผ ๋‹ฌ์„ฑํ•˜๋ฉฐ, tiered KV cache buffering๊ณผ transition-minimizing scheduling์„ ํ†ตํ•ด ์žฌ์ƒค๋”ฉ ๋น„์šฉ๊นŒ์ง€ ํšจ๊ณผ์ ์œผ๋กœ ์ค„์ž…๋‹ˆ๋‹ค.


โœ… ๋…ผ๋ฌธ์˜ ์ฃผ์š” ๊ธฐ์—ฌ

๊ตฌ๋ถ„๋‚ด์šฉ
๋ฌธ์ œ์ prefill๊ณผ decode๋Š” ์„ฑ๊ฒฉ์ด ๋‹ฌ๋ผ ๋™์ผํ•œ ๋ณ‘๋ ฌํ™” ์ „๋žต ์‚ฌ์šฉ ์‹œ ๋น„ํšจ์œจ ์ดˆ๋ž˜
ํ•ด๊ฒฐ์ฑ…๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋‹จ๊ณ„๋ณ„๋กœ ๋™์ ์œผ๋กœ ์กฐ์ • (model re-sharding)
์ถ”๊ฐ€ ๊ธฐ๋ฒ•- tiered KV cache buffering (CPU ๋ฉ”๋ชจ๋ฆฌ ๋ณด์กฐ ์ €์žฅ์†Œ ์‚ฌ์šฉ)
- transition-minimizing scheduling (stage ์ „ํ™˜ ์ตœ์†Œํ™”)
ํšจ๊ณผvLLM ๋Œ€๋น„ ํ‰๊ท  1.36ร—, ์ตœ๋Œ€ 1.78ร— throughput ํ–ฅ์ƒ

โš™๏ธ ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜: ๋™์  Model Re-sharding ๊ณผ์ • ์˜ˆ์‹œ

๐Ÿ”ง ์˜ˆ์‹œ ์ž…๋ ฅ

  • LLM: LLaMA2-13B
  • ์‹œ์Šคํ…œ: 8 ร— L4 GPU (PCIe ์—ฐ๊ฒฐ)
  • ์ž…๋ ฅ ์š”์ฒญ ์ˆ˜: 16๊ฐœ
  • ๋ณ‘๋ ฌํ™” ์ „๋žต: prefill์— pipeline parallelism (PP4), decode์— tensor parallelism (TP4)

๐Ÿ”„ ์ „์ฒด ์ถ”๋ก  ์ฒ˜๋ฆฌ ํ๋ฆ„

  1. Prefill ๋‹จ๊ณ„ (PP4 ์ ์šฉ):

    • ์ž…๋ ฅ ์ „์ฒด ์‹œํ€€์Šค๊ฐ€ GPU 4๋Œ€์— pipeline ํ˜•ํƒœ๋กœ ๋ถ„๋ฐฐ๋จ
    • ๋ชจ๋ธ ๋ ˆ์ด์–ด๋„ ์ˆœ์ฐจ์ ์œผ๋กœ ๊ฐ GPU์— ๋‚˜๋‰˜์–ด ์ฒ˜๋ฆฌ
    • ํ†ต์‹  ์˜ค๋ฒ„ํ—ค๋“œ ์ž‘์Œ, weight loading ํšจ์œจ์ ์ž„
  2. KV ์บ์‹œ ์ƒ์„ฑ ๋ฐ CPU๋กœ ์˜คํ”„๋กœ๋“œ:

    • ๊ฐ GPU๋Š” ์ž์‹ ์˜ shard KV๋ฅผ CPU ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ์— ์ €์žฅ
  3. ๋ชจ๋ธ weight์™€ KV cache ์žฌ์ƒค๋”ฉ:

    • pipeline ๊ตฌ์กฐ๋กœ ๋‚˜๋‰˜์—ˆ๋˜ ๋ชจ๋ธ weight โ†’ tensor ๊ตฌ์กฐ๋กœ ์žฌ๋ฐฐ์น˜
    • KV ์บ์‹œ๋„ TP ๊ตฌ์กฐ๋กœ ์žฌ๋ถ„๋ฐฐ๋จ (CPUโ†’GPU ๋น„๋™๊ธฐ ์ „์†ก)
  4. Decode ๋‹จ๊ณ„ (TP4 ์ ์šฉ):

    • ๋ชจ๋“  GPU๊ฐ€ ๋™์ผ weight shard๋กœ ๋ณ‘๋ ฌ ์ฒ˜๋ฆฌ
    • ํ•œ ํ† ํฐ์”ฉ ์ƒ์„ฑ โ†’ GPU ๊ฐ„ AllReduce ํ†ต์‹  ๋ฐœ์ƒ
    • weight loading ๋ณ‘๋ ฌ ์ฒ˜๋ฆฌ๋กœ ํšจ์œจ์„ฑ ๊ทน๋Œ€ํ™”๋จ
  5. ๋น„๋™๊ธฐ ์ฒ˜๋ฆฌ ๋ฐ stage ์ „ํ™˜ ์ตœ์ ํ™”:

    • CPUโ†’GPU KV ์ „์†ก์€ prefetch thread๋กœ ๋น„๋™๊ธฐ ์ˆ˜ํ–‰
    • KV๊ฐ€ ๊ฝ‰ ์ฐฐ ๋•Œ๋งŒ decode๋กœ ์ „ํ™˜ํ•˜์—ฌ ์žฌ์ƒค๋”ฉ ํšŸ์ˆ˜ ์ตœ์†Œํ™”

๐Ÿ“ˆ ์„ฑ๋Šฅ ๋น„๊ต (vLLM vs Seesaw)

ํ™˜๊ฒฝ๋ชจ๋ธvLLM (req/s)Seesaw (req/s)์†๋„ ํ–ฅ์ƒ ๋ฐฐ์ˆ˜
A10, 4GPU15B1.0 (๊ธฐ์ค€)1.45+45%
L4, 4GPU15B1.01.29+29%
A100 (PCIe), 8GPU70B1.01.46+46%
A100 (NVLink), 8GPU70B1.01.13+13%

๐Ÿง  ๋…ผ๋ฌธ์˜ ๋…์ฐฝ์„ฑ

ํ•ญ๋ชฉSeesaw์˜ ์ฐจ๋ณ„์ 
๋ณ‘๋ ฌํ™” ์ „๋žตstage ๋ณ„๋กœ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋ฐ”๊ฟ€ ์ˆ˜ ์žˆ๋„๋ก ์„ค๊ณ„ (TP โ†” PP)
Re-sharding overhead ์ฒ˜๋ฆฌCPU๋ฅผ ์ค‘๊ฐ„ ์ €์žฅ์†Œ๋กœ ํ™œ์šฉํ•œ tiered buffering ๋„์ž…
Scheduling ์ตœ์ ํ™”transition-minimizing scheduler ์„ค๊ณ„
์ ์šฉ ์œ ์—ฐ์„ฑ๊ธฐ์กด ์‹œ์Šคํ…œ(vLLM, TensorRT-LLM)์—๋„ ์‰ฝ๊ฒŒ ํ†ตํ•ฉ ๊ฐ€๋Šฅ

โš ๏ธ ํ•œ๊ณ„์  ๋ฐ ์ œ์•ฝ

  1. ์žฌ์ƒค๋”ฉ ์˜ค๋ฒ„ํ—ค๋“œ ์กด์žฌ

    • CPUโ†”GPU ์‚ฌ์ด์˜ ๋ฐ์ดํ„ฐ ์ด๋™์€ ์—ฌ์ „ํžˆ ๋ณ‘๋ชฉ ๊ฐ€๋Šฅ์„ฑ ์žˆ์Œ
    • ํ•ด๊ฒฐ์„ ์œ„ํ•ด ๋น„๋™๊ธฐ prefetching ์‚ฌ์šฉํ•˜์ง€๋งŒ ํ•˜๋“œ์›จ์–ด ์˜์กด๋„ ๋†’์Œ
  2. ๋ฉ”๋ชจ๋ฆฌ ์‚ฌ์šฉ ์ฆ๊ฐ€

    • CPU, GPU ๋ชจ๋‘์—์„œ KV cache๋ฅผ ๋ณด๊ด€ํ•ด์•ผ ํ•˜๋ฏ€๋กœ ๋” ๋งŽ์€ ๋ฉ”๋ชจ๋ฆฌ ์š”๊ตฌ
  3. ๋ณต์žกํ•œ ์•„ํ‚คํ…์ฒ˜ ๊ตฌ์„ฑ

    • worker-scheduler ๊ตฌ์กฐ์™€ ๋™๊ธฐํ™” ํ•„์š”๋กœ ์ธํ•ด ์‹œ์Šคํ…œ ๊ตฌํ˜„ ๋ณต์žก๋„ ์ฆ๊ฐ€
  4. online-serving ํ™˜๊ฒฝ์—์„œ๋Š” ๋œ ์ ํ•ฉ

    • offline throughput ์ค‘์‹ฌ ์„ค๊ณ„ โ†’ latency ๋ฏผ๊ฐํ•œ ์‹œ์Šคํ…œ์—๋Š” ๋ถ€์ ์ ˆ

ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜์„ ์˜ˆ์‹œ๋กœ ๋”ฐ๋ผ๊ฐ€๊ธฐ

๋ฌผ์–ด๋ณธ ๊ฒƒ: ๋…ผ๋ฌธ์˜ ํ•ต์‹ฌ์•Œ๊ณ ๋ฆฌ์ฆ˜์„ ์˜ˆ์‹œ๋ฅผ ๋“ค์–ด ์•„์ฃผ ์ž์„ธํ•˜๊ฒŒ ์„ค๋ช…ํ•ด์ฃผ๊ณ  ์ •๋ฆฌํ•ด ์ค˜

๋…ผ๋ฌธ “Seesaw: High-throughput LLM Inference via Model Re-sharding”์˜ ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜์€ ๋‹ค์Œ 3๊ฐ€์ง€ ๊ตฌ์„ฑ ์š”์†Œ๋กœ ์š”์•ฝ๋ฉ๋‹ˆ๋‹ค:


โœ… ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜ ๊ตฌ์„ฑ ์š”์•ฝ

๊ตฌ์„ฑ ์š”์†Œ์—ญํ• 
1. Dynamic Model Re-shardingPrefill ๋‹จ๊ณ„์—์„  Pipeline Parallelism (PP), Decode ๋‹จ๊ณ„์—์„  Tensor Parallelism (TP)์„ ์‚ฌ์šฉํ•˜๋„๋ก ๋ชจ๋ธ weight์™€ KV cache์˜ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋™์ ์œผ๋กœ ๋ณ€๊ฒฝ
2. Tiered KV Cache BufferingGPU ๋ฉ”๋ชจ๋ฆฌ ํ•œ๊ณ„๋ฅผ ๊ทน๋ณตํ•˜๊ธฐ ์œ„ํ•ด CPU ๋ฉ”๋ชจ๋ฆฌ๋ฅผ ํ™œ์šฉํ•˜์—ฌ KV cache๋ฅผ ์ž„์‹œ ์ €์žฅ
3. Transition-Minimizing Scheduling์ž์ฃผ stage๋ฅผ ์˜ค๊ฐ€๋Š” transition overhead๋ฅผ ์ค„์ด๊ธฐ ์œ„ํ•ด ์ผ์ • ์‹œ์ (๋ฒ„ํผ๊ฐ€ ๊ฐ€๋“ ์ฐผ์„ ๋•Œ)๊นŒ์ง€๋งŒ prefill ์ˆ˜ํ–‰ ํ›„ decode๋กœ ์ „ํ™˜

๐ŸŽฏ ์˜ˆ์‹œ๋กœ ๋ณธ ์•Œ๊ณ ๋ฆฌ์ฆ˜ ๋™์ž‘ (๋ชจ๋ธ: LLaMA2-13B, GPU 4๋Œ€ ์‚ฌ์šฉ)

์˜ˆ์‹œ ์ž…๋ ฅ

  • ์ด ์š”์ฒญ: 16๊ฐœ์˜ sequence
  • Prefill ๋ณ‘๋ ฌํ™” ์ „๋žต: PP4 (Pipeline Parallelism with 4-way split)
  • Decode ๋ณ‘๋ ฌํ™” ์ „๋žต: TP4 (Tensor Parallelism with 4-way split)

๐Ÿ”„ ์ „์ฒด ์•Œ๊ณ ๋ฆฌ์ฆ˜ ํ”Œ๋กœ์šฐ

โ‘  Prefill ๋‹จ๊ณ„: Pipeline Parallelism ์ ์šฉ

  • ๋ชจ๋ธ ๋ ˆ์ด์–ด๋ฅผ GPU 4๊ฐœ์— ์ˆœ์ฐจ์ ์œผ๋กœ ๋‚˜๋ˆ ์„œ ๋ฐฐ์น˜ (L/4 ๋ ˆ์ด์–ด์”ฉ).
  • 16๊ฐœ์˜ ์‹œํ€€์Šค๋ฅผ micro-batch๋กœ ๋‚˜๋ˆ ์„œ ํŒŒ์ดํ”„๋ผ์ธ์— ํˆฌ์ž….
  • ๊ฐ ์‹œํ€€์Šค๋Š” ํ•œ ๋ฒˆ์— ์—ฌ๋Ÿฌ ํ† ํฐ์„ ์ฒ˜๋ฆฌํ•˜๋ฏ€๋กœ GPU compute ์ž์›์„ ์ž˜ ํ™œ์šฉํ•จ.
  • ์ด๋•Œ ์ƒ์„ฑ๋œ KV Cache๋Š” GPU์— ์ €์žฅํ•˜์ง€ ์•Š๊ณ  CPU ๋ฉ”๋ชจ๋ฆฌ๋กœ ์˜คํ”„๋กœ๋“œํ•จ.
PLAINTEXT
[GPU1] L1 | [GPU2] L2 | [GPU3] L3 | [GPU4] L4
       โ†“         โ†“         โ†“         โ†“
Sequence 1 โ†’ Sequence 2 โ†’ ... โ†’ Sequence 16

โ†’ ๊ฐ ์‹œํ€€์Šค๋ณ„ KV ์บ์‹œ๊ฐ€ CPU ๋ฉ”๋ชจ๋ฆฌ์— ์ €์žฅ๋จ


โ‘ก KV Cache Tiered Buffering: GPUโ†’CPU๋กœ KV ์ด๋™

  • KV ์บ์‹œ๋Š” ๊ฐ GPU์—์„œ ๋กœ์ปฌ shard๋กœ ๊ณ„์‚ฐ๋œ ํ›„ CPU ๊ณต์œ  ๋ฉ”๋ชจ๋ฆฌ์— ์ €์žฅ๋จ
  • HND(H, N, D) ํฌ๋งท์œผ๋กœ ์ €์žฅ โ†’ TP์—์„œ head ์ถ• ๊ธฐ์ค€์œผ๋กœ ํšจ์œจ์  ์ ‘๊ทผ ๊ฐ€๋Šฅ
  • CPU ๋ฉ”๋ชจ๋ฆฌ๋Š” KV ์บ์‹œ๋ฅผ ์™„์ „ํžˆ ๋‹ด์„ ์ˆ˜ ์žˆ์„ ๋•Œ๊นŒ์ง€ ๊ณ„์† ์ฑ„์›€

โ‘ข Model Re-sharding: ๋ชจ๋ธ weight ์žฌ๊ตฌ์„ฑ

  • Prefill์ด ๋๋‚˜๋ฉด ๋‹ค์Œ ๋‹จ๊ณ„์ธ decode๋ฅผ ์œ„ํ•ด model weight์˜ sharding ๋ฐฉ์‹์„ TP4๋กœ ๋ณ€๊ฒฝ
  • TP๋Š” ๋ชจ๋“  GPU๊ฐ€ ๊ฐ™์€ ๋ ˆ์ด์–ด๋ฅผ ๋ถ„์‚ฐ ๊ณ„์‚ฐํ•˜๋Š” ๋ฐฉ์‹์ด๋ฏ€๋กœ, weight๋ฅผ TP ๊ธฐ์ค€์œผ๋กœ ์žฌ๋ฐฐ์น˜ํ•ด์•ผ ํ•จ
  • ๊ธฐ์กด PP์—์„œ ๋ถ„ํ• ๋œ weight๋ฅผ CPU์—์„œ ๋‹ค์‹œ ๋ถˆ๋Ÿฌ์™€ GPU์— TP ๋ฐฉ์‹์œผ๋กœ ๋กœ๋“œ

โ‘ฃ Decode ๋‹จ๊ณ„: Tensor Parallelism ์ ์šฉ

  • KV ์บ์‹œ๋Š” CPU๋กœ๋ถ€ํ„ฐ GPU๋กœ ๋น„๋™๊ธฐ์ ์œผ๋กœ ๋กœ๋“œ๋จ (prefetch thread ์‚ฌ์šฉ)
  • GPU๋“ค์€ ๋™์ผํ•œ layer ์ „์ฒด๋ฅผ ๋ณ‘๋ ฌ๋กœ ์ฒ˜๋ฆฌํ•˜๋ฉฐ ํ† ํฐ์„ ํ•œ ๊ฐœ์”ฉ ์ƒ์„ฑ (auto-regressive)
  • TP๋Š” weight๊ฐ€ GPU๊ฐ„ ๋ถ„์‚ฐ๋˜์–ด ์žˆ์–ด ํ•œ ๋ฒˆ๋งŒ ๋กœ๋”ฉํ•ด๋„ ์—ฌ๋Ÿฌ ์‹œํ€€์Šค ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅ
  • ๋Œ€๋Ÿ‰์˜ ์‹œํ€€์Šค๋ฅผ ๋™์‹œ์— ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅํ•˜์—ฌ throughput ๊ทน๋Œ€ํ™”
PLAINTEXT
[GPU1] --- TP โ†’ Layer Computation (1/4)
[GPU2] --- TP โ†’ Layer Computation (1/4)
[GPU3] --- TP โ†’ Layer Computation (1/4)
[GPU4] --- TP โ†’ Layer Computation (1/4)

โ†’ CPU์—์„œ KV shard๋ฅผ ๋ฐ›์•„์„œ TP ๋ฐฉ์‹์œผ๋กœ Decode ์ˆ˜ํ–‰


โ‘ค Transition-Minimizing Scheduling

  • ์ „ํ™˜ ์กฐ๊ฑด: CPU KV cache๊ฐ€ ๊ฐ€๋“ ์ฐผ์„ ๋•Œ๋งŒ prefill โ†’ decode๋กœ ์ „ํ™˜
  • decode ์ค‘์—๋Š” prefetch thread๊ฐ€ ๋ฏธ๋ฆฌ KV๋ฅผ ์˜ฎ๊ฒจ ๋†“์Œ โ†’ compute์™€ comm overlap
  • ์ด๋ ‡๊ฒŒ ํ•˜๋ฉด frequent reshards๋ฅผ ๋ฐฉ์ง€ํ•˜๋ฉด์„œ๋„ decode throughput์„ ์œ ์ง€

๐Ÿ“Œ ์ „์ฒด ๋™์ž‘ ํ๋ฆ„ ์š”์•ฝ ๋‹ค์ด์–ด๊ทธ๋žจ

PLAINTEXT
[์‹œ์ž‘] โ†’ Prefill(PP) โ†’ KV Cache(CPU์— ์ €์žฅ) โ†’ Model Re-shard(TP) โ†’ Decode(TP) โ†’ ๋
            โ†‘                                                   โ†“
            โ””โ”€โ”€โ”€โ”€โ”€โ”€โ”€ Transition-minimizing (๋ฒ„ํผ ๊ฝ‰ ์ฐจ๋ฉด ์ „ํ™˜) โ”€โ”˜

๐Ÿง  ์ •๋ฆฌ

๋‹จ๊ณ„๋ณ‘๋ ฌํ™” ์ „๋žต์žฅ์ ๋‹จ์ 
PrefillPipeline Parallelism๋‚ฎ์€ ํ†ต์‹  ๋น„์šฉ, ํ† ํฐ ๋ณ‘๋ ฌ ์ฒ˜๋ฆฌmicro-batch overhead
DecodeTensor Parallelism๋น ๋ฅธ weight loading, ๋†’์€ throughputAllReduce ํ†ต์‹  ๋น„์šฉ
KV ์ €์žฅTiered BufferingGPU ๋ฉ”๋ชจ๋ฆฌ ํ•œ๊ณ„ ๊ทน๋ณตCPUโ†”GPU ์ „์†ก ์†๋„ ๋ณ‘๋ชฉ
์Šค์ผ€์ค„๋งTransition-Minimizedre-shard ํšŸ์ˆ˜ ์ตœ์†Œํ™”๋ณต์žกํ•œ ์Šค์ผ€์ค„๋Ÿฌ ๋กœ์ง ํ•„์š”

์‹คํ—˜ ๊ฒฐ๊ณผ์™€ ๋‹ค๋ฅธ ๋ฐฉ๋ฒ•๋ก ๊ณผ์˜ ์ฐจ์ด

๋ฌผ์–ด๋ณธ ๊ฒƒ: ์ด ๋…ผ๋ฌธ์—์„œ ์ œ์‹œํ•œ ๊ฒฐ๊ณผ๋ฅผ ์ž์„ธํ•˜๊ฒŒ ๋ณด๊ณ  ๋‹ค๋ฅธ ๋ฐฉ๋ฒ•๋ก ์— ๋น„ํ•˜์—ฌ ํŠน์ถœ๋‚œ ์ ๊ณผ ๋…ผ๋ฌธ์—์„œ ์ œ๊ธฐํ•˜๋Š” ์–ด๋– ํ•œ ๋ฐฉ๋ฒ•์ด ์ด๋Ÿฌํ•œ ๊ฒฐ๊ณผ๋ฅผ ๋„์ถœํ•˜๊ฒŒ ๋˜์—ˆ๋Š”์ง€ ๋…ผ๋ฌธ์—์„œ ์ œ์‹œํ•˜๋Š” ์ด์œ ์™€ ๋„ˆ์˜ ์ƒ๊ฐ์„ ์•Œ๋ ค์ค˜

๋…ผ๋ฌธ Seesaw: High-throughput LLM Inference via Model Re-sharding์€ ๋‹ค์–‘ํ•œ ํ•˜๋“œ์›จ์–ด ๋ฐ ๋ชจ๋ธ ์„ค์ •์—์„œ SOTA ์ถ”๋ก  ์—”์ง„์ธ vLLM ๋Œ€๋น„ ์ตœ๋Œ€ 1.78ร—, ํ‰๊ท  1.36ร—์˜ throughput ํ–ฅ์ƒ์„ ๋ณด์—ฌ์ค๋‹ˆ๋‹ค. ์ด ์„ฑ๊ณผ๋Š” ๋‹จ์ผ ๋ณ‘๋ ฌํ™” ์ „๋žต์— ๊ณ ์ •๋œ ๊ธฐ์กด ์—”์ง„๋“ค๊ณผ ์ฐจ๋ณ„๋˜๋Š” ๋™์  ๋ณ‘๋ ฌํ™” ์ „๋žต ์กฐํ•ฉ (PP + TP)๊ณผ transition-aware scheduling ๋•๋ถ„์ž…๋‹ˆ๋‹ค.


๐Ÿ“Š Seesaw ์„ฑ๋Šฅ ๊ฒฐ๊ณผ ์ •๋ฆฌ (๋…ผ๋ฌธ 6์žฅ ๊ธฐ์ค€)

๋ชจ๋ธGPUvLLM ThroughputSeesaw Throughput๊ฐœ์„  ๋ฐฐ์ˆ˜์„ค๋ช…
LLaMA3-15B4 ร— A101.0 (baseline)1.45+45%chunked prefill ์กฐํ•ฉ๋ณด๋‹ค ์šฐ์ˆ˜
CodeLLaMA-34B8 ร— A101.01.78+78%prefill/decoding stage ๋ถ„๋ฆฌ๋กœ ์ด๋“ ๊ทน๋Œ€ํ™”
LLaMA2-70B8 ร— A100 PCIe1.01.46+46%TP์—์„œ ๋ฐœ์ƒํ•˜๋Š” all-reduce ๋ณ‘๋ชฉ ์™„ํ™”
LLaMA2-70B8 ร— A100 NVLink1.01.13+13%NVLink ํ™˜๊ฒฝ์—์„œ๋„ ์—ฌ์ „ํžˆ ๊ฐœ์„  ํ™•์ธ

โ€ป chunked prefill์„ ์‚ฌ์šฉํ•œ vLLM (TP2+PP2) ๋Œ€๋น„ Seesaw๊ฐ€ ์šฐ์ˆ˜ํ•˜๋‹ค๋Š” ์ ์€ ๋ˆˆ์—ฌ๊ฒจ๋ณผ ํฌ์ธํŠธ์ž…๋‹ˆ๋‹ค.


๐Ÿ† ํŠน์ถœ๋‚œ ์  (Compared to other methods)

๋น„๊ต ๋Œ€์ƒํ•œ๊ณ„Seesaw์˜ ์šฐ์ˆ˜ํ•œ ์ 
vLLM๋‹จ์ผ ๋ณ‘๋ ฌํ™” ์ „๋žต๋งŒ ์‚ฌ์šฉ ๊ฐ€๋Šฅ (ex. TP-only, PP-only)stage๋ณ„ ๋ณ‘๋ ฌํ™” ์ „๋žต ๋™์  ๋ณ€๊ฒฝ (PP โ†’ TP)
DistServe / Mooncake (๋ถ„๋ฆฌ ๋ฐฐ์น˜ ๊ธฐ๋ฐ˜)prefill/decode ์ฒ˜๋ฆฌ๋Ÿ‰ ๋ถˆ๊ท ํ˜•, resource ๋‚ญ๋น„, GPU duplication๊ฐ™์€ ๋ฆฌ์†Œ์Šค ๋‚ด์—์„œ stage ๊ฐ„ ๋ณ‘๋ ฌํ™” ๋ฐฉ์‹ ์ „ํ™˜
Chunked Prefill (Sarathi, DeepSpeed-FastGen)chunk ์‚ฌ์ด decode-only ๋‹จ๊ณ„ ๋ฐœ์ƒ, ์ตœ์  chunk size ์ฐพ๊ธฐ ์–ด๋ ค์›€prefill๊ณผ decode๋ฅผ ์™„์ „ํžˆ ๋ถ„๋ฆฌํ•˜์—ฌ ๋” ํฐ batching ๊ฐ€๋Šฅ
๊ธฐ์กด hybrid parallelism๋ณ‘๋ ฌํ™” ๋ฐฉ์‹ ์กฐํ•ฉ์€ ๊ณ ์ •์ ์ด๊ณ  stage-aware๊ฐ€ ์•„๋‹˜Seesaw๋Š” ๊ฐ stage๋ณ„ ์ตœ์  ๋ณ‘๋ ฌ ์ „๋žต์„ ๋™์ ์œผ๋กœ ์„ ํƒ

๐Ÿ“Œ ์„ฑ๋Šฅ ๊ฐœ์„ ์„ ์œ ๋„ํ•œ Seesaw์˜ ๋ฐฉ๋ฒ•๋ก ๊ณผ ๋…ผ๋ฌธ ๋‚ด ๊ทผ๊ฑฐ

1. Dynamic Model Re-sharding

  • ๊ทผ๊ฑฐ: prefill์€ communication-bound โ†’ pipeline์ด ์œ ๋ฆฌ, decode๋Š” weight-loading-bound โ†’ tensor๊ฐ€ ์œ ๋ฆฌ
  • ๋…ผ๋ฌธ ์‹คํ—˜: Figure 1, Figure 3, 12์—์„œ ๋ณด์—ฌ์ง€๋“ฏ ๊ฐ๊ฐ์˜ ๋ณ‘๋ ฌํ™” ๋ฐฉ์‹์ด stage์— ๋”ฐ๋ผ ์œ ๋ถˆ๋ฆฌ๊ฐ€ ๋‹ค๋ฆ„
  • ๋‚ด ์ƒ๊ฐ: ๊ธฐ์กด LLM inference ํ”„๋ ˆ์ž„์›Œํฌ๋Š” ํ†ต์ผ๋œ ๋ณ‘๋ ฌํ™” ๋ฐฉ์‹์„ ์‚ฌ์šฉํ•ด suboptimal ์ „๋žต์œผ๋กœ ์„ฑ๋Šฅ์„ ๋‚ญ๋น„ํ–ˆ์Œ. Seesaw๋Š” ์ด๋Ÿฐ โ€œ๋ณ‘๋ ฌํ™” ์ „๋žต ์„ ํƒ ๋ฌธ์ œโ€๋ฅผ stage-aware ๊ด€์ ์—์„œ ์žฌ์ •์˜ํ•œ ์ ์ด ํ˜์‹ ์ .

2. Tiered KV Cache Buffering

  • ๊ทผ๊ฑฐ: GPU ๋ฉ”๋ชจ๋ฆฌ ํ•œ๊ณ„๋ฅผ CPU ๋ฉ”๋ชจ๋ฆฌ๋กœ offload ํ•˜์—ฌ, prefill ์‹œ ๋Œ€๊ทœ๋ชจ ์‹œํ€€์Šค๋ฅผ ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅ
  • ๋…ผ๋ฌธ ์‹คํ—˜: Figure 6, 7์—์„œ Tiered KV ์บ์‹œ ๊ตฌ์กฐ ๋ฐ swap-in/out ๋™์ž‘ ์„ค๋ช…. GPU์— decode ์บ์‹œ๋ฅผ ๋ฏธ๋ฆฌ prefetchํ•จ
  • ๋‚ด ์ƒ๊ฐ: ์ด ๊ตฌ์กฐ ๋•๋ถ„์— GPU์—์„œ โ€œprefill batching โ†’ decode batchingโ€ ๊ฐ„์— ๋ณ‘๋ชฉ ์—†์ด ๋ถ€๋“œ๋Ÿฝ๊ฒŒ ๋„˜์–ด๊ฐˆ ์ˆ˜ ์žˆ์Œ. ํŠนํžˆ PCIe ํ™˜๊ฒฝ์—์„œ๋Š” ํ•„์ˆ˜์ 

3. Transition-minimizing Scheduling

  • ๊ทผ๊ฑฐ: frequent reshards๋Š” ์˜ค๋ฒ„ํ—ค๋“œ ์œ ๋ฐœ โ†’ decode ์ „ํ™˜์€ CPU ๋ฒ„ํผ๊ฐ€ ๊ฐ€๋“ ์ฐผ์„ ๋•Œ๋งŒ ๋ฐœ์ƒ
  • ๋…ผ๋ฌธ ์‹คํ—˜: Figure 2(a~c), ablation ์‹คํ—˜ (Section 6.3)์—์„œ re-shard ํšŸ์ˆ˜ ์ค„์ด๋ฉด์„œ๋„ decode throughput ์œ ์ง€๋จ
  • ๋‚ด ์ƒ๊ฐ: ๊ธฐ์กด scheduling์€ latency ์ตœ์†Œํ™”๋ฅผ ์ค‘์‹ฌ์œผ๋กœ ์„ค๊ณ„๋˜์–ด throughput์—๋Š” ์ ํ•ฉํ•˜์ง€ ์•Š์•˜์Œ. Seesaw๋Š” ์ด๋ฅผ ์˜คํžˆ๋ ค throughput-maximizing ๊ด€์ ์—์„œ ์žฌ์„ค๊ณ„ํ•œ ์ ์ด ๊ฐ•๋ ฅํ•œ ์ฐจ๋ณ„์ .

๐Ÿง  ์ถ”๊ฐ€ ๋ถ„์„ ๋ฐ ํ•œ๊ณ„

  • Seesaw๋Š” โ€˜offline batch inferenceโ€™์—๋Š” ํƒ์›”ํ•˜์ง€๋งŒ โ€˜interactive onlineโ€™ inference์—๋Š” latency ์ฆ๊ฐ€ ์œ„ํ—˜ ์กด์žฌ

    • decode batch๊ฐ€ ์ถฉ๋ถ„ํžˆ ์ปค์•ผ ์„ฑ๋Šฅ ์ด์ ์ด ๊ทน๋Œ€ํ™”๋จ
  • re-shard์˜ CPUโ†”GPU ์ „์†ก์ด NVLink ์—†์„ ๊ฒฝ์šฐ bottleneck์ด ๋  ์ˆ˜๋„ ์žˆ์Œ

    • ๋…ผ๋ฌธ์—์„œ๋„ PCIe vs NVLink ์„ฑ๋Šฅ ์ฐจ์ด๋ฅผ ๋ถ„๋ช…ํžˆ ์ œ์‹œ (A100 ์‹คํ—˜์—์„œ 46% โ†’ 13%๋กœ ๊ฐ์†Œ)

๐Ÿ”š ๊ฒฐ๋ก 

Seesaw๋Š” “LLM ์ถ”๋ก  ๋‹จ๊ณ„๋ณ„๋กœ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋งž์ถค ์ตœ์ ํ™”"ํ•œ๋‹ค๋Š” ๊ฐœ๋…์„ ์ตœ์ดˆ๋กœ ๊ณ ์•ˆํ•˜๊ณ , ์ด๋ฅผ tiered memory + intelligent scheduler์™€ ๊ฒฐํ•ฉํ•˜์—ฌ ๊ธฐ์กด ์ถ”๋ก  ์‹œ์Šคํ…œ๋“ค์˜ ๋‹จ์ ์„ ํ•ด๊ฒฐํ•œ ์„ค๊ณ„์ž…๋‹ˆ๋‹ค. ํŠนํžˆ ๋Œ€๊ทœ๋ชจ ์˜คํ”„๋ผ์ธ ์ฒ˜๋ฆฌ, resource constraint ํ™˜๊ฒฝ์—์„œ ์‹ค์งˆ์ ์ธ ์„ฑ๋Šฅ ๊ฐœ์„ ์„ ์ž…์ฆํ•œ ์ ์ด ์‚ฐ์—…์  ๊ฐ€์น˜๊ฐ€ ํฝ๋‹ˆ๋‹ค.

ํ•œ ์ค„ ์š”์•ฝ: ๋‹จ์ˆœํ•œ ๋ณ‘๋ ฌํ™” ์กฐํ•ฉ์ด ์•„๋‹ˆ๋ผ, โ€œ์–ธ์ œโ€ ๊ทธ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ์“ฐ๋А๋ƒ๊ฐ€ ์ค‘์š”ํ•˜๋‹ค๋Š” ์ ์„ ์‹คํ—˜๊ณผ ์ด๋ก  ๋ชจ๋‘์—์„œ ๊ฐ•๋ ฅํžˆ ์ž…์ฆํ•œ ๋…ผ๋ฌธ์ž…๋‹ˆ๋‹ค.

์ž…๋ ฅ ๋ฐ์ดํ„ฐ์™€ ์ถ”๋ก  ๊ณผ์ •, ๋ชจ๋ธ ๊ตฌ์กฐ์™€ ์ž์› ์š”๊ตฌ๋Ÿ‰

๋ฌผ์–ด๋ณธ ๊ฒƒ: ์ด ๋…ผ๋ฌธ์˜ ์ž…๋ ฅ๋ฐ์ดํ„ฐ์™€ ์ถ”๋ก  ๊ณผ์ •์— ๋Œ€ํ•ด์„œ ์˜ˆ์‹œ๋ฅผ ๋“ค์–ด ์•„์ฃผ ์ž์„ธํ•˜๊ฒŒ ์„ค๋ช…ํ•ด ์ฃผ๊ฒ ๋‹ˆ? ์ถ”๊ฐ€์ ์œผ๋กœ ๋ชจ๋ธ์•„ํ‚คํ…์ฒ˜์˜ ๊ตฌ์„ฑ ๋ฐ ๋ชจ๋ธ์„ ๊ตฌ์„ฑํ•˜๋Š” ์—ฐ์‚ฐ๊ณผ ๋ฉ”๋ชจ๋ฆฌ ์š”๊ตฌ๋Ÿ‰ ์ปดํ“จํŒ… ์š”๊ตฌ๋Ÿ‰ ๋Œ€ํ•ด์„œ๋„ ์„ค๋ช…ํ•ด์ค˜

์ข‹์Šต๋‹ˆ๋‹ค. ์•„๋ž˜๋Š” ๋…ผ๋ฌธ Seesaw์˜ ์ „์ฒด ์ž…๋ ฅโ€“์ถ”๋ก โ€“๋ชจ๋ธ ๊ตฌ์กฐโ€“๋ฆฌ์†Œ์Šค ์‚ฌ์šฉ์— ๋Œ€ํ•œ ์ƒ์„ธ ์„ค๋ช…์ž…๋‹ˆ๋‹ค. ์˜ˆ์‹œ์™€ ํ•จ๊ป˜ ๋‹จ๊ณ„๋ณ„๋กœ ๊ตฌ์„ฑํ–ˆ์Šต๋‹ˆ๋‹ค.


๐Ÿงช 1. ์ž…๋ ฅ ๋ฐ์ดํ„ฐ ๊ตฌ์กฐ ๋ฐ ์˜ˆ์‹œ

๋…ผ๋ฌธ์—์„œ๋Š” ๋‘ ๊ฐ€์ง€ ๋Œ€ํ‘œ์ ์ธ throughput-oriented inference workload๋ฅผ ์‚ฌ์šฉํ–ˆ์Šต๋‹ˆ๋‹ค:

๋ฐ์ดํ„ฐ์…‹์„ค๋ช…Input ๊ธธ์ดOutput ๊ธธ์ด
ShareGPT๋Œ€ํ™” ๊ธฐ๋ก ๊ธฐ๋ฐ˜32~256 ํ† ํฐ32~512 ํ† ํฐ
arXiv summarization๋…ผ๋ฌธ ์š”์•ฝ ํƒœ์Šคํฌ์ตœ๋Œ€ 3000 ํ† ํฐ32~128 ํ† ํฐ

๐Ÿ“Œ ์˜ˆ์‹œ ์ž…๋ ฅ (arXiv summarization ๊ธฐ๋ฐ˜)

JSON
{
  "input": "In this paper, we propose a novel model...",
  "max_output_tokens": 128
}

โ†’ ํ•˜๋‚˜์˜ sequence๊ฐ€ ๋งŒ๋“ค์–ด์ง€๊ณ , inference ์‹œ์Šคํ…œ์—๋Š” ๋‹ค์Œ๊ณผ ๊ฐ™์ด ์š”์ฒญ์ด ๋“ค์–ด๊ฐ‘๋‹ˆ๋‹ค:

  • prompt length = 2048 tokens
  • max decode length = 128 tokens
  • ์ด 1 sequence = 2048 + 128 = 2176 tokens (๋‹จ, decode๋Š” 1 token์”ฉ step-by-step ์ƒ์„ฑ)

๐Ÿง  2. ์ถ”๋ก (inference) ๊ณผ์ •์˜ ๋‹จ๊ณ„๋ณ„ ์ฒ˜๋ฆฌ

Seesaw๋Š” inference๋ฅผ ๋‹ค์Œ ๋‘ ๋‹จ๊ณ„๋กœ ๋‚˜๋ˆ„์–ด ์ฒ˜๋ฆฌํ•ฉ๋‹ˆ๋‹ค.

โ–ถ Step 1: Prefill ๋‹จ๊ณ„

  • ๋ชจ๋“  ์ž…๋ ฅ ํ† ํฐ(2048๊ฐœ)์„ ํ•œ ๋ฒˆ์— transformer๋กœ forward pass ํ•จ.
  • ๊ฐ attention layer์—์„œ Query, Key, Value(QKV)๋ฅผ ๊ณ„์‚ฐํ•˜๊ณ , KV๋Š” ์บ์‹œ๋กœ ์ €์žฅ๋จ.
  • ๋ณ‘๋ ฌํ™” ์ „๋žต: Pipeline Parallelism (PP) ์ ์šฉ โ†’ ๋ ˆ์ด์–ด๋ฅผ ์—ฌ๋Ÿฌ GPU์— ๋ถ„์‚ฐ.

โ–ถ Step 2: Decode ๋‹จ๊ณ„

  • ํ† ํฐ์„ 1๊ฐœ์”ฉ ๋ฐ˜๋ณต์ ์œผ๋กœ ์ƒ์„ฑํ•จ (auto-regressive).
  • ์ด์ „ ๋‹จ๊ณ„์—์„œ ์ €์žฅํ•œ KV ์บ์‹œ๋ฅผ ์‚ฌ์šฉํ•˜์—ฌ Q์™€ Attention๋งŒ ๊ณ„์‚ฐ
  • ๋ณ‘๋ ฌํ™” ์ „๋žต: Tensor Parallelism (TP) ์ ์šฉ โ†’ weight sharding ๋ฐ AllReduce.

๐Ÿ”„ ์ „์ฒด ์˜ˆ์‹œ ์ถ”๋ก  ํ๋ฆ„

TEXT
์ž…๋ ฅ: 2048 ํ† ํฐ ๋…ผ๋ฌธ ํ…์ŠคํŠธ โ†’ [Prefill with PP4]
โ†’ ๊ฐ GPU๋Š” ์ผ๋ถ€ ๋ ˆ์ด์–ด๋งŒ ๋‹ด๋‹น (ex. 4๊ฐœ ๋ ˆ์ด์–ด ์ค‘ GPU1: L1, GPU2: L2 ...)
โ†’ Q, K, V ๊ณ„์‚ฐ โ†’ KV๋Š” CPU์— ์ €์žฅ

โ†’ [Transition & re-shard to TP4]
โ†’ Decode ์‹œ์ž‘: ํ† ํฐ 1๊ฐœ ์ƒ์„ฑ โ†’ ๋‹ค์Œ ํ† ํฐ ์ƒ์„ฑ ...
โ†’ ๋ฐ˜๋ณต์ ์œผ๋กœ CPU์—์„œ KV๋ฅผ GPU๋กœ ๊ฐ€์ ธ์™€ decode

๐Ÿ—๏ธ 3. ๋ชจ๋ธ ์•„ํ‚คํ…์ฒ˜ ๊ตฌ์„ฑ

๋…ผ๋ฌธ ์‹คํ—˜์—๋Š” ๋‹ค์Œ ๋ชจ๋ธ์ด ์‚ฌ์šฉ๋ฉ๋‹ˆ๋‹ค:

๋ชจ๋ธ๋ช…ํฌ๊ธฐ์•„ํ‚คํ…์ฒ˜
LLaMA315BGQA ๊ธฐ๋ฐ˜ decoder-only
CodeLLaMA34BGQA ๊ธฐ๋ฐ˜ decoder-only
LLaMA270BGQA ๊ธฐ๋ฐ˜ decoder-only

๋ชจ๋‘ transformer decoder-only ๊ตฌ์กฐ๋กœ ๋˜์–ด ์žˆ์œผ๋ฉฐ, ํ•ต์‹ฌ ๊ตฌ์„ฑ์€ ์•„๋ž˜์™€ ๊ฐ™์Šต๋‹ˆ๋‹ค:

๐Ÿ”น ํ•œ Layer ๊ตฌ์„ฑ

  1. LayerNorm
  2. Multi-head Attention (QKV ๊ณ„์‚ฐ, GQA ๊ตฌ์กฐ)
  3. Residual Add
  4. FeedForward Network (2~4ร— hidden size)
  5. ๋˜ ๋‹ค๋ฅธ LayerNorm + Residual

โฑ๏ธ 4. ์—ฐ์‚ฐ๋Ÿ‰ ๋ฐ ๋ฉ”๋ชจ๋ฆฌ ์š”๊ตฌ๋Ÿ‰ ๋ถ„์„

1) Prefill ๋‹จ๊ณ„:

  • ์—ฐ์‚ฐ๋Ÿ‰:

    • Attention: O(nยฒd) โ†’ n์€ sequence length, d๋Š” hidden dim
    • FFN: O(nยทdยฒ)
  • ๋ฉ”๋ชจ๋ฆฌ:

    • Activations + KV Cache ๋ชจ๋‘ ํ•„์š”
    • GPU๋‹น model weight shard + ์ค‘๊ฐ„ ํ™œ์„ฑ๊ฐ’ ์ €์žฅ
  • ๋ณ‘๋ชฉ:

    • Communication (AllReduce)๋ณด๋‹ค computation์ด ์šฐ์„ธ (์ด์œ : ๊ธด ์‹œํ€€์Šค)
    • โ†’ PP๊ฐ€ ์œ ๋ฆฌ

2) Decode ๋‹จ๊ณ„:

  • ์—ฐ์‚ฐ๋Ÿ‰:

    • Attention: O(nยทd) (n์€ ๊ณผ๊ฑฐ ํ† ํฐ ์ˆ˜ โ†’ ์ฆ๊ฐ€)
    • FFN: O(dยฒ)
  • ๋ฉ”๋ชจ๋ฆฌ:

    • GPU ๋ฉ”๋ชจ๋ฆฌ๋‹น KV Cache ์ €์žฅ ํ•„์š”
    • Activations๋Š” ๋งค์šฐ ์ž‘์Œ (1 step)
  • ๋ณ‘๋ชฉ:

    • weight loading ๋ฐ broadcasting
    • โ†’ TP๊ฐ€ ์œ ๋ฆฌ (shard๋กœ ๋‚˜๋ˆ ์„œ weight ์‚ฌ์šฉ ๊ฐ€๋Šฅ)

๐Ÿ“ฆ ์‹ค์ œ ๊ณ„์‚ฐ ์˜ˆ (LLaMA2-70B ๊ธฐ์ค€)

  • ๋ชจ๋ธ ํŒŒ๋ผ๋ฏธํ„ฐ ์ˆ˜: 70B
  • hidden dim: ์•ฝ 8192
  • layer ์ˆ˜: 80
  • FFN ๋น„์œจ: 4ร— โ†’ FFN dim โ‰ˆ 32768

โ›ฝ ์ถ”๋ก  ์—ฐ์‚ฐ๋Ÿ‰ (1๊ฐœ ์‹œํ€€์Šค)

  • Prefill (2048 ํ† ํฐ):

    • Attention: 2048ยฒร—8192 โ‰ˆ 34B FLOPs/layer
    • FFN: 2048ร—8192ร—32768 โ‰ˆ 549B FLOPs/layer
    • ์ „์ฒด: (34B + 549B) ร— 80 layers = 46.8 TFLOPs
  • Decode (128 steps):

    • ๊ฐ step์—์„œ Attention + FFN = ~6.5B FLOPs/layer
    • ์ „์ฒด: ~6.5B ร— 128 ร— 80 = 66.5 TFLOPs

โ†’ decode๋Š” step-by-step์ด์ง€๋งŒ ์ „์ฒด FLOP ํ•ฉ์‚ฐ์€ prefill๋ณด๋‹ค ๋งŽ์Œ โ†’ ํ•˜์ง€๋งŒ ๋ณ‘๋ ฌ์„ฑ์ด ๋‚ฎ์•„, ์‹ค์ œ throughput์—์„  prefill์ด ๋น ๋ฆ„, decode๊ฐ€ ๋ณ‘๋ชฉ


๐Ÿงพ ์ •๋ฆฌ ์š”์•ฝ

ํ•ญ๋ชฉPrefillDecode
๋ณ‘๋ ฌํ™” ์ „๋žตPipeline (PP)Tensor (TP)
์ฒ˜๋ฆฌ ๋‹จ์œ„Prompt ์ „์ฒด (์ˆ˜์ฒœ ํ† ํฐ)1 ํ† ํฐ ๋ฐ˜๋ณต
๊ณ„์‚ฐ๋Ÿ‰๋งค์šฐ ํผ (O(nยฒ))์ž‘์Œ (O(nd))
๋ณ‘๋ชฉํ†ต์‹  (AllReduce)weight loading
์บ์‹œ์—†์Œ โ†’ ์ƒ์„ฑKV cache ํ™œ์šฉ
์ตœ์ ํ™”PP๊ฐ€ ํšจ์œจ์ TP๊ฐ€ ํšจ์œจ์ 

โœ… ๊ฒฐ๋ก 

์ด ๋…ผ๋ฌธ์—์„œ ์ฒ˜๋ฆฌํ•˜๋Š” ์ž…๋ ฅ ๋ฐ์ดํ„ฐ๋Š” ํ‰๊ท ์ ์œผ๋กœ ์ˆ˜์ฒœ ํ† ํฐ์˜ prompt + ์ˆ˜๋ฐฑ ํ† ํฐ์˜ ์ƒ์„ฑ์œผ๋กœ ๊ตฌ์„ฑ๋˜์–ด ์žˆ์œผ๋ฉฐ, ์ด ๊ตฌ์กฐ์  ํŠน์„ฑ์„ ๋ฐ˜์˜ํ•˜์—ฌ prefill๊ณผ decode์˜ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋‹ค๋ฅด๊ฒŒ ์ ์šฉํ•˜๋Š” ๊ฒƒ์ด ํ•ต์‹ฌ์ž…๋‹ˆ๋‹ค. Seesaw๋Š” ์ด ์ฐจ์ด๋ฅผ ์ •ํ™•ํžˆ ๋ถ„์„ํ•˜๊ณ , ๋™์  re-sharding + ๊ณ„์ธต์  ์บ์‹œ + ์Šค์ผ€์ค„๋ง์œผ๋กœ ์ด๋ฅผ ํšจ๊ณผ์ ์œผ๋กœ ๊ตฌํ˜„ํ–ˆ๊ธฐ์—, ๊ธฐ์กด static-parallelism ๋ฐฉ์‹ ๋Œ€๋น„ ํƒ์›”ํ•œ throughput ํ–ฅ์ƒ์„ ์ด๋ค˜์Šต๋‹ˆ๋‹ค.

ํ•œ๊ณ„๋ฅผ ๋„˜์œผ๋ ค๋Š” ์—ฐ๊ตฌ ํ๋ฆ„

๋ฌผ์–ด๋ณธ ๊ฒƒ: ์ด ๋…ผ๋ฌธ์˜ ํ•œ๊ณ„๋ฅผ ๊ทน๋ณตํ•˜๊ธฐ ์œ„ํ•œ ๋ฐฉ๋ฒ•์œผ๋กœ ์–ด๋–ค ์—ฐ๊ตฌํ๋ฆ„์ด ์žˆ๋Š”์ง€ ์ •๋ฆฌํ•ด์„œ ์ž์„ธํ•˜๊ฒŒ ์„ค๋ช…ํ•ด ์ค˜

๋…ผ๋ฌธ “Seesaw: High-throughput LLM Inference via Model Re-sharding”์€ ๊ณ ์ • ๋ณ‘๋ ฌํ™” ์ „๋žต์˜ ๋น„ํšจ์œจ์„ฑ์„ ๊ทน๋ณตํ–ˆ์ง€๋งŒ, ์—ฌ์ „ํžˆ ์‹œ์Šคํ…œ์ , ์•Œ๊ณ ๋ฆฌ์ฆ˜์  ํ•œ๊ณ„์ ์ด ์กด์žฌํ•ฉ๋‹ˆ๋‹ค. ์ด์— ๋”ฐ๋ผ ์ด ํ•œ๊ณ„๋ฅผ ๊ทน๋ณตํ•˜๋ ค๋Š” ํ›„์† ์—ฐ๊ตฌ ํ๋ฆ„์€ ํฌ๊ฒŒ 4๊ฐ€์ง€ ๋ฐฉํ–ฅ์œผ๋กœ ๋ฐœ์ „ํ•˜๊ณ  ์žˆ์Šต๋‹ˆ๋‹ค:


โœ… ์š”์•ฝ: ํ•œ๊ณ„์ ๊ณผ ๋Œ€์‘ ์—ฐ๊ตฌ ํ๋ฆ„

Seesaw์˜ ํ•œ๊ณ„์ ๋Œ€์‘ํ•˜๋Š” ์—ฐ๊ตฌ ํ๋ฆ„ํ•ต์‹ฌ ์•„์ด๋””์–ด
โ‘  ๋นˆ๋ฒˆํ•œ re-sharding overheadLayer-level Adaptive Re-sharding์„ธ๋ถ„ํ™”๋œ weight migration, ๋˜๋Š” re-shard-free execution
โ‘ก CPU โ†” GPU ๊ฐ„ KV cache ์ „์†ก ๋ณ‘๋ชฉKV Cache Compression & Prefetch Optimization์••์ถ• + ๋น„๋™๊ธฐ prefetch์˜ ์ตœ์ ํ™”
โ‘ข offline inference ์ „์šฉ โ†’ online์—๋Š” ๋ถ€์ ํ•ฉLatency-aware Scheduling & Prefetchingonline + offline hybrid ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅํ•œ ์Šค์ผ€์ค„๋Ÿฌ
โ‘ฃ statically configured ๋ณ‘๋ ฌ ์กฐํ•ฉReinforcement Learning ๊ธฐ๋ฐ˜ Auto-parallelism์‹œ์ ๋ณ„ ์ตœ์  ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ํ•™์Šต ๊ธฐ๋ฐ˜์œผ๋กœ ํƒ์ƒ‰

๐Ÿ” 1. Re-sharding Overhead ์™„ํ™”: Adaptive Sharding

Seesaw์˜ ํ•œ๊ณ„

  • stage ๊ฐ„ ์ „ํ™˜ ์‹œ model weight, KV cache๋ฅผ GPU ๊ฐ„ ์žฌ๋ฐฐ์น˜ํ•ด์•ผ ํ•˜๋ฉฐ,
  • ์ด๋Š” memory copy + data layout transformation์ด ํ•„์š”ํ•˜์—ฌ latency ์ฆ๊ฐ€

๋Œ€์‘ ์—ฐ๊ตฌ

  • Splitwise (ISCA 2024): phase splitting ๊ธฐ๋ฒ• ์ œ์•ˆ

    • ํ•œ ๋ฒˆ์˜ shard ์žฌ๋ฐฐ์น˜๋กœ ์—ฌ๋Ÿฌ ๋‹จ๊ณ„ ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅํ•˜๋„๋ก layer ๋ฐฐ์น˜๋ฅผ ์กฐ์ •
  • Dynamic Layer Migration: ์ž์ฃผ ๋ณ€ํ•˜์ง€ ์•Š๋Š” layer๋Š” ๊ณ ์ •, ๋‚˜๋จธ์ง€๋งŒ shard ์ด๋™

๋ฏธ๋ž˜ ๋ฐฉํ–ฅ

  • layer-level granularity์—์„œ sharding-free hybrid execution ์„ค๊ณ„
  • GPU-localํ•œ KV cache reuse ๋ฐฉ์‹ ์—ฐ๊ตฌ

๐Ÿ” 2. KV ์บ์‹œ ๋ฉ”๋ชจ๋ฆฌ ๋ณ‘๋ชฉ ๊ทน๋ณต: Compression & Scheduling

Seesaw์˜ ํ•œ๊ณ„

  • CPU์—์„œ GPU๋กœ KV cache๋ฅผ ์˜ฎ๊ธธ ๋•Œ PCIe bandwidth (16GB/s ์ˆ˜์ค€)๊ฐ€ ๋ณ‘๋ชฉ
  • ๋Œ€๋Ÿ‰ ์‹œํ€€์Šค๋ฅผ ํ•œ๋ฒˆ์— decode ์‹œ prefetch๊ฐ€ ์ œ๋•Œ ๋„์ฐฉํ•˜์ง€ ์•Š์œผ๋ฉด stall ๋ฐœ์ƒ

๋Œ€์‘ ์—ฐ๊ตฌ

  • FastDecoder (2024): 8-bit quantization์„ ์ ์šฉํ•œ KV ์บ์‹œ ์••์ถ•
  • NanoFlow (2024): swap-in/swap-out ํƒ€์ด๋ฐ์„ token-level prefetch๋กœ ์„ธ๋ถ„ํ™”
  • DistAttention (2024): multi-level attention cache hierarchy ์ œ์•ˆ (L2-like KV cache)

๋ฏธ๋ž˜ ๋ฐฉํ–ฅ

  • LRU ๊ธฐ๋ฐ˜ cache eviction ์ •์ฑ…์ด ์ ์šฉ๋œ GPU KV cache ๊ด€๋ฆฌ์ž
  • ์••์ถ•๋ฅ  โ†” latency tradeoff์— ์ตœ์ ํ™”๋œ adaptive ์••์ถ• ์•Œ๊ณ ๋ฆฌ์ฆ˜

๐Ÿ” 3. ์˜จ๋ผ์ธ/์ธํ„ฐ๋ž™ํ‹ฐ๋ธŒ ํ™˜๊ฒฝ ๋ถ€์ ํ•ฉ ๋ฌธ์ œ: Hybrid Scheduling

Seesaw์˜ ํ•œ๊ณ„

  • offline-only ์‹œ์Šคํ…œ์œผ๋กœ ์„ค๊ณ„๋จ โ†’ prompt latency ๊ณ ๋ ค ์•ˆ ํ•จ
  • interactive workload (ex. chatbot)์˜ ๊ฒฝ์šฐ delay๊ฐ€ ์ฆ๊ฐ€ํ•  ์ˆ˜ ์žˆ์Œ

๋Œ€์‘ ์—ฐ๊ตฌ

  • Sarathi-Serve (2024): chunked prefill๊ณผ decode๋ฅผ ์„ž๋Š” piggyback scheduling
  • Slice-level Scheduling (Cheng et al., 2024): GPU ๋‚ด idle slice ๋‹จ์œ„๋กœ prefill/decoding ๋ถ„๋ฐฐ

๋ฏธ๋ž˜ ๋ฐฉํ–ฅ

  • multi-objective scheduling (throughput + latency) ๊ฐ€๋Šฅํ•˜๋„๋ก heuristic ์Šค์ผ€์ค„๋Ÿฌ ๊ณ ๋„ํ™”
  • request priority ๊ธฐ๋ฐ˜ prompt-to-decode dispatch routing ์‹œ์Šคํ…œ

๐Ÿ” 4. ์ •์ ์ธ ๋ณ‘๋ ฌํ™” ๊ตฌ์„ฑ์˜ ํ•œ๊ณ„: Auto-parallelism ํƒ์ƒ‰

Seesaw์˜ ํ•œ๊ณ„

  • ๋ณ‘๋ ฌํ™” ์กฐํ•ฉ (PPx โ†’ TPy)์€ ์‚ฌ๋žŒ์ด ์‹คํ—˜์ ์œผ๋กœ ์ฐพ์Œ โ†’ ์ตœ์ ํ™” ๋ถˆ์™„์ „
  • workload distribution, token length, output ratio์— ๋”ฐ๋ผ ์ตœ์  ์กฐํ•ฉ์ด ๋ฐ”๋€œ

๋Œ€์‘ ์—ฐ๊ตฌ

  • AutoServe (๊ฐ€์นญ, ์•„์ง ๋ฏธ๋ฐœํ‘œ): workload ํŠน์„ฑ์„ ์‹ค์‹œ๊ฐ„์œผ๋กœ ๊ฐ์ง€ํ•˜์—ฌ ๋ณ‘๋ ฌ ์กฐํ•ฉ์„ ์กฐ์ •
  • RL ๊ธฐ๋ฐ˜ Controller: inference latency / throughput์„ reward๋กœ ํ•˜๋Š” ๋ณ‘๋ ฌํ™” ํƒ์ƒ‰๊ธฐ

๋ฏธ๋ž˜ ๋ฐฉํ–ฅ

  • token ๊ธธ์ด, memory usage ์˜ˆ์ธก โ†’ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋™์ ์œผ๋กœ ์ž๋™ ์ ์šฉํ•˜๋Š” system compiler + scheduler ํ†ตํ•ฉ

๐Ÿง  ์ข…ํ•ฉ ์ •๋ฆฌ

์นดํ…Œ๊ณ ๋ฆฌ๊ธฐ์กด ํ•œ๊ณ„ํ–ฅํ›„ ๋ณด์™„ ๊ฐ€๋Šฅ ๋ฐฉํ–ฅ
Re-shardingํ†ต์งธ shard ์ด๋™partial migration, layer-adaptive
KV cachePCIe ๋ณ‘๋ชฉ, ์ „์†ก ์ง€์—ฐ์••์ถ• + L2-style ์บ์‹œ ๊ณ„์ธตํ™”
Schedulingoffline only, prefill-decode ๊ณ ์ •latency-aware, priority-aware hybrid
๋ณ‘๋ ฌํ™” ์ „๋žต์ˆ˜๋™ ํŠœ๋‹AutoML ๋˜๋Š” RL ๊ธฐ๋ฐ˜ ๋™์  ์กฐ์ •

๐Ÿ“Œ ๋งˆ๋ฌด๋ฆฌ ์ƒ๊ฐ

Seesaw๋Š” ํ™•์‹คํ•œ ๊ตฌ์กฐ์  ๊ธฐ์—ฌ๋ฅผ ํ–ˆ์ง€๋งŒ, ์—ฌ์ „ํžˆ โ€œํ†ต์‹ /๋ฉ”๋ชจ๋ฆฌ ๋ณ‘๋ชฉโ€ ๋ฐ โ€œ์‹คํ–‰ ํ™˜๊ฒฝ ๋‹ค์–‘์„ฑโ€์—์„œ ์™„์ „ํ•œ ํ•ด๋ฒ•์€ ์•„๋‹™๋‹ˆ๋‹ค. ์œ„์—์„œ ์ •๋ฆฌํ•œ ํ›„์† ์—ฐ๊ตฌ๋“ค์€ Seesaw์˜ ํ•ต์‹ฌ ์„ค๊ณ„๋ฅผ ๊ธฐ๋ฐ˜์œผ๋กœ ๋” ์ •๊ตํ•˜๊ณ  ๋ฒ”์šฉ์ ์ธ ๊ณ ์„ฑ๋Šฅ ์ถ”๋ก  ์‹œ์Šคํ…œ์œผ๋กœ ํ™•์žฅํ•˜๊ณ ์ž ํ•˜๋Š” ํ๋ฆ„์ž…๋‹ˆ๋‹ค.

Q : Seesaw์—์„œ prefill๊ณผ decode ๋‹จ๊ณ„์˜ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๊ฐ๊ฐ ๋‹ค๋ฅด๊ฒŒ ์„ ํƒํ•˜๋Š” ๊ฒƒ์ด throughput ํ–ฅ์ƒ์— ์–ด๋–ป๊ฒŒ ๊ธฐ์—ฌํ•˜๋ฉฐ, ์ด๋ฅผ ์œ„ํ•ด ์–ด๋–ค ์กฐ๊ฑด์—์„œ pipeline parallelism ๋˜๋Š” tensor parallelism์ด ์œ ๋ฆฌํ•œ์ง€ ์ˆ˜์‹์ด๋‚˜ ์‹คํ–‰ ์˜ˆ์‹œ๋ฅผ ํ†ตํ•ด ์„ค๋ช…ํ•  ์ˆ˜ ์žˆ๋Š”๊ฐ€?

์ข‹์€ ์งˆ๋ฌธ์ž…๋‹ˆ๋‹ค. ์ด ์งˆ๋ฌธ์€ Seesaw์˜ ํ•ต์‹ฌ ์„ค๊ณ„์ฒ ํ•™์ธ “๋‹จ๊ณ„๋ณ„ ๋ณ‘๋ ฌํ™” ์ „๋žต ๋ถ„๋ฆฌ”๊ฐ€ ์™œ ํ•„์š”ํ•œ๊ฐ€, ๊ทธ๋ฆฌ๊ณ  ์–ด๋–ค ๊ฒฝ์šฐ์— TP ๋˜๋Š” PP๊ฐ€ ์œ ๋ฆฌํ•œ๊ฐ€๋ฅผ ์ˆ˜ํ•™์ ์ด๊ณ  ์‹œ์Šคํ…œ์ ์œผ๋กœ ์„ค๋ช…ํ•  ์ˆ˜ ์žˆ์–ด์•ผ ํ•ฉ๋‹ˆ๋‹ค.


โœ… ๊ฒฐ๋ก ๋ถ€ํ„ฐ ๋งํ•˜๋ฉด

Prefill ๋‹จ๊ณ„๋Š” pipeline parallelism(PP)์ด ์œ ๋ฆฌํ•˜๊ณ , Decode ๋‹จ๊ณ„๋Š” tensor parallelism(TP)์ด ์œ ๋ฆฌํ•ฉ๋‹ˆ๋‹ค.

Seesaw๋Š” ์ด ๋‘ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ stage๋ณ„๋กœ ๋‚˜๋ˆ ์„œ ์ ์šฉํ•˜๊ณ , ์ด๋ฅผ ํ†ตํ•ด ์ „๋ฐ˜์ ์ธ throughput์„ 1.36~1.78ร— ํ–ฅ์ƒ์‹œํ‚ต๋‹ˆ๋‹ค.


๐Ÿ” ์™œ ๋‹จ๊ณ„๋ณ„๋กœ ๋ณ‘๋ ฌํ™” ์ „๋žต์„ ๋‚˜๋ˆ ์•ผ ํ• ๊นŒ?

โ–ถ Prefill ๋‹จ๊ณ„์˜ ํŠน์ง•:

  • Prompt ์ „์ฒด ์‹œํ€€์Šค (n โ‰ˆ 512~3000)๋ฅผ ํ•œ ๋ฒˆ์— ์ฒ˜๋ฆฌ
  • Token ์ˆ˜๊ฐ€ ๋งŽ์•„์„œ ์—ฐ์‚ฐ๋Ÿ‰์ด O(nยฒ) (self-attention)
  • GPU compute ์ž์›์ด ๊ฝ‰ ์ฐจ๊ฒŒ ํ™œ์šฉ๋จ โ†’ communication์ด ๋ณ‘๋ชฉ

โ†’ โ‡’ compute-heavy + communication-bound โ†’ PP๊ฐ€ ์œ ๋ฆฌ

โ–ถ Decode ๋‹จ๊ณ„์˜ ํŠน์ง•:

  • 1 token์”ฉ ์ˆœ์ฐจ ์ƒ์„ฑ โ†’ ์—ฐ์‚ฐ๋Ÿ‰ ์ ๊ณ  O(n) ์ˆ˜์ค€
  • Token ์ˆ˜ ์ ๊ณ , ๋ฐ˜๋ณต์  โ†’ weight loading overhead๊ฐ€ ์ƒ๋Œ€์ ์œผ๋กœ ํผ
  • ์ž‘์€ ์—ฐ์‚ฐ๋Ÿ‰ โ†’ load โ†’ compute โ†’ idle โ†’ load ๋ฐ˜๋ณต๋จ

โ†’ โ‡’ memory-transfer-bound โ†’ TP๊ฐ€ ์œ ๋ฆฌ


๐Ÿงฎ ์ˆ˜์‹ ๊ธฐ๋ฐ˜ ์„ค๋ช… (๋…ผ๋ฌธ Appendix A ๊ธฐ๋ฐ˜)

๋…ผ๋ฌธ์—์„œ๋Š” ๋‹ค์Œ๊ณผ ๊ฐ™์€ throughput inverse ๋ชจ๋ธ์„ ์ œ์•ˆํ•ฉ๋‹ˆ๋‹ค:

throughputโปยน โˆ T_linear_dm / TP + (T_attn_dm + T_comp) / (DP ยท TP ยท PP) + T_comm(TP) / (DP ยท PP)

ํ•ญ๋ชฉ์˜๋ฏธ
T_linear_dmweight loading ์‹œ๊ฐ„
T_attn_dmKV cache ์ „์†ก ์‹œ๊ฐ„
T_compcompute ์‹œ๊ฐ„
T_commAllReduce ๋“ฑ ํ†ต์‹  ์‹œ๊ฐ„

์ด ์‹์—์„œ ์ค‘์š”ํ•œ ๊ฑด TP ์ฆ๊ฐ€ โ†’ T_linear_dm ๊ฐ์†Œ (โ†’ TP๋Š” decode์— ์ข‹์Œ), ํ•˜์ง€๋งŒ TP ์ฆ๊ฐ€ โ†’ T_comm ์ฆ๊ฐ€ (โ†’ PP๋Š” prefill์— ์ข‹์Œ)


๐Ÿ“Œ ์‹คํ–‰ ์˜ˆ์‹œ ๋น„๊ต

Case 1: Prefill with 2048 tokens (LLaMA2-13B)

์ „๋žตํŠน์ง•์˜ˆ์ƒ ์„ฑ๋Šฅ
TP4weight loading์€ ๋น ๋ฆ„, ํ•˜์ง€๋งŒ all-reduce ๋งŽ์ŒโŒ ๋А๋ฆผ
PP4๋ ˆ์ด์–ด๋ฅผ 4 GPU์— ๋‚˜๋ˆ  ์ˆœ์ฐจ ์ฒ˜๋ฆฌ, ํ†ต์‹  ์ ์Œโœ… ๋น ๋ฆ„

Case 2: Decode with 128 steps

์ „๋žตํŠน์ง•์˜ˆ์ƒ ์„ฑ๋Šฅ
TP4weight shard๋ฅผ ๋‚˜๋ˆ  ๊ฐ€์ง โ†’ ๋ฐ˜๋ณต loading ์—†์ด ์ฒ˜๋ฆฌโœ… ๋น ๋ฆ„
PP4๋งค step๋งˆ๋‹ค pipeline ์‹œ์ž‘ํ•ด์•ผ ํ•˜๋ฏ€๋กœ ์˜ค๋ฒ„ํ—ค๋“œ ํผโŒ ๋А๋ฆผ

โ†’ Seesaw๋Š” prefill์—” PP, decode์—” TP๋ฅผ ๊ฐ๊ฐ ์„ ํƒ โ†’ ์ตœ์  ์กฐํ•ฉ


๐Ÿ“Š ์‹คํ—˜ ๊ฒฐ๊ณผ๋กœ ๋ณธ ์‹ค์ œ ์ฐจ์ด (๋…ผ๋ฌธ Figure 12)

์ „๋žตPrefill ์†๋„Decode ์†๋„Total ์‹œ๊ฐ„
TP4๋А๋ฆผ๋น ๋ฆ„๋А๋ฆผ
PP4๋น ๋ฆ„๋А๋ฆผ๋А๋ฆผ
Seesaw(PPโ†’TP)๋น ๋ฆ„๋น ๋ฆ„โœ… ๊ฐ€์žฅ ๋น ๋ฆ„

๐Ÿง  ์ •๋ฆฌ

  • prefill์€ ๊ณ„์‚ฐ๋Ÿ‰์ด ๋งŽ๊ณ  all-reduce cost๊ฐ€ dominant โ†’ pipeline ๋ณ‘๋ ฌํ™”๋กœ ํ†ต์‹  ์˜ค๋ฒ„ํ—ค๋“œ ๊ฐ์†Œ
  • decode๋Š” ์—ฐ์‚ฐ๋Ÿ‰ ์ ๊ณ  weight ๋ฐ˜๋ณต ๋กœ๋”ฉ์ด ๋ณ‘๋ชฉ โ†’ tensor ๋ณ‘๋ ฌํ™”๋กœ weight shard ํ™œ์šฉ
  • Seesaw๋Š” ์ด๋ฅผ runtime์—์„œ stage-awareํ•˜๊ฒŒ ํŒ๋‹จํ•˜๊ณ , ๋™์ ์œผ๋กœ re-shardํ•จ์œผ๋กœ์จ throughput์„ ํ–ฅ์ƒ์‹œํ‚ด

Q : Tiered KV cache buffering๊ณผ transition-minimizing scheduling์ด ์–ด๋–ป๊ฒŒ ์ž‘๋™ํ•˜๋ฉฐ, ์ด๋Ÿฌํ•œ ์„ค๊ณ„๊ฐ€ CPUโ†”GPU ์ „์†ก ๋ณ‘๋ชฉ๊ณผ re-sharding ์˜ค๋ฒ„ํ—ค๋“œ๋ฅผ ์–ด๋–ป๊ฒŒ ์™„ํ™”ํ•˜๋Š”์ง€, ์‹ค์ œ ์ถ”๋ก  ์ฒ˜๋ฆฌ ๊ณผ์ •์— ๋”ฐ๋ผ ๊ตฌ์ฒด์ ์œผ๋กœ ์„ค๋ช…ํ•  ์ˆ˜ ์žˆ๋Š”๊ฐ€?

์•„์ฃผ ์ข‹์€ ์งˆ๋ฌธ์ž…๋‹ˆ๋‹ค. ์ด ์งˆ๋ฌธ์€ Seesaw ๋…ผ๋ฌธ์˜ ์„ฑ๋Šฅ ํ•ต์‹ฌ ์—”์ง„์ธ Tiered KV Cache Buffering๊ณผ Transition-minimizing Scheduling์˜ ์ž‘๋™ ๋ฐฉ์‹๊ณผ ๊ทธ ํšจ๊ณผ๋ฅผ ์‹ค์ œ ํ๋ฆ„ ์ค‘์‹ฌ์œผ๋กœ ์ดํ•ดํ•ด์•ผ ๋‹ตํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค.


โœ… ๊ฒฐ๋ก  ์š”์•ฝ

Tiered KV Cache Buffering์€ CPU ๋ฉ”๋ชจ๋ฆฌ๋ฅผ KV Cache์˜ ๋ณด์กฐ ์ €์žฅ์†Œ๋กœ ์‚ฌ์šฉํ•˜์—ฌ GPU ๋ฉ”๋ชจ๋ฆฌ ํ•œ๊ณ„๋ฅผ ๊ทน๋ณตํ•˜๊ณ , Transition-minimizing Scheduling์€ prefillโ†”decode ์‚ฌ์ด์˜ ์ „ํ™˜ ๋นˆ๋„๋ฅผ ์ตœ์†Œํ™”ํ•˜์—ฌ re-sharding ๋น„์šฉ์„ ์ค„์ž…๋‹ˆ๋‹ค.

์ด ๋‘ ์„ค๊ณ„๋Š” CPUโ†”GPU ๊ฐ„ ๋ฐ์ดํ„ฐ ๋ณ‘๋ชฉ๊ณผ ์žฆ์€ re-shard ๋น„์šฉ์ด๋ผ๋Š” ๋‘ ๊ฐ€์ง€ ๋ณ‘๋ชฉ์„ ๊ตฌ์กฐ์ ์œผ๋กœ ์™„ํ™”ํ•ฉ๋‹ˆ๋‹ค.


๐Ÿง  ์™œ ์ด ์„ค๊ณ„๊ฐ€ ํ•„์š”ํ–ˆ๋Š”๊ฐ€?

๊ธฐ์กด continuous batching ๋ฌธ์ œ:

  • prefill๊ณผ decode๋ฅผ ๋ฒˆ๊ฐˆ์•„ ์ˆ˜ํ–‰ํ•ด์•ผ ํ•˜๋ฏ€๋กœ โ†’ stage ์ „ํ™˜(re-sharding)์ด ์ž์ฃผ ๋ฐœ์ƒ

  • stage ์ „ํ™˜๋งˆ๋‹ค:

    • model weight ์žฌ๋ฐฐ์น˜
    • KV cache re-sharding (layout ๋ณ€๊ฒฝ)
    • GPUโ†”CPU ๋ฐ์ดํ„ฐ ์ด๋™

โ†’ ์ „ํ™˜์ด ์žฆ์„์ˆ˜๋ก ์ „์ฒด throughput์€ ๊ธ‰๊ฐ (Figure 2a ์ฐธ์กฐ)


๐Ÿ”„ Tiered KV Cache Buffering: ์–ด๋–ป๊ฒŒ ์ž‘๋™ํ•˜๋Š”๊ฐ€?

๊ธฐ๋ณธ ์•„์ด๋””์–ด:

  • KV Cache๋ฅผ GPU์— ์ €์žฅํ•˜์ง€ ์•Š๊ณ  CPU ๋ฉ”๋ชจ๋ฆฌ์— Tiered ๋ฐฉ์‹์œผ๋กœ ์ €์žฅ
  • Prefill์€ ์—ฐ์†์ ์œผ๋กœ ์ˆ˜ํ–‰ํ•˜๊ณ , ๊ทธ ๊ฒฐ๊ณผ(KV)๋ฅผ ๋ชจ๋‘ CPU์— ์ €์žฅ
  • Decode๋Š” ์ด CPU ๋ฒ„ํผ์—์„œ ํ•„์š”ํ•œ KV๋งŒ ๋น„๋™๊ธฐ๋กœ ๋ถˆ๋Ÿฌ์˜ด

์‹ค์ œ ์ž‘๋™ ํ๋ฆ„:

TEXT
[Prefill Stage]
1. Input Sequence โ†’ Prefill ์‹คํ–‰ (PP ๋ฐฉ์‹)
2. ๊ฐ GPU๊ฐ€ ์ž๊ธฐ shard์˜ KV ๊ณ„์‚ฐ
3. GPUโ†’Pinned Memoryโ†’CPU Shared Memory ๋กœ KV ์บ์‹œ ์ €์žฅ

[Decode Stage]
4. Scheduler๋Š” CPU buffer์—์„œ decode์šฉ ์‹œํ€€์Šค๋ฅผ ์„ ํƒ
5. Worker๋Š” Prefetch Thread๋กœ GPU์— ํ•„์š”ํ•œ shard๋งŒ ๋น„๋™๊ธฐ๋กœ Load
6. Load ์™„๋ฃŒ๋œ ์‹œํ€€์Šค๋ถ€ํ„ฐ TP ๊ธฐ๋ฐ˜ Decode ์‹คํ–‰

ํ•ต์‹ฌ ์žฅ์น˜: Shared CPU memory ๊ตฌ์กฐ๋ฅผ ํ™œ์šฉํ•˜์—ฌ GPU ๊ฐ„ KV ์ด๋™ ์—†์ด ์žฌshard ๊ฐ€๋Šฅ


โฑ๏ธ Transition-minimizing Scheduling: ์–ด๋–ป๊ฒŒ ์ž‘๋™ํ•˜๋Š”๊ฐ€?

๋ฌธ์ œ์ 

  • ๊ธฐ์กด ์Šค์ผ€์ค„๋Ÿฌ๋Š” decode๊ฐ€ ๋๋‚˜์ž๋งˆ์ž prefill ์‹œ์ž‘ โ†’ ๋งค๋ฒˆ stage ์ „ํ™˜
  • โ†’ weight ์ด๋™ + KV ์บ์‹œ ์ด๋™์ด ๋„ˆ๋ฌด ์ž์ฃผ ๋ฐœ์ƒ โ†’ throughput ๊ธ‰๊ฐ

๊ฐœ์„ ๋œ ๋ฐฉ์‹:

์กฐ๊ฑด์Šค์ผ€์ค„๋Ÿฌ ๋™์ž‘
CPU KV ๋ฒ„ํผ๊ฐ€ ๊ฐ€๋“ ์ฐผ์„ ๋•Œ๋งŒprefill โ†’ decode ์ „ํ™˜
CPU KV ๋ฒ„ํผ๊ฐ€ ๋น„์—ˆ์„ ๋•Œ๋งŒdecode โ†’ prefill ์ „ํ™˜

โ†’ ํ•œ๋ฒˆ์— ๋งŽ์€ ์‹œํ€€์Šค๋ฅผ prefillํ•œ ํ›„์— decode๋งŒ ์—ฐ์† ์ˆ˜ํ–‰ํ•จ โ†’ decode ์ˆ˜ํ–‰ ์ค‘์—๋Š” ์ƒˆ prefill์„ ํ•˜์ง€ ์•Š์Œ


๐Ÿ“Š ํšจ๊ณผ: ๋ณ‘๋ชฉ ์™„ํ™”

ํ•ญ๋ชฉ๊ธฐ์กด ๋ฐฉ์‹Seesaw ๋ฐฉ์‹
GPU memory ์ œํ•œdecode ์ค‘ ๋ฉ”๋ชจ๋ฆฌ ๋ถ€์กฑ ๋ฐœ์ƒCPU๋กœ ๋ถ„์‚ฐ ์ €์žฅ ๊ฐ€๋Šฅ
stage ์ „ํ™˜ ์ˆ˜์ˆ˜๋ฐฑ ํšŒ ๊ฐ€๋Šฅ์ˆ˜์‹ญ ํšŒ๋กœ ์ค„์ž„
re-shard ์‹œ๊ฐ„์ „ํ™˜๋งˆ๋‹ค ๋ฐœ์ƒ์ „ํ™˜ ํšŸ์ˆ˜๋ฅผ ์ค„์—ฌ amortize
GPU idle ์‹œ๊ฐ„frequent swappingcontinuous compute ๊ฐ€๋Šฅ

๐Ÿ“Œ ์‹ค์ œ ์ถ”๋ก  ์˜ˆ์‹œ

ํ™˜๊ฒฝ:

  • ์‹œํ€€์Šค 128๊ฐœ
  • prefill: PP4, decode: TP4

์Šค์ผ€์ค„ ๋™์ž‘:

  1. Scheduler๋Š” 128๊ฐœ ์‹œํ€€์Šค๋ฅผ prefill stage๋กœ ์—ฐ์† ์ฒ˜๋ฆฌ
  2. ๊ฐ ์‹œํ€€์Šค์˜ KV ์บ์‹œ๋Š” GPU โ†’ CPU๋กœ ์ด๋™
  3. CPU ๋ฒ„ํผ๊ฐ€ ๊ฐ€๋“ ์ฐจ๋ฉด prefill ์ค‘๋‹จ โ†’ re-shard ์‹คํ–‰
  4. TP4๋กœ ๋ชจ๋ธ weight shard ์žฌ๋ฐฐ์น˜
  5. Worker๋Š” CPU buffer์—์„œ ์‹œํ€€์Šค ํ•˜๋‚˜์”ฉ GPU๋กœ ๋น„๋™๊ธฐ ๋กœ๋”ฉ
  6. decode๊ฐ€ ๋๋‚œ ์‹œํ€€์Šค๋Š” GPU์—์„œ KV ํ•ด์ œ
  7. CPU buffer๊ฐ€ ๋‹ค ๋น„์›Œ์ง€๋ฉด ๋‹ค์‹œ prefill๋กœ ์ „ํ™˜

โ†’ ์ „ํ™˜์€ ๋‹จ 1ํšŒ, throughput ์†์‹ค ์—†์ด ๋Œ€๋Ÿ‰ ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅ


๐Ÿง  ์ •๋ฆฌ ์š”์•ฝ

์š”์†Œํšจ๊ณผ
Tiered KV BufferingGPU ๋ฉ”๋ชจ๋ฆฌ ํ•œ๊ณ„ ๊ทน๋ณต, ๋Œ€๊ทœ๋ชจ prefill ๊ฐ€๋Šฅ
Transition-minimizing Schedulingre-shard ํšŸ์ˆ˜ ๊ฐ์†Œ, ์ „์ฒด ์ฒ˜๋ฆฌ๋Ÿ‰ ์ฆ๊ฐ€
Prefetch ThreadCPUโ†’GPU ์ „์†ก ์ง€์—ฐ์„ compute์™€ overlap ์ฒ˜๋ฆฌ
Shared Memory ๊ธฐ๋ฐ˜ KV ์บ์‹œTP/PP ๊ฐ„ re-shard๋ฅผ ๋น ๋ฅด๊ฒŒ ์ฒ˜๋ฆฌ ๊ฐ€๋Šฅ

๋ผ์ด์„ ์Šค

์ž‘์„ฑ์ž: Jaehun Ryu

๋งํฌ: https://jaehun.me/posts/seesaw-high-throughput-llm-inference-via-model-re-sharding/

๋ผ์ด์„ ์Šค: CC BY 4.0

์ด ์ €์ž‘๋ฌผ์€ ํฌ๋ฆฌ์—์ดํ‹ฐ๋ธŒ ์ปค๋จผ์ฆˆ ์ €์ž‘์žํ‘œ์‹œ 4.0 ๊ตญ์ œ ๋ผ์ด์„ ์Šค์— ๋”ฐ๋ผ ์ด์šฉํ•  ์ˆ˜ ์žˆ์Šต๋‹ˆ๋‹ค. ์ถœ์ฒ˜๋ฅผ ๋ฐํžˆ๋ฉด ์ƒ์—…์  ๋ชฉ์ ์„ ํฌํ•จํ•ด ์ž์œ ๋กญ๊ฒŒ ์ด์šฉ ๊ฐ€๋Šฅํ•ฉ๋‹ˆ๋‹ค.

๋Œ“๊ธ€