๊ฐ์ ๊ณผ ๋ ์ฐฝ์ฑ, ํต์ฌ ์๊ณ ๋ฆฌ์ฆ, ๊ทธ๋ฆฌ๊ณ ํ๊ณ
๋ฌผ์ด๋ณธ ๊ฒ: ์ด ๋ ผ๋ฌธ์ ์์ฃผ ์์ธํ๊ฒ ์ฝ๊ณ ๋ ผ๋ฌธ์ ๊ฐ์ ๊ณผ ๋ ์ฐฝ์ ์ธ ์ง์ ์ ์ค๋ช ํด์ฃผ๊ณ ํต์ฌ ์๊ณ ๋ฆฌ์ฆ์ ์์ ์ ๋ ฅ์ ๋ค์ด์ ์ ์ฒด์ ์ธ ๊ณผ์ ์ ์ค๋ช ํด์ค ์ถ๊ฐ์ ์ผ๋ก ๋ ผ๋ฌธ์ ํ๊ณ์ ์ ๋ํด์๋ ์๋ ค์ค
๐ ๊ฒฐ๋ก ์์ฝ
๋ ผ๋ฌธ โ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)
๐ ์ ์ฒด ์ถ๋ก ์ฒ๋ฆฌ ํ๋ฆ
Prefill ๋จ๊ณ (PP4 ์ ์ฉ):
- ์ ๋ ฅ ์ ์ฒด ์ํ์ค๊ฐ GPU 4๋์ pipeline ํํ๋ก ๋ถ๋ฐฐ๋จ
- ๋ชจ๋ธ ๋ ์ด์ด๋ ์์ฐจ์ ์ผ๋ก ๊ฐ GPU์ ๋๋์ด ์ฒ๋ฆฌ
- ํต์ ์ค๋ฒํค๋ ์์, weight loading ํจ์จ์ ์
KV ์บ์ ์์ฑ ๋ฐ CPU๋ก ์คํ๋ก๋:
- ๊ฐ GPU๋ ์์ ์ shard KV๋ฅผ CPU ๊ณต์ ๋ฉ๋ชจ๋ฆฌ์ ์ ์ฅ
๋ชจ๋ธ weight์ KV cache ์ฌ์ค๋ฉ:
- pipeline ๊ตฌ์กฐ๋ก ๋๋์๋ ๋ชจ๋ธ weight โ tensor ๊ตฌ์กฐ๋ก ์ฌ๋ฐฐ์น
- KV ์บ์๋ TP ๊ตฌ์กฐ๋ก ์ฌ๋ถ๋ฐฐ๋จ (CPUโGPU ๋น๋๊ธฐ ์ ์ก)
Decode ๋จ๊ณ (TP4 ์ ์ฉ):
- ๋ชจ๋ GPU๊ฐ ๋์ผ weight shard๋ก ๋ณ๋ ฌ ์ฒ๋ฆฌ
- ํ ํ ํฐ์ฉ ์์ฑ โ GPU ๊ฐ AllReduce ํต์ ๋ฐ์
- weight loading ๋ณ๋ ฌ ์ฒ๋ฆฌ๋ก ํจ์จ์ฑ ๊ทน๋ํ๋จ
๋น๋๊ธฐ ์ฒ๋ฆฌ ๋ฐ stage ์ ํ ์ต์ ํ:
- CPUโGPU KV ์ ์ก์ prefetch thread๋ก ๋น๋๊ธฐ ์ํ
- KV๊ฐ ๊ฝ ์ฐฐ ๋๋ง decode๋ก ์ ํํ์ฌ ์ฌ์ค๋ฉ ํ์ ์ต์ํ
๐ ์ฑ๋ฅ ๋น๊ต (vLLM vs Seesaw)
| ํ๊ฒฝ | ๋ชจ๋ธ | vLLM (req/s) | Seesaw (req/s) | ์๋ ํฅ์ ๋ฐฐ์ |
|---|---|---|---|---|
| A10, 4GPU | 15B | 1.0 (๊ธฐ์ค) | 1.45 | +45% |
| L4, 4GPU | 15B | 1.0 | 1.29 | +29% |
| A100 (PCIe), 8GPU | 70B | 1.0 | 1.46 | +46% |
| A100 (NVLink), 8GPU | 70B | 1.0 | 1.13 | +13% |
๐ง ๋ ผ๋ฌธ์ ๋ ์ฐฝ์ฑ
| ํญ๋ชฉ | Seesaw์ ์ฐจ๋ณ์ |
|---|---|
| ๋ณ๋ ฌํ ์ ๋ต | stage ๋ณ๋ก ๋ณ๋ ฌํ ์ ๋ต์ ๋ฐ๊ฟ ์ ์๋๋ก ์ค๊ณ (TP โ PP) |
| Re-sharding overhead ์ฒ๋ฆฌ | CPU๋ฅผ ์ค๊ฐ ์ ์ฅ์๋ก ํ์ฉํ tiered buffering ๋์ |
| Scheduling ์ต์ ํ | transition-minimizing scheduler ์ค๊ณ |
| ์ ์ฉ ์ ์ฐ์ฑ | ๊ธฐ์กด ์์คํ (vLLM, TensorRT-LLM)์๋ ์ฝ๊ฒ ํตํฉ ๊ฐ๋ฅ |
โ ๏ธ ํ๊ณ์ ๋ฐ ์ ์ฝ
์ฌ์ค๋ฉ ์ค๋ฒํค๋ ์กด์ฌ
- CPUโGPU ์ฌ์ด์ ๋ฐ์ดํฐ ์ด๋์ ์ฌ์ ํ ๋ณ๋ชฉ ๊ฐ๋ฅ์ฑ ์์
- ํด๊ฒฐ์ ์ํด ๋น๋๊ธฐ prefetching ์ฌ์ฉํ์ง๋ง ํ๋์จ์ด ์์กด๋ ๋์
๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ ์ฆ๊ฐ
- CPU, GPU ๋ชจ๋์์ KV cache๋ฅผ ๋ณด๊ดํด์ผ ํ๋ฏ๋ก ๋ ๋ง์ ๋ฉ๋ชจ๋ฆฌ ์๊ตฌ
๋ณต์กํ ์ํคํ ์ฒ ๊ตฌ์ฑ
- worker-scheduler ๊ตฌ์กฐ์ ๋๊ธฐํ ํ์๋ก ์ธํด ์์คํ ๊ตฌํ ๋ณต์ก๋ ์ฆ๊ฐ
online-serving ํ๊ฒฝ์์๋ ๋ ์ ํฉ
- offline throughput ์ค์ฌ ์ค๊ณ โ latency ๋ฏผ๊ฐํ ์์คํ ์๋ ๋ถ์ ์
ํต์ฌ ์๊ณ ๋ฆฌ์ฆ์ ์์๋ก ๋ฐ๋ผ๊ฐ๊ธฐ
๋ฌผ์ด๋ณธ ๊ฒ: ๋ ผ๋ฌธ์ ํต์ฌ์๊ณ ๋ฆฌ์ฆ์ ์์๋ฅผ ๋ค์ด ์์ฃผ ์์ธํ๊ฒ ์ค๋ช ํด์ฃผ๊ณ ์ ๋ฆฌํด ์ค
๋ ผ๋ฌธ “Seesaw: High-throughput LLM Inference via Model Re-sharding”์ ํต์ฌ ์๊ณ ๋ฆฌ์ฆ์ ๋ค์ 3๊ฐ์ง ๊ตฌ์ฑ ์์๋ก ์์ฝ๋ฉ๋๋ค:
โ ํต์ฌ ์๊ณ ๋ฆฌ์ฆ ๊ตฌ์ฑ ์์ฝ
| ๊ตฌ์ฑ ์์ | ์ญํ |
|---|---|
| 1. Dynamic Model Re-sharding | Prefill ๋จ๊ณ์์ Pipeline Parallelism (PP), Decode ๋จ๊ณ์์ Tensor Parallelism (TP)์ ์ฌ์ฉํ๋๋ก ๋ชจ๋ธ weight์ KV cache์ ๋ณ๋ ฌํ ์ ๋ต์ ๋์ ์ผ๋ก ๋ณ๊ฒฝ |
| 2. Tiered KV Cache Buffering | GPU ๋ฉ๋ชจ๋ฆฌ ํ๊ณ๋ฅผ ๊ทน๋ณตํ๊ธฐ ์ํด 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 ๋ฉ๋ชจ๋ฆฌ๋ก ์คํ๋ก๋ํจ.
[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 ๊ทน๋ํ
[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์ ์ ์ง
๐ ์ ์ฒด ๋์ ํ๋ฆ ์์ฝ ๋ค์ด์ด๊ทธ๋จ
[์์] โ Prefill(PP) โ KV Cache(CPU์ ์ ์ฅ) โ Model Re-shard(TP) โ Decode(TP) โ ๋
โ โ
โโโโโโโโ Transition-minimizing (๋ฒํผ ๊ฝ ์ฐจ๋ฉด ์ ํ) โโ๐ง ์ ๋ฆฌ
| ๋จ๊ณ | ๋ณ๋ ฌํ ์ ๋ต | ์ฅ์ | ๋จ์ |
|---|---|---|---|
| Prefill | Pipeline Parallelism | ๋ฎ์ ํต์ ๋น์ฉ, ํ ํฐ ๋ณ๋ ฌ ์ฒ๋ฆฌ | micro-batch overhead |
| Decode | Tensor Parallelism | ๋น ๋ฅธ weight loading, ๋์ throughput | AllReduce ํต์ ๋น์ฉ |
| KV ์ ์ฅ | Tiered Buffering | GPU ๋ฉ๋ชจ๋ฆฌ ํ๊ณ ๊ทน๋ณต | CPUโGPU ์ ์ก ์๋ ๋ณ๋ชฉ |
| ์ค์ผ์ค๋ง | Transition-Minimized | re-shard ํ์ ์ต์ํ | ๋ณต์กํ ์ค์ผ์ค๋ฌ ๋ก์ง ํ์ |
์คํ ๊ฒฐ๊ณผ์ ๋ค๋ฅธ ๋ฐฉ๋ฒ๋ก ๊ณผ์ ์ฐจ์ด
๋ฌผ์ด๋ณธ ๊ฒ: ์ด ๋ ผ๋ฌธ์์ ์ ์ํ ๊ฒฐ๊ณผ๋ฅผ ์์ธํ๊ฒ ๋ณด๊ณ ๋ค๋ฅธ ๋ฐฉ๋ฒ๋ก ์ ๋นํ์ฌ ํน์ถ๋ ์ ๊ณผ ๋ ผ๋ฌธ์์ ์ ๊ธฐํ๋ ์ด๋ ํ ๋ฐฉ๋ฒ์ด ์ด๋ฌํ ๊ฒฐ๊ณผ๋ฅผ ๋์ถํ๊ฒ ๋์๋์ง ๋ ผ๋ฌธ์์ ์ ์ํ๋ ์ด์ ์ ๋์ ์๊ฐ์ ์๋ ค์ค
๋ ผ๋ฌธ Seesaw: High-throughput LLM Inference via Model Re-sharding์ ๋ค์ํ ํ๋์จ์ด ๋ฐ ๋ชจ๋ธ ์ค์ ์์ SOTA ์ถ๋ก ์์ง์ธ vLLM ๋๋น ์ต๋ 1.78ร, ํ๊ท 1.36ร์ throughput ํฅ์์ ๋ณด์ฌ์ค๋๋ค. ์ด ์ฑ๊ณผ๋ ๋จ์ผ ๋ณ๋ ฌํ ์ ๋ต์ ๊ณ ์ ๋ ๊ธฐ์กด ์์ง๋ค๊ณผ ์ฐจ๋ณ๋๋ ๋์ ๋ณ๋ ฌํ ์ ๋ต ์กฐํฉ (PP + TP)๊ณผ transition-aware scheduling ๋๋ถ์ ๋๋ค.
๐ Seesaw ์ฑ๋ฅ ๊ฒฐ๊ณผ ์ ๋ฆฌ (๋ ผ๋ฌธ 6์ฅ ๊ธฐ์ค)
| ๋ชจ๋ธ | GPU | vLLM Throughput | Seesaw Throughput | ๊ฐ์ ๋ฐฐ์ | ์ค๋ช |
|---|---|---|---|---|---|
| LLaMA3-15B | 4 ร A10 | 1.0 (baseline) | 1.45 | +45% | chunked prefill ์กฐํฉ๋ณด๋ค ์ฐ์ |
| CodeLLaMA-34B | 8 ร A10 | 1.0 | 1.78 | +78% | prefill/decoding stage ๋ถ๋ฆฌ๋ก ์ด๋ ๊ทน๋ํ |
| LLaMA2-70B | 8 ร A100 PCIe | 1.0 | 1.46 | +46% | TP์์ ๋ฐ์ํ๋ all-reduce ๋ณ๋ชฉ ์ํ |
| LLaMA2-70B | 8 ร A100 NVLink | 1.0 | 1.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 ๊ธฐ๋ฐ)
{
"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.
๐ ์ ์ฒด ์์ ์ถ๋ก ํ๋ฆ
์
๋ ฅ: 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. ๋ชจ๋ธ ์ํคํ ์ฒ ๊ตฌ์ฑ
๋ ผ๋ฌธ ์คํ์๋ ๋ค์ ๋ชจ๋ธ์ด ์ฌ์ฉ๋ฉ๋๋ค:
| ๋ชจ๋ธ๋ช | ํฌ๊ธฐ | ์ํคํ ์ฒ |
|---|---|---|
| LLaMA3 | 15B | GQA ๊ธฐ๋ฐ decoder-only |
| CodeLLaMA | 34B | GQA ๊ธฐ๋ฐ decoder-only |
| LLaMA2 | 70B | GQA ๊ธฐ๋ฐ decoder-only |
๋ชจ๋ transformer decoder-only ๊ตฌ์กฐ๋ก ๋์ด ์์ผ๋ฉฐ, ํต์ฌ ๊ตฌ์ฑ์ ์๋์ ๊ฐ์ต๋๋ค:
๐น ํ Layer ๊ตฌ์ฑ
- LayerNorm
- Multi-head Attention (QKV ๊ณ์ฐ, GQA ๊ตฌ์กฐ)
- Residual Add
- FeedForward Network (2~4ร hidden size)
- ๋ ๋ค๋ฅธ 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๊ฐ ๋ณ๋ชฉ
๐งพ ์ ๋ฆฌ ์์ฝ
| ํญ๋ชฉ | Prefill | Decode |
|---|---|---|
| ๋ณ๋ ฌํ ์ ๋ต | 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 overhead | Layer-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 & Prefetching | online + 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 cache | PCIe ๋ณ๋ชฉ, ์ ์ก ์ง์ฐ | ์์ถ + L2-style ์บ์ ๊ณ์ธตํ |
| Scheduling | offline 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_dm | weight loading ์๊ฐ |
| T_attn_dm | KV cache ์ ์ก ์๊ฐ |
| T_comp | compute ์๊ฐ |
| T_comm | AllReduce ๋ฑ ํต์ ์๊ฐ |
์ด ์์์ ์ค์ํ ๊ฑด TP ์ฆ๊ฐ โ T_linear_dm ๊ฐ์ (โ TP๋ decode์ ์ข์), ํ์ง๋ง TP ์ฆ๊ฐ โ T_comm ์ฆ๊ฐ (โ PP๋ prefill์ ์ข์)
๐ ์คํ ์์ ๋น๊ต
Case 1: Prefill with 2048 tokens (LLaMA2-13B)
| ์ ๋ต | ํน์ง | ์์ ์ฑ๋ฅ |
|---|---|---|
| TP4 | weight loading์ ๋น ๋ฆ, ํ์ง๋ง all-reduce ๋ง์ | โ ๋๋ฆผ |
| PP4 | ๋ ์ด์ด๋ฅผ 4 GPU์ ๋๋ ์์ฐจ ์ฒ๋ฆฌ, ํต์ ์ ์ | โ ๋น ๋ฆ |
Case 2: Decode with 128 steps
| ์ ๋ต | ํน์ง | ์์ ์ฑ๋ฅ |
|---|---|---|
| TP4 | weight 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๋ง ๋น๋๊ธฐ๋ก ๋ถ๋ฌ์ด
์ค์ ์๋ ํ๋ฆ:
[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 swapping | continuous compute ๊ฐ๋ฅ |
๐ ์ค์ ์ถ๋ก ์์
ํ๊ฒฝ:
- ์ํ์ค 128๊ฐ
- prefill: PP4, decode: TP4
์ค์ผ์ค ๋์:
- Scheduler๋ 128๊ฐ ์ํ์ค๋ฅผ prefill stage๋ก ์ฐ์ ์ฒ๋ฆฌ
- ๊ฐ ์ํ์ค์ KV ์บ์๋ GPU โ CPU๋ก ์ด๋
- CPU ๋ฒํผ๊ฐ ๊ฐ๋ ์ฐจ๋ฉด prefill ์ค๋จ โ re-shard ์คํ
- TP4๋ก ๋ชจ๋ธ weight shard ์ฌ๋ฐฐ์น
- Worker๋ CPU buffer์์ ์ํ์ค ํ๋์ฉ GPU๋ก ๋น๋๊ธฐ ๋ก๋ฉ
- decode๊ฐ ๋๋ ์ํ์ค๋ GPU์์ KV ํด์
- CPU buffer๊ฐ ๋ค ๋น์์ง๋ฉด ๋ค์ prefill๋ก ์ ํ
โ ์ ํ์ ๋จ 1ํ, throughput ์์ค ์์ด ๋๋ ์ฒ๋ฆฌ ๊ฐ๋ฅ
๐ง ์ ๋ฆฌ ์์ฝ
| ์์ | ํจ๊ณผ |
|---|---|
| Tiered KV Buffering | GPU ๋ฉ๋ชจ๋ฆฌ ํ๊ณ ๊ทน๋ณต, ๋๊ท๋ชจ prefill ๊ฐ๋ฅ |
| Transition-minimizing Scheduling | re-shard ํ์ ๊ฐ์, ์ ์ฒด ์ฒ๋ฆฌ๋ ์ฆ๊ฐ |
| Prefetch Thread | CPUโGPU ์ ์ก ์ง์ฐ์ compute์ overlap ์ฒ๋ฆฌ |
| Shared Memory ๊ธฐ๋ฐ KV ์บ์ | TP/PP ๊ฐ re-shard๋ฅผ ๋น ๋ฅด๊ฒ ์ฒ๋ฆฌ ๊ฐ๋ฅ |
๋๊ธ