HMT: Hierarchical Memory Training with Dynamic Low-Rank Optimization and Hybrid Activation Compression
์์ VRAM ํ๊ฒฝ์์ ๊ฑฐ๋ ์ธ์ด๋ชจ๋ธ ํ์ต์ ์ํ ๊ณ์ธตํ ๋ฉ๋ชจ๋ฆฌ ํ์ต ์๊ณ ๋ฆฌ์ฆ
๋๊ท๋ชจ ์ธ์ด๋ชจ๋ธ ํ์ต์ ๋ชจ๋ธ ํ๋ผ๋ฏธํฐ, gradient, optimizer state, activation์ผ๋ก ์ธํด ๋ง๋ํ GPU VRAM์ ์๊ตฌํ๋ค. ๊ธฐ์กด ์ ๊ทผ์ LoRA/QLoRA, ZeRO offload, gradient checkpointing, CPU/NVMe offload๋ฅผ ํตํด ๋ฉ๋ชจ๋ฆฌ ๋ถ์กฑ์ ์ํํ์ง๋ง, ๋ง์ ๊ฒฝ์ฐ ํ์ต ์์ ๋ ์ ํ, PCIe/NVMe ๋ณ๋ชฉ, ์ฌ๊ณ์ฐ ๋น์ฉ ์ฆ๊ฐ๋ผ๋ ํ๊ณ๋ฅผ ๊ฐ์ง๋ค. ๋ณธ ๋ ผ๋ฌธ์ ์์ VRAM ํ๊ฒฝ์์ ๊ฑฐ๋ LLM์ ๋ ํจ์จ์ ์ผ๋ก ํ์ตํ๊ธฐ ์ํ ๊ณ์ธตํ ๋ฉ๋ชจ๋ฆฌ ํ์ต ์๊ณ ๋ฆฌ์ฆ HMT๋ฅผ ์ ์ํ๋ค. HMT๋ ๋จ์ํ ๋ฐ์ดํฐ๋ฅผ CPU๋ ๋์คํฌ๋ก ๋ด๋ฆฌ๋ ๋ฐฉ์์ด ์๋๋ผ, ํ์ต ์ค ์์ฑ๋๋ gradient, optimizer state, activation์ ํํ ์์ฒด๋ฅผ ์๊ฒ ๋ง๋ ๋ค. ๊ตฌ์ฒด์ ์ผ๋ก๋ ๋์ ์ ๋ญํฌ gradient projection, APOLLO/GaLore ๊ณ์ด optimizer state ์์ถ, layer-wise hybrid activation compression, CPU basis cache, NVMe checkpoint-only policy๋ฅผ ๊ฒฐํฉํ๋ค. ๋ชฉํ๋ offload ์ค์ฌ ํ์ต๋ณด๋ค ๋์ ์ฒ๋ฆฌ๋์ ์ ์งํ๋ฉด์ full-parameter fine-tuning์ ๊ฐ๊น์ด ํ์ต ์์ ๋๋ฅผ ํ๋ณดํ๋ ๊ฒ์ด๋ค.
๊ฑฐ๋ ์ธ์ด๋ชจ๋ธ์ ํ๋ผ๋ฏธํฐ ์๊ฐ ์ฆ๊ฐํ ์๋ก ์ถ๋ก ๊ณผ ํ์ต ๋ชจ๋์์ ๋ฉ๋ชจ๋ฆฌ ์๊ตฌ๋์ด ๊ธ๊ฒฉํ ์ฆ๊ฐํ๋ค. ํนํ ํ์ต ๋จ๊ณ์์๋ ๋จ์ํ ๋ชจ๋ธ weight๋ง ์ ์ฅํ๋ ๊ฒ์ด ์๋๋ผ gradient, optimizer state, activation๊น์ง ์ ์ฅํด์ผ ํ๋ฏ๋ก ๋ฉ๋ชจ๋ฆฌ ์ฌ์ฉ๋์ด ์ถ๋ก ๋ณด๋ค ํจ์ฌ ํฌ๋ค. AdamW optimizer๋ฅผ ์ฌ์ฉํ๋ ์ผ๋ฐ์ ์ธ mixed precision ํ์ต์์๋ optimizer state๊ฐ ๋ชจ๋ธ weight๋ณด๋ค ๋ ํฐ ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์ฐจ์งํ ์ ์๋ค.
์์ VRAM ํ๊ฒฝ์์ ์ฌ์ฉ๋๋ ๋ํ์ ์ธ ๋ฐฉ๋ฒ์ ๋ค์๊ณผ ๊ฐ๋ค.
์ฒซ์งธ, LoRA ๋๋ QLoRA๋ ์๋ณธ ๋ชจ๋ธ weight๋ฅผ ๊ณ ์ ํ๊ณ ์์ adapter๋ง ํ์ตํจ์ผ๋ก์จ trainable parameter์ optimizer state๋ฅผ ์ค์ธ๋ค. ๊ทธ๋ฌ๋ weight update ๊ณต๊ฐ์ด ์ ํ๋๊ธฐ ๋๋ฌธ์ full fine-tuning๊ณผ ๋ค๋ฅธ ํ์ต ๋์ญํ์ ๊ฐ๋๋ค.
๋์งธ, DeepSpeed ZeRO, FSDP, CPU/NVMe offload๋ optimizer state, gradient, parameter๋ฅผ CPU RAM ๋๋ NVMe๋ก ์ด๋์์ผ GPU VRAM ์ฌ์ฉ๋์ ์ค์ธ๋ค. ๊ทธ๋ฌ๋ CPUโGPU ๋๋ NVMeโGPU ์ ์ก์ด ๋น๋ฒํด์ง์๋ก PCIe์ ๋์คํฌ I/O๊ฐ ๋ณ๋ชฉ์ด ๋๋ค.
์ ์งธ, gradient checkpointing์ activation ์ ์ฅ๋์ ์ค์ด๋ ๋์ backward ๊ณผ์ ์์ forward ๊ณ์ฐ์ ๋ค์ ์ํํ๋ค. ์ด ๋ฐฉ์์ ๋ฉ๋ชจ๋ฆฌ ์ ๊ฐ ํจ๊ณผ๋ ํฌ์ง๋ง ํ์ต ์๊ฐ์ด ์ฆ๊ฐํ๋ค.
์ต๊ทผ์๋ ๊ธฐ์กด ๋ฐฉ์๊ณผ ๋ค๋ฅธ ์ ๊ทผ์ด ๋ฑ์ฅํ๊ณ ์๋ค. GaLore๋ weight space๊ฐ ์๋๋ผ gradient space์์ low-rank projection์ ์ํํ์ฌ full-parameter learning์ ๊ฐ๊น์ด ํ์ต์ ์ ์งํ๋ฉด์ optimizer state ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์ค์ธ๋ค. ํด๋น ์ฐ๊ตฌ๋ optimizer state ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์ต๋ 65.5% ์ค์ด๊ณ , 8-bit GaLore์์๋ optimizer memory๋ฅผ ์ต๋ 82.5%, ์ ์ฒด training memory๋ฅผ 63.3% ์ค์๋ค๊ณ ๋ณด๊ณ ํ๋ค. ๋ํ 24GB RTX 4090์์ 7B ๋ชจ๋ธ pretraining ๊ฐ๋ฅ์ฑ์ ๋ณด์๋ค๊ณ ์ค๋ช ํ๋ค.
APOLLO๋ AdamW์ learning-rate adaptation์ ์ค๋ณต์ฑ์ด ์๋ค๋ ๊ด์ฐฐ์์ ์ถ๋ฐํ์ฌ, auxiliary low-rank optimizer state๋ก AdamW ์์ค์ ์ฑ๋ฅ์ SGD ์์ค ๋ฉ๋ชจ๋ฆฌ ๋น์ฉ์ ๊ฐ๊น๊ฒ ๋ฌ์ฑํ๋ ค๋ ๋ฐฉ๋ฒ์ด๋ค. ํนํ APOLLO-Mini๋ rank-1 ๋ณํ์์๋ ๊ฐํ ๋ฉ๋ชจ๋ฆฌ ํจ์จ์ ๋ชฉํ๋ก ํ๋ค.
CompAct๋ optimizer๊ฐ ์๋๋ผ backward์ ํ์ํ activation compute graph๋ฅผ ๋์์ผ๋ก ํ๋ฉฐ, compressed activation์ ์ ์ฅํด pretraining์์ GPU peak memory๋ฅผ 25~30%, fine-tuning์์ 50% ์ค์๋ค๊ณ ๋ณด๊ณ ํ๋ค.
๋ณธ ๋ ผ๋ฌธ์ ์ด๋ฌํ ํ๋ฆ์ ๊ฒฐํฉํ์ฌ, ์์ VRAM ํ๊ฒฝ์์ ๊ฑฐ๋ LLM์ ํ์ตํ๊ธฐ ์ํ ์๋ก์ด ํตํฉ ์๊ณ ๋ฆฌ์ฆ์ ์ ์ํ๋ค.
๋ณธ ์ฐ๊ตฌ์ ๋ชฉํ๋ ๋ค์๊ณผ ๊ฐ๋ค.
์์ VRAM ํ๊ฒฝ์์ ๊ฑฐ๋ LLM์ ํ์ตํ๋, ๋จ์ CPU/NVMe offload์ ์์กดํ์ง ์๊ณ ํ์ต ์ค ์์ฑ๋๋ ์ฃผ์ ๋ฉ๋ชจ๋ฆฌ ํญ๋ชฉ์ ์ํ์ ์ผ๋ก ์์ถํ๋ค.
๋์ ๋ฉ๋ชจ๋ฆฌ ํญ๋ชฉ์ ๋ค์๊ณผ ๊ฐ๋ค.
- Model parameters
- Gradients
- Optimizer states
- Activations
- Temporary buffers
์ฐ๋ฆฌ๊ฐ ์ต์ ํํ๋ ค๋ ๋ชฉ์ ํจ์๋ ๋ค์๊ณผ ๊ฐ์ด ์ ์ํ ์ ์๋ค.
minimize peak_gpu_memory
maximize tokens_per_second
maintain validation_loss_quality
์ฆ, ๋จ์ํ ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์ค์ด๋ ๊ฒ์ด ์๋๋ผ ๋ฉ๋ชจ๋ฆฌ, ์๋, ํ์ต ํ์ง์ ๊ท ํ์ ์ต์ ํํ๋ค.
๊ธฐ์กด ์์ VRAM ํ์ต ๋ฐฉ์์ ํฌ๊ฒ ๋ ๋ถ๋ฅ๋ก ๋๋๋ค.
์ฒซ ๋ฒ์งธ๋ ํ์ต ๋ฒ์๋ฅผ ์ค์ด๋ ๋ฐฉ์์ด๋ค. LoRA/QLoRA๊ฐ ๋ํ์ ์ด๋ค. ์ด ๋ฐฉ์์ ๋งค์ฐ ์ค์ฉ์ ์ด์ง๋ง, ์๋ณธ weight ์ ์ฒด๋ฅผ ์ง์ ์ ๋ฐ์ดํธํ์ง ์๋๋ค.
๋ ๋ฒ์งธ๋ ํ์ต ์ํ๋ฅผ ์ธ๋ถ ๋ฉ๋ชจ๋ฆฌ๋ก ์ด๋ํ๋ ๋ฐฉ์์ด๋ค. ZeRO-Offload, ZeRO-Infinity, CPU/NVMe offload๊ฐ ์ฌ๊ธฐ์ ํด๋นํ๋ค. ์ด ๋ฐฉ์์ ๋ชจ๋ธ์ ์คํ ๊ฐ๋ฅํ๊ฒ ๋ง๋ค์ง๋ง, GPU๊ฐ ๊ณ์ฐ์ ๊ธฐ๋ค๋ฆฌ๋ ์๊ฐ์ด ์ฆ๊ฐํ ์ ์๋ค.
๋ณธ ๋ ผ๋ฌธ์์๋ ๋ค์ ๊ด์ ์ ์ทจํ๋ค.
์์ VRAM ๋ฌธ์ ๋ฅผ ํด๊ฒฐํ๊ธฐ ์ํด์๋ ๋ฉ๋ชจ๋ฆฌ๋ฅผ ์ธ๋ถ๋ก ๋ฐ์ด๋ด๋ ๊ฒ๋ณด๋ค, ํ์ต ์ํ ์์ฒด๋ฅผ ๋ ์์ ํํ์ผ๋ก ๋ฐ๊พธ๋ ๊ฒ์ด ์ฐ์ ๋์ด์ผ ํ๋ค.
HMT๋ ๋ค์ ๋ค ๊ฐ์ง ํต์ฌ ๊ตฌ์ฑ์์๋ก ์ด๋ฃจ์ด์ง๋ค.
HMT = Dynamic Low-Rank Gradient Projection
+ Memory-Efficient Optimizer State
+ Hybrid Activation Compression
+ Hierarchical Basis and Checkpoint Storage
Transformer์ linear layer weight๋ฅผ ๋ค์๊ณผ ๊ฐ์ด ๋๋ค.
W_l โ R^{out ร in}
์ผ๋ฐ์ ์ธ ํ์ต์์๋ backward ๊ณผ์ ์์ full gradient๊ฐ ์์ฑ๋๋ค.
G_l = โL / โW_l
HMT๋ ์ด gradient๋ฅผ ๊ทธ๋๋ก optimizer์ ์ ๋ฌํ์ง ์๋๋ค. ๋์ gradient๋ฅผ ์ ์ฐจ์ ๋ถ๋ถ๊ณต๊ฐ์ผ๋ก projectํ๋ค.
G_l โ P_l ยท ฤ_l ยท Q_l^T
์ฌ๊ธฐ์:
G_l : layer l์ ์๋ gradient
P_l : output ๋ฐฉํฅ projection basis
Q_l : input ๋ฐฉํฅ projection basis
ฤ_l : low-rank gradient
rank r : r << min(out, in)
low-rank gradient๋ ๋ค์๊ณผ ๊ฐ์ด ๊ณ์ฐํ๋ค.
ฤ_l = P_l^T G_l Q_l
optimizer๋ full gradient G_l์ด ์๋๋ผ ฤ_l์ ๋ํ state๋ง ์ ์งํ๋ค. ๋ฐ๋ผ์ AdamW์ momentum, variance๋ full matrix ํฌ๊ธฐ๊ฐ ์๋๋ผ low-rank matrix ํฌ๊ธฐ๋ก ์ ์ฅ๋๋ค.
๊ธฐ์กด low-rank ๋ฐฉ์์ ๋ณดํต layer๋ง๋ค ๊ณ ์ rank๋ฅผ ์ฌ์ฉํ๋ค. ๊ทธ๋ฌ๋ ๋ชจ๋ layer์ gradient spectrum์ด ๋์ผํ์ง ์๋ค. ์ผ๋ถ layer๋ rank 32๋ก๋ ์ถฉ๋ถํ ์ ์๊ณ , ์ผ๋ถ MLP layer๋ rank 128 ์ด์์ด ํ์ํ ์ ์๋ค.
HMT๋ gradient energy ratio๋ฅผ ๊ธฐ๋ฐ์ผ๋ก layer๋ณ rank๋ฅผ ๋์ ์ผ๋ก ์ ํํ๋ค.
E(r) = sum_{i=1}^{r} ฯ_i^2 / sum_{i=1}^{n} ฯ_i^2
์ฌ๊ธฐ์ ฯ_i๋ gradient์ singular value์ด๋ค.
rank ์ ํ ๊ท์น์ ๋ค์๊ณผ ๊ฐ๋ค.
if E(64) >= ฯ:
rank = 64
elif E(128) >= ฯ:
rank = 128
elif E(256) >= ฯ:
rank = 256
else:
rank = full ๋๋ high-rank fallback๋ณดํต ฯ = 0.90 ~ 0.98 ์ฌ์ด์์ ์คํํ๋ค.
์ค์ ๊ตฌํ์์๋ ๋งค step๋ง๋ค SVD๋ฅผ ์ํํ๋ฉด ๋๋ฌด ๋๋ฆฌ๋ฏ๋ก ๋ค์์ ์ฌ์ฉํ๋ค.
- randomized SVD
- power iteration
- update interval
- gradient norm ๊ธฐ๋ฐ rank ์์ธก
- layer type๋ณ rank prior
AdamW๋ parameter๋ง๋ค ๋ณดํต ๋ค์ state๋ฅผ ์ ์ฅํ๋ค.
m_t : first moment
v_t : second moment
HMT์์๋ full weight shape์ ๋ํด m_t, v_t๋ฅผ ์ ์ฅํ์ง ์๋๋ค. ๋์ low-rank gradient ๊ณต๊ฐ์์ optimizer state๋ฅผ ์ ์ฅํ๋ค.
mฬ_t, vฬ_t โ R^{r ร r}
๋๋ APOLLO-style๋ก channel-wise ๋๋ tensor-wise scaling์ ๊ทผ์ฌํ๋ค.
APOLLO๋ AdamW์ adaptive scaling์ full parameter ๋จ์๊ฐ ์๋๋ผ ๊ตฌ์กฐํ๋ learning-rate scaling์ผ๋ก ๊ทผ์ฌํ๋ ๋ฐฉ์์ด๋ฉฐ, auxiliary low-rank state๋ฅผ ์ฌ์ฉํ๋ค.
HMT์ optimizer๋ ๋ค์ ๋ ๋ชจ๋๋ฅผ ์ง์ํ๋ค.
- Mode A: GaLore-style low-rank AdamW
- Mode B: APOLLO-style approximate gradient scaling
์ด๊ธฐ ๊ตฌํ์์๋ Mode A๊ฐ ๋ ์ง๊ด์ ์ด๋ค. ์ดํ ์ฑ๋ฅ์ด ์์ ๋๋ฉด APOLLO-style scaling์ ์ถ๊ฐํ๋ค.
ํ๋ จ ์ค activation์ backward๋ฅผ ์ํด ์ ์ฅ๋๋ค. ๊ธด sequence๋ ํฐ batch์์๋ activation์ด optimizer state๋ณด๋ค ํฐ ๋ณ๋ชฉ์ด ๋ ์ ์๋ค. CompAct๋ compressed activation์ ํตํด peak memory๋ฅผ ์ค์ผ ์ ์์์ ๋ณด์๋ค.
HMT๋ ๋ชจ๋ activation์ ๋์ผํ ์ ์ฑ ์ ์ ์ฉํ์ง ์๋๋ค. Transformer ๋ด๋ถ ๊ตฌ์ฑ์์๋ณ๋ก ๋ค๋ฅธ ์ ์ฑ ์ ์ฌ์ฉํ๋ค.
| ๊ตฌ์ฑ์์ | ์ ์ฑ |
|---|---|
| Embedding activation | keep BF16 |
| Attention Q/K/V | recompute |
| Attention output | compress FP8 ๋๋ INT8 |
| MLP intermediate | compress INT8 |
| Residual stream | keep BF16 ๋๋ FP8 |
| LayerNorm input/stat | keep BF16 |
| LM head | keep BF16 |
์ด ์ ์ฑ ์ ์ด์ ๋ ๋ค์๊ณผ ๊ฐ๋ค.
Attention activation์ ์ฌ๊ณ์ฐ ๋น์ฉ์ด ์๋์ ์ผ๋ก ํ์ฉ ๊ฐ๋ฅํ ๊ฒฝ์ฐ๊ฐ ๋ง๊ณ , MLP intermediate๋ ํฌ๊ธฐ๊ฐ ์ปค์ compression ์ด๋์ด ํฌ๋ค. LayerNorm๊ณผ residual stream์ ์์น ์์ ์ฑ์ ๋ฏผ๊ฐํ๋ฏ๋ก ๊ณผ๋ํ๊ฒ ์์ถํ์ง ์๋๋ค.
๊ธฐ์กด offload๋ weight, optimizer state, gradient๋ฅผ CPU๋ NVMe๋ก ์ด๋ํ๋ค. HMT๋ ์ด์ ๋ค๋ฅด๊ฒ CPU RAM์ projection basis cache๋ก ์ฌ์ฉํ๋ค.
์ฆ, CPU์๋ ๋ค์๋ง ์ ์ฅํ๋ค.
P_l,Q_l์ ์ค๋๋ ๋ฒ์ - layer๋ณ rank history
- gradient spectrum statistics
- optimizer low-rank history
GPU์๋ ํ์ฌ step์์ ํ์ํ basis๋ง ์ ์งํ๋ค.
์ด ๋ฐฉ์์ ์ฅ์ ์ CPUโGPU ์ ์ก๋์ด weight offload๋ณด๋ค ํจ์ฌ ์๋ค๋ ์ ์ด๋ค.
NVMe๋ ํ์ต ์ค frequent offload ๋์์ผ๋ก ์ฌ์ฉํ์ง ์๋๋ค. HMT์์ NVMe๋ ๋ค์ ์ฉ๋๋ก ์ ํํ๋ค.
- adapter ๋๋ low-rank optimizer checkpoint
- projection basis snapshot
- training state snapshot
- dataset cache
์ฆ, NVMe๋ฅผ "๋๋ฆฐ VRAM"์ผ๋ก ์ฐ์ง ์๊ณ , "์ ๋น๋ ์์ ์ ์ฅ์"๋ก๋ง ์ฌ์ฉํ๋ค.
Algorithm 1. HMT Training Loop
Input:
model ฮธ
dataset D
memory policy ฯ
optimizer ฮฉ
rank threshold ฯ
projection update interval K
Initialize:
Load model in BF16 or 8-bit
Initialize projection bases P_l, Q_l for selected layers
Initialize low-rank optimizer states
Initialize CPU basis cache
Initialize activation compression policy
For each training step t:
1. Load batch x_t
2. Forward pass:
For each layer l:
Apply layer-specific activation policy:
- keep
- recompute
- compress_int8
- compress_fp8
3. Compute loss L_t
4. Backward pass:
For each trainable linear layer l:
Compute gradient G_l
If t mod K == 0:
Estimate gradient spectrum
Select rank r_l dynamically
Update projection bases P_l, Q_l
Store old bases in CPU cache
Project gradient:
ฤ_l = P_l^T G_l Q_l
Release full gradient G_l from GPU memory
5. Optimizer step:
Update low-rank optimizer states mฬ_l, vฬ_l
Reconstruct update:
ฮW_l = P_l Update(ฤ_l) Q_l^T
Apply weight update:
W_l โ W_l - ฮท ฮW_l
6. Periodic checkpoint:
Save model delta, low-rank states, basis snapshots to NVMe
Output:
trained model ฮธ
low-rank optimizer states
final projection bases
๊ฐ์ฅ ์ค์ํ ๋ถ๋ถ์ projection basis๋ฅผ ์ด๋ป๊ฒ ๊ฐฑ์ ํ๋๋๋ค.
๋งค step๋ง๋ค full SVD๋ฅผ ํ๋ฉด ๋๋ฆฌ๋ค. ๋ฐ๋ผ์ HMT๋ K step๋ง๋ค๋ง basis๋ฅผ ๊ฐฑ์ ํ๋ค.
Input:
gradient G_l
previous basis P_l, Q_l
target energy threshold ฯ
rank candidates R = {32, 64, 128, 256}
Process:
1. Randomized range finder๋ก approximate singular vectors ๊ณ์ฐ
2. singular value energy ratio ๊ณ์ฐ
3. ์ต์ rank r_l ์ ํ
4. P_l, Q_l ๊ฐฑ์
5. ์ด์ basis๋ CPU cache๋ก ์ด๋
Pseudo-code:
@torch.no_grad()
def update_projection_basis(grad, rank_candidates, energy_threshold):
# grad: [out_dim, in_dim]
# ์ค์ ๊ตฌํ์์๋ torch.linalg.svd ๋์ randomized SVD๋ฅผ ์ฌ์ฉํด์ผ ํจ
U, S, Vh = torch.linalg.svd(grad.float(), full_matrices=False)
energy = torch.cumsum(S ** 2, dim=0) / torch.sum(S ** 2)
selected_rank = rank_candidates[-1]
for r in rank_candidates:
if energy[r - 1] >= energy_threshold:
selected_rank = r
break
P = U[:, :selected_rank].to(grad.dtype)
Q = Vh[:selected_rank, :].T.to(grad.dtype)
return P, Q, selected_rank๊ฐ๋ฐ ์ด๊ธฐ์๋ ์์ฒ๋ผ SVD๋ก ๊ฒ์ฆํ๊ณ , ์ดํ randomized SVD ๋๋ Triton kernel๋ก ์ต์ ํํ๋ค.
Low-rank gradient๋ฅผ ๊ณ์ฐํ ๋ค optimizer state๋ ฤ์ ๋ํด์๋ง ์ ์งํ๋ค.
class HMTLowRankAdamW:
def __init__(
self,
params,
lr=2e-5,
betas=(0.9, 0.95),
eps=1e-8,
weight_decay=0.01,
state_dtype=torch.float16,
):
self.params = list(params)
self.lr = lr
self.beta1, self.beta2 = betas
self.eps = eps
self.weight_decay = weight_decay
self.state_dtype = state_dtype
self.state = {}
@torch.no_grad()
def step_layer(self, weight, grad, projector):
# grad: full gradient, temporary
P, Q = projector.P, projector.Q
# project full gradient to low-rank space
low_grad = P.T @ grad @ Q
sid = id(weight)
if sid not in self.state:
self.state[sid] = {
"step": 0,
"m": torch.zeros_like(low_grad, dtype=self.state_dtype),
"v": torch.zeros_like(low_grad, dtype=self.state_dtype),
}
st = self.state[sid]
st["step"] += 1
m = st["m"]
v = st["v"]
m.mul_(self.beta1).add_(low_grad, alpha=1 - self.beta1)
v.mul_(self.beta2).addcmul_(low_grad, low_grad, value=1 - self.beta2)
m_hat = m / (1 - self.beta1 ** st["step"])
v_hat = v / (1 - self.beta2 ** st["step"])
low_update = m_hat / (torch.sqrt(v_hat.float()) + self.eps)
low_update = low_update.to(weight.dtype)
# reconstruct update to full weight shape
update = P @ low_update @ Q.T
if self.weight_decay > 0:
weight.mul_(1 - self.lr * self.weight_decay)
weight.add_(update, alpha=-self.lr)
# full grad๋ ์ฆ์ ์ ๊ฑฐ ๊ฐ๋ฅ
grad = None์ค์ ๊ตฌํ์์๋ P @ low_update @ Q.T๊ฐ ๋ณ๋ชฉ์ด ๋ ์ ์์ผ๋ฏ๋ก Triton kernel๋ก fused reconstruction-update๋ฅผ ๊ตฌํํ๋ ๊ฒ์ด ์ข๋ค.
Activation compression์ custom autograd๋ก ๊ตฌํํ๋ค.
์ด๊ธฐ ๋ฒ์ ์ INT8 block-wise quantization์ ์ฌ์ฉํ๋ค.
class BlockwiseInt8Compressor:
def __init__(self, block_size=256):
self.block_size = block_size
def compress(self, x):
orig_shape = x.shape
flat = x.reshape(-1, self.block_size)
scale = flat.abs().amax(dim=-1, keepdim=True).clamp(min=1e-6) / 127.0
q = torch.round(flat / scale).clamp(-128, 127).to(torch.int8)
meta = {
"orig_shape": orig_shape,
"scale": scale.to(torch.float16),
}
return q, meta
def decompress(self, q, meta):
x = q.float() * meta["scale"].float()
return x.reshape(meta["orig_shape"])custom linear:
class CompressedLinearFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, x, weight, bias, compressor):
y = torch.nn.functional.linear(x, weight, bias)
packed_x, meta = compressor.compress(x)
ctx.save_for_backward(weight)
ctx.packed_x = packed_x
ctx.meta = meta
ctx.compressor = compressor
ctx.has_bias = bias is not None
return y
@staticmethod
def backward(ctx, grad_y):
(weight,) = ctx.saved_tensors
x = ctx.compressor.decompress(ctx.packed_x, ctx.meta)
grad_x = grad_y @ weight
grad_w = grad_y.reshape(-1, grad_y.shape[-1]).T @ x.reshape(-1, x.shape[-1])
grad_b = None
if ctx.has_bias:
reduce_dims = tuple(range(grad_y.ndim - 1))
grad_b = grad_y.sum(dim=reduce_dims)
return grad_x, grad_w, grad_b, None์ฃผ์ํ ์ ์ ์ ์ฝ๋๋ ์ฐ๊ตฌ prototype์ด๋ค. ์ค์ LLM์ linear layer shape์ tensor layout์ ๋ง๊ฒ ์์ ํด์ผ ํ๋ฉฐ, ์ฑ๋ฅ์ ์ํด Triton kernelํ๊ฐ ํ์ํ๋ค.
์๋๋ ์ง์ ๊ฐ๋ฐ ๊ฐ๋ฅํ ๋จ๊ณ๋ณ ๊ณํ์ด๋ค.
๋ชฉํ
๊ธฐ์กด ๋ฐฉ๋ฒ๊ณผ ๋น๊ตํ ๊ธฐ์ค์ ์ ๋ง๋ ๋ค.
๊ตฌํ ์ธ์ด
- Python
- PyTorch
- Hugging Face Transformers
- YAML
๊ตฌํ ๋ด์ฉ
- Llama ๊ณ์ด 1B~3B ๋ชจ๋ธ ๋ก๋ฉ
- BF16 ํ์ต ๋ฃจํ ์์ฑ
- AdamW baseline
- QLoRA baseline
- VRAM, tokens/sec, loss curve ๊ธฐ๋ก
ํต์ฌ ์๊ณ ๋ฆฌ์ฆ
Standard causal language modeling training loop
์ฐ์ถ๋ฌผ
train_baseline.pyconfigs/baseline_adamw.yamlconfigs/baseline_qlora.yamlmemory_profiler.py
์์ ๊ตฌ์กฐ
hmt_train/
train_baseline.py
configs/
baseline_adamw.yaml
baseline_qlora.yaml
hmt/
data.py
model_loader.py
profiler.py
๋ชฉํ
full gradient๋ฅผ ์์ฑํ๋ optimizer state๋ low-rank๋ก ์ ์ฅํ๋ค.
๊ตฌํ ์ธ์ด
- Python
- PyTorch
๊ตฌํ ์๊ณ ๋ฆฌ์ฆ
- Dynamic Low-Rank Gradient Projection
- Low-Rank AdamW
๊ฐ๋ฐ ์์
nn.Linearlayer๋ง ๋์์ผ๋ก ์ ์ - backward ํ gradient hook์์
G_lํ๋ณด - SVD ๊ธฐ๋ฐ
P_l,Q_l๊ณ์ฐ ฤ_l = P_l^T G_l Q_l๊ณ์ฐ- low-rank AdamW state ์ ์ฅ
ฮW_l = P_l Update(ฤ_l) Q_l^T๋ก weight update- full gradient ์ฆ์ ์ญ์
ํต์ฌ ํ์ผ
hmt/optim/projector.pyhmt/optim/lowrank_adamw.pyhmt/trainer.py
์ฑ๊ณต ๊ธฐ์ค
- AdamW ๋๋น GPU memory ๊ฐ์
- loss๊ฐ ๋ฐ์ฐํ์ง ์์
- ์์ ๋ชจ๋ธ์์ baseline๊ณผ ์ ์ฌํ validation loss
๋ชฉํ
layer๋ณ gradient spectrum์ ๋ฐ๋ผ rank๋ฅผ ์๋ ์กฐ์ ํ๋ค.
๊ตฌํ ์ธ์ด
- Python
- PyTorch
๊ตฌํ ์๊ณ ๋ฆฌ์ฆ
- Gradient spectrum estimation
- Energy-based rank selection
- Rank scheduling
๊ฐ๋ฐ ์์
- rank ํ๋ณด๊ตฐ ์ ์:
[32, 64, 128, 256] - projection update interval K ์ค์
- K step๋ง๋ค spectrum ์ถ์
- energy threshold ฯ ๊ธฐ์ค์ผ๋ก rank ์ ํ
- layer๋ณ rank log ์ ์ฅ
ํต์ฌ ํ์ผ
hmt/optim/rank_scheduler.pyhmt/optim/spectrum.py
rank scheduler ์์
class EnergyRankScheduler:
def __init__(self, candidates=(32, 64, 128, 256), threshold=0.95):
self.candidates = candidates
self.threshold = threshold
def select_rank(self, singular_values):
energy = torch.cumsum(singular_values ** 2, dim=0)
energy = energy / energy[-1]
for r in self.candidates:
if r <= len(energy) and energy[r - 1] >= self.threshold:
return r
return self.candidates[-1]์ฑ๊ณต ๊ธฐ์ค
- ๊ณ ์ rank ๋๋น ๋น์ทํ loss
- ํ๊ท rank ๊ฐ์
- tokens/sec ๋๋ memory ํจ์จ ํฅ์
๋ชฉํ
gradient checkpointing๋ณด๋ค ๋น ๋ฅด๊ฑฐ๋, ๋น์ทํ ์๋์์ ๋ ๋ฎ์ activation memory๋ฅผ ๋ฌ์ฑํ๋ค.
๊ตฌํ ์ธ์ด
- Python
- PyTorch custom autograd
๊ตฌํ ์๊ณ ๋ฆฌ์ฆ
- Blockwise INT8 activation compression
- Hybrid activation policy
๊ฐ๋ฐ ์์
- MLP activation๋ง INT8 ์์ถ
- attention activation์ recompute ์ ์ง
- residual๊ณผ layernorm์ BF16 ์ ์ง
- custom autograd Function ์์ฑ
- memory/์๋ ๋น๊ต
ํต์ฌ ํ์ผ
hmt/memory/activation_compress.pyhmt/autograd/compressed_linear.pyhmt/memory/policy.py
์ ์ฑ ์์
activation:
embedding: keep_bf16
attention_qkv: recompute
attention_output: compress_int8
mlp_intermediate: compress_int8
residual: keep_bf16
layernorm: keep_bf16์ฑ๊ณต ๊ธฐ์ค
- no compression ๋๋น peak VRAM ๊ฐ์
- gradient checkpointing ๋๋น step time ๊ฐ์ ๋๋ ์ ์ฌ
- loss degradation ํ์ฉ ๋ฒ์ ๋ด ์ ์ง
๋ชฉํ
CPU RAM์ weight offload๊ฐ ์๋๋ผ projection basis cache๋ก ์ฌ์ฉํ๋ค.
๊ตฌํ ์ธ์ด
- Python
- PyTorch
๊ตฌํ ์๊ณ ๋ฆฌ์ฆ
- Basis lifecycle management
- Asynchronous CPU-GPU transfer
- Pinned memory staging
๊ฐ๋ฐ ์์
- ์ค๋๋
P_l,Q_l์ CPU pinned memory๋ก ์ด๋ - ํ์ฌ step์ ํ์ํ basis๋ง GPU์ ์ ์ง
- layer๋ณ basis cache hit/miss ๊ธฐ๋ก
- basis prefetch ๊ตฌํ
ํต์ฌ ํ์ผ
hmt/memory/cpu_basis_cache.pyhmt/memory/staging.py
basis cache ์์
class CPUBasisCache:
def __init__(self, max_entries=1024):
self.cache = {}
self.max_entries = max_entries
def put(self, key, P, Q):
self.cache[key] = {
"P": P.detach().to("cpu", non_blocking=True).pin_memory(),
"Q": Q.detach().to("cpu", non_blocking=True).pin_memory(),
}
def get(self, key, device):
item = self.cache[key]
P = item["P"].to(device, non_blocking=True)
Q = item["Q"].to(device, non_blocking=True)
return P, Q์ฑ๊ณต ๊ธฐ์ค
- GPU์ ์ ์ง๋๋ basis memory ๊ฐ์
- CPU transfer overhead๊ฐ ์ ์ฒด step time์ ์์ ๋น์จ๋ก ์ ์ง
๋ชฉํ
Python/PyTorch prototype์์ ๋ณ๋ชฉ์ด ๋๋ ์ฐ์ฐ์ GPU kernel๋ก ์ต์ ํํ๋ค.
๊ตฌํ ์ธ์ด
- Python
- Triton
์ต์ ํ ๋์
- blockwise activation quantization
- blockwise activation dequantization
- low-rank projection matmul
- fused reconstruction + weight update
- optimizer state update
ํต์ฌ ํ์ผ
hmt/kernels/quantize.pyhmt/kernels/dequantize.pyhmt/kernels/lowrank_update.py
์ฐ์ ์์
- activation compress/decompress
- reconstructed update ์ ์ฉ
- low-rank gradient projection
- fused optimizer update
์ฑ๊ณต ๊ธฐ์ค
- Python implementation ๋๋น tokens/sec ํฅ์
- GPU utilization ์์น
torch.profiler์์ kernel launch overhead ๊ฐ์
๋ชฉํ
GaLore-style optimizer ์ธ์ APOLLO-style approximate gradient scaling์ ๊ตฌํํ๋ค.
๊ตฌํ ์ธ์ด
- Python
- PyTorch
- Triton (optional)
๊ตฌํ ์๊ณ ๋ฆฌ์ฆ
- Approximated Gradient Scaling
- Channel-wise learning-rate scaling
- Low-rank auxiliary optimizer state
๊ฐ๋ฐ ์์
- tensor-wise scaling ๋ฒ์ ๊ตฌํ
- channel-wise scaling ๋ฒ์ ๊ตฌํ
- rank-1 APOLLO-Mini ์คํ์ผ ๊ตฌํ
- GaLore-style optimizer์ ๋น๊ต
ํต์ฌ ํ์ผ
hmt/optim/apollo.pyhmt/optim/scaling.py
์ฑ๊ณต ๊ธฐ์ค
- GaLore๋ณด๋ค optimizer state memory ๊ฐ์
- AdamW ๋๋ GaLore์ ์ ์ฌํ loss curve
- ํ์ต ์์ ์ฑ ํ๋ณด
๋ชฉํ
๋ชจ๋ ์ ์ฑ ์ YAML config๋ก ์ ์ดํ ์ ์๋ ํ์ต ํ๋ ์์ํฌ๋ฅผ ๋ง๋ ๋ค.
๊ตฌํ ์ธ์ด
- Python
- PyTorch
- YAML
- OmegaConf
๊ตฌํ ๋ด์ฉ
- model loading
- dataset packing
- dynamic rank optimizer
- activation compression
- CPU basis cache
- NVMe checkpoint-only save
- profiling
- experiment logging
ํ๋ก์ ํธ ๊ตฌ์กฐ
hmt_train/
train.py
configs/
hmt_1b.yaml
hmt_3b.yaml
hmt_7b.yaml
hmt/
model_loader.py
data.py
trainer.py
optim/
projector.py
lowrank_adamw.py
rank_scheduler.py
apollo.py
memory/
activation_compress.py
cpu_basis_cache.py
policy.py
checkpoint.py
autograd/
compressed_linear.py
kernels/
quantize.py
dequantize.py
lowrank_update.py
utils/
profiler.py
logging.py
- Baseline A: BF16 AdamW
- Baseline B: QLoRA
- Baseline C: GaLore fixed-rank
- Baseline D: APOLLO
- Proposed: HMT dynamic-rank + hybrid activation compression
- Peak GPU memory
- Average GPU memory
- tokens/sec
- step time
- validation loss
- CPU RAM usage
- PCIe traffic
- checkpoint size
- rank distribution
- activation compression error
์ด๊ธฐ ์คํ์ ์์ ๋ชจ๋ธ๋ถํฐ ์์ํด์ผ ํ๋ค.
- Phase 1: 125M ~ 350M
- Phase 2: 1B
- Phase 3: 3B
- Phase 4: 7B
- Phase 5: 13B ์ด์
์ฒ์๋ถํฐ 7B ์ด์์ผ๋ก ๊ฐ๋ฉด ๋๋ฒ๊น ๋น์ฉ์ด ๋๋ฌด ํฌ๋ค.
| ์์ญ | ์ธ์ด/๋๊ตฌ | ์ด์ |
|---|---|---|
| ํ์ต ๋ฃจํ | Python | ์คํ ์๋ |
| ๋ชจ๋ธ ์คํ | PyTorch | LLM ์ํ๊ณ ํธํ์ฑ |
| config | YAML + OmegaConf | ์คํ ์ฌํ์ฑ |
| optimizer prototype | Python + PyTorch | ๋๋ฒ๊น ์ฉ์ด |
| activation compression | PyTorch custom autograd | backward ์ ์ด |
| ๊ณ ์ฑ๋ฅ kernel | Triton | CUDA๋ณด๋ค ๋น ๋ฅธ ๊ฐ๋ฐ |
| ์ต์ข ๊ทนํ ์ต์ ํ | CUDA C++ | ํ์ํ ๋๋ง |
| ๋ฐ์ดํฐ ์ ์ฒ๋ฆฌ | Python ๋๋ Rust | ๋์ฉ๋์ด๋ฉด Rust ๊ณ ๋ ค |
| ๋ก๊น | TensorBoard / W&B | ์คํ ์ถ์ |
| profiling | torch.profiler, Nsight Systems | ๋ณ๋ชฉ ๋ถ์ |
- Python/PyTorch๋ก correctness ๊ฒ์ฆ
- ์์ ๋ชจ๋ธ์์ loss curve ํ์ธ
- memory profiler๋ก ์ค์ ์ ๊ฐ ํ์ธ
- ๋ณ๋ชฉ kernel๋ง Triton์ผ๋ก ์ด๋
- APOLLO-style scaling ์ถ๊ฐ
- CPU basis cache ์ถ๊ฐ
- 3B/7B ๋ชจ๋ธ๋ก ํ์ฅ
- ํ์ํ ๋๋ง CUDA extension ์์ฑ
๋ณธ ๋ฐฉ๋ฒ์ ์์ ๊ธฐ์ฌ์ ์ ๋ค์๊ณผ ๊ฐ๋ค.
์ฒซ์งธ, CPU/NVMe offload ์ค์ฌ์ด ์๋๋ผ ํ์ต ์ํ ํํ ์์ฒด๋ฅผ ์์ถํ๋ค.
๋์งธ, LoRA์ฒ๋ผ weight update ๊ณต๊ฐ์ ๊ณ ์ adapter๋ก ์ ํํ์ง ์๊ณ , full-parameter update์ ๊ฐ๊น์ด gradient-space update๋ฅผ ์ํํ๋ค.
์ ์งธ, ๋ชจ๋ layer์ ๊ฐ์ ๋ฉ๋ชจ๋ฆฌ ์ ์ฑ ์ ์ ์ฉํ์ง ์๊ณ , attention, MLP, residual, layernorm, lm_head์ ์๋ก ๋ค๋ฅธ ์ ์ฑ ์ ์ ์ฉํ๋ค.
๋ท์งธ, CPU RAM์ weight offload ์ ์ฅ์๊ฐ ์๋๋ผ projection basis cache๋ก ์ฌ์ฉํ๋ค.
๋ค์ฏ์งธ, NVMe๋ ํ์ต ์ค ๋น๋ฒํ swap ๊ณต๊ฐ์ด ์๋๋ผ checkpoint-only ์ ์ฅ์๋ก ์ ํํ๋ค.
activation compression๊ณผ low-rank gradient projection์ ๋ชจ๋ ํ์ต ํ์ง์ ์ํฅ์ ์ค ์ ์๋ค. ํนํ LayerNorm, residual stream, lm_head๋ ์์ถ์ ๋ฏผ๊ฐํ ์ ์์ผ๋ฏ๋ก BF16 ์ ์ง๊ฐ ์์ ํ๋ค.
basis ๊ฐฑ์ ์ SVD๋ฅผ ์ฌ์ฉํ๋ฉด ๋งค์ฐ ๋๋ฆด ์ ์๋ค. ๋ฐ๋ผ์ ์ค์ ๊ตฌํ์์๋ randomized SVD, power iteration, update interval์ด ํ์๋ค.
low-rank update๋ฅผ ๋ค์ full weight shape์ผ๋ก ๋ณต์ํ๋ ๊ณผ์ ์ด ๋ณ๋ชฉ์ด ๋ ์ ์๋ค. ์ด ๋ถ๋ถ์ Triton fused kernel๋ก ์ต์ ํํด์ผ ํ๋ค.
Hugging Face Trainer๋ฅผ ๊ทธ๋๋ก ์ฌ์ฉํ๊ธฐ ์ด๋ ต๋ค. ์ง์ training loop, optimizer step, gradient hook, activation hook์ ์ ์ดํด์ผ ํ๋ค.
๋ณธ ๋ ผ๋ฌธ์ ์์ VRAM ํ๊ฒฝ์์ ๊ฑฐ๋ ์ธ์ด๋ชจ๋ธ์ ํ์ตํ๊ธฐ ์ํ ๊ณ์ธตํ ๋ฉ๋ชจ๋ฆฌ ํ์ต ์๊ณ ๋ฆฌ์ฆ HMT๋ฅผ ์ ์ํ์๋ค. ๊ธฐ์กด ๋ฐฉ์์ ์ฃผ๋ก ๋ชจ๋ธ ์ํ๋ฅผ CPU๋ NVMe๋ก ์ด๋ํ๊ฑฐ๋, ํ์ต ๊ฐ๋ฅํ ํ๋ผ๋ฏธํฐ ์๋ฅผ ์ค์ด๋ ๋ฐฉ์์ ์์กดํ๋ค. ๋ฐ๋ฉด HMT๋ gradient, optimizer state, activation์ ํํ ์์ฒด๋ฅผ ์ค์ฌ GPU VRAM ์๊ตฌ๋์ ๋ฎ์ถ๋ค.
ํต์ฌ ์์ด๋์ด๋ ๋ค์๊ณผ ๊ฐ๋ค.
- gradient๋ full matrix๋ก ์ค๋ ๋ณด๊ดํ์ง ์๊ณ low-rank space๋ก projectํ๋ค.
- optimizer state๋ full parameter shape์ด ์๋๋ผ low-rank shape์ผ๋ก ์ ์งํ๋ค.
- activation์ layer type๋ณ๋ก keep, recompute, compress๋ฅผ ๋ค๋ฅด๊ฒ ์ ์ฉํ๋ค.
- CPU RAM์ weight offload๊ฐ ์๋๋ผ projection basis cache๋ก ์ฌ์ฉํ๋ค.
- NVMe๋ checkpoint-only ์ ์ฅ์๋ก ์ ํํ๋ค.
- ์ฑ๋ฅ ๋ณ๋ชฉ์ Triton kernel๋ก ๋จ๊ณ์ ์ผ๋ก ์ต์ ํํ๋ค.
์ง์ ๊ฐ๋ฐ ๊ด์ ์์๋ Python + PyTorch๋ก prototype์ ๋ง๋ค๊ณ , ๋ณ๋ชฉ์ด ํ์ธ๋ compression/projection/update ์ฐ์ฐ๋ง Triton์ผ๋ก ์ฎ๊ธฐ๋ ๋ฐฉ์์ด ๊ฐ์ฅ ํ์ค์ ์ด๋ค. CUDA C++๋ ์ต์ข ๋จ๊ณ์์๋ง ํ์ํ๋ค.
| Stage | ๋ด์ฉ | ์ธ์ด | ์๊ณ ๋ฆฌ์ฆ |
|---|---|---|---|
| Stage 0 | Baseline ๊ตฌ์ถ | Python, PyTorch | AdamW, QLoRA |
| Stage 1 | Low-rank optimizer ๊ตฌํ | Python, PyTorch | GaLore-style gradient projection |
| Stage 2 | Dynamic rank ์ถ๊ฐ | Python, PyTorch | energy-based rank scheduling |
| Stage 3 | Activation compression ์ถ๊ฐ | Python, PyTorch custom autograd | blockwise INT8/FP8 activation compression |
| Stage 4 | CPU basis cache ์ถ๊ฐ | Python, PyTorch | pinned memory basis staging |
| Stage 5 | Triton ์ต์ ํ | Python, Triton | fused quantize/dequantize/projection/update |
| Stage 6 | APOLLO-style optimizer ์ถ๊ฐ | Python, PyTorch, Triton (optional) | approximate gradient scaling |
| Stage 7 | HMT ํตํฉ trainer ์์ฑ | Python, PyTorch, YAML | dynamic low-rank optimizer + hybrid activation policy |