Skip to content

Folders and files

NameName
Last commit message
Last commit date

Latest commit

ย 

History

19 Commits
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 
ย 

Repository files navigation

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์— ๊ฐ€๊นŒ์šด ํ•™์Šต ์ž์œ ๋„๋ฅผ ํ™•๋ณดํ•˜๋Š” ๊ฒƒ์ด๋‹ค.


1. ์„œ๋ก 

๊ฑฐ๋Œ€ ์–ธ์–ด๋ชจ๋ธ์€ ํŒŒ๋ผ๋ฏธํ„ฐ ์ˆ˜๊ฐ€ ์ฆ๊ฐ€ํ• ์ˆ˜๋ก ์ถ”๋ก ๊ณผ ํ•™์Šต ๋ชจ๋‘์—์„œ ๋ฉ”๋ชจ๋ฆฌ ์š”๊ตฌ๋Ÿ‰์ด ๊ธ‰๊ฒฉํžˆ ์ฆ๊ฐ€ํ•œ๋‹ค. ํŠนํžˆ ํ•™์Šต ๋‹จ๊ณ„์—์„œ๋Š” ๋‹จ์ˆœํžˆ ๋ชจ๋ธ 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์„ ํ•™์Šตํ•˜๊ธฐ ์œ„ํ•œ ์ƒˆ๋กœ์šด ํ†ตํ•ฉ ์•Œ๊ณ ๋ฆฌ์ฆ˜์„ ์ œ์•ˆํ•œ๋‹ค.


2. ๋ฌธ์ œ ์ •์˜

2.1 ๋ชฉํ‘œ

๋ณธ ์—ฐ๊ตฌ์˜ ๋ชฉํ‘œ๋Š” ๋‹ค์Œ๊ณผ ๊ฐ™๋‹ค.

์ž‘์€ VRAM ํ™˜๊ฒฝ์—์„œ ๊ฑฐ๋Œ€ LLM์„ ํ•™์Šตํ•˜๋˜, ๋‹จ์ˆœ CPU/NVMe offload์— ์˜์กดํ•˜์ง€ ์•Š๊ณ  ํ•™์Šต ์ค‘ ์ƒ์„ฑ๋˜๋Š” ์ฃผ์š” ๋ฉ”๋ชจ๋ฆฌ ํ•ญ๋ชฉ์„ ์ˆ˜ํ•™์ ์œผ๋กœ ์••์ถ•ํ•œ๋‹ค.

๋Œ€์ƒ ๋ฉ”๋ชจ๋ฆฌ ํ•ญ๋ชฉ์€ ๋‹ค์Œ๊ณผ ๊ฐ™๋‹ค.

  1. Model parameters
  2. Gradients
  3. Optimizer states
  4. Activations
  5. Temporary buffers

์šฐ๋ฆฌ๊ฐ€ ์ตœ์ ํ™”ํ•˜๋ ค๋Š” ๋ชฉ์  ํ•จ์ˆ˜๋Š” ๋‹ค์Œ๊ณผ ๊ฐ™์ด ์ •์˜ํ•  ์ˆ˜ ์žˆ๋‹ค.

minimize   peak_gpu_memory
maximize   tokens_per_second
maintain   validation_loss_quality

์ฆ‰, ๋‹จ์ˆœํžˆ ๋ฉ”๋ชจ๋ฆฌ๋ฅผ ์ค„์ด๋Š” ๊ฒƒ์ด ์•„๋‹ˆ๋ผ ๋ฉ”๋ชจ๋ฆฌ, ์†๋„, ํ•™์Šต ํ’ˆ์งˆ์˜ ๊ท ํ˜•์„ ์ตœ์ ํ™”ํ•œ๋‹ค.

2.2 ๊ธฐ์กด ๋ฐฉ์‹์˜ ํ•œ๊ณ„

๊ธฐ์กด ์ž‘์€ VRAM ํ•™์Šต ๋ฐฉ์‹์€ ํฌ๊ฒŒ ๋‘ ๋ถ€๋ฅ˜๋กœ ๋‚˜๋‰œ๋‹ค.

์ฒซ ๋ฒˆ์งธ๋Š” ํ•™์Šต ๋ฒ”์œ„๋ฅผ ์ค„์ด๋Š” ๋ฐฉ์‹์ด๋‹ค. LoRA/QLoRA๊ฐ€ ๋Œ€ํ‘œ์ ์ด๋‹ค. ์ด ๋ฐฉ์‹์€ ๋งค์šฐ ์‹ค์šฉ์ ์ด์ง€๋งŒ, ์›๋ณธ weight ์ „์ฒด๋ฅผ ์ง์ ‘ ์—…๋ฐ์ดํŠธํ•˜์ง€ ์•Š๋Š”๋‹ค.

๋‘ ๋ฒˆ์งธ๋Š” ํ•™์Šต ์ƒํƒœ๋ฅผ ์™ธ๋ถ€ ๋ฉ”๋ชจ๋ฆฌ๋กœ ์ด๋™ํ•˜๋Š” ๋ฐฉ์‹์ด๋‹ค. ZeRO-Offload, ZeRO-Infinity, CPU/NVMe offload๊ฐ€ ์—ฌ๊ธฐ์— ํ•ด๋‹นํ•œ๋‹ค. ์ด ๋ฐฉ์‹์€ ๋ชจ๋ธ์„ ์‹คํ–‰ ๊ฐ€๋Šฅํ•˜๊ฒŒ ๋งŒ๋“ค์ง€๋งŒ, GPU๊ฐ€ ๊ณ„์‚ฐ์„ ๊ธฐ๋‹ค๋ฆฌ๋Š” ์‹œ๊ฐ„์ด ์ฆ๊ฐ€ํ•  ์ˆ˜ ์žˆ๋‹ค.

๋ณธ ๋…ผ๋ฌธ์—์„œ๋Š” ๋‹ค์Œ ๊ด€์ ์„ ์ทจํ•œ๋‹ค.

์ž‘์€ VRAM ๋ฌธ์ œ๋ฅผ ํ•ด๊ฒฐํ•˜๊ธฐ ์œ„ํ•ด์„œ๋Š” ๋ฉ”๋ชจ๋ฆฌ๋ฅผ ์™ธ๋ถ€๋กœ ๋ฐ€์–ด๋‚ด๋Š” ๊ฒƒ๋ณด๋‹ค, ํ•™์Šต ์ƒํƒœ ์ž์ฒด๋ฅผ ๋” ์ž‘์€ ํ‘œํ˜„์œผ๋กœ ๋ฐ”๊พธ๋Š” ๊ฒƒ์ด ์šฐ์„ ๋˜์–ด์•ผ ํ•œ๋‹ค.


3. ์ œ์•ˆ ๋ฐฉ๋ฒ•: HMT

HMT๋Š” ๋‹ค์Œ ๋„ค ๊ฐ€์ง€ ํ•ต์‹ฌ ๊ตฌ์„ฑ์š”์†Œ๋กœ ์ด๋ฃจ์–ด์ง„๋‹ค.

HMT = Dynamic Low-Rank Gradient Projection
    + Memory-Efficient Optimizer State
    + Hybrid Activation Compression
    + Hierarchical Basis and Checkpoint Storage

3.1 Dynamic Low-Rank Gradient Projection

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 ํฌ๊ธฐ๋กœ ์ €์žฅ๋œ๋‹ค.

3.2 ๋™์  rank ์„ ํƒ

๊ธฐ์กด 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๋ฅผ ์ˆ˜ํ–‰ํ•˜๋ฉด ๋„ˆ๋ฌด ๋А๋ฆฌ๋ฏ€๋กœ ๋‹ค์Œ์„ ์‚ฌ์šฉํ•œ๋‹ค.

  1. randomized SVD
  2. power iteration
  3. update interval
  4. gradient norm ๊ธฐ๋ฐ˜ rank ์˜ˆ์ธก
  5. layer type๋ณ„ rank prior

3.3 Memory-Efficient Optimizer State

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์„ ์ถ”๊ฐ€ํ•œ๋‹ค.

3.4 Hybrid Activation Compression

ํ›ˆ๋ จ ์ค‘ 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์€ ์ˆ˜์น˜ ์•ˆ์ •์„ฑ์— ๋ฏผ๊ฐํ•˜๋ฏ€๋กœ ๊ณผ๋„ํ•˜๊ฒŒ ์••์ถ•ํ•˜์ง€ ์•Š๋Š”๋‹ค.

3.5 CPU Basis Cache

๊ธฐ์กด 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๋ณด๋‹ค ํ›จ์”ฌ ์ž‘๋‹ค๋Š” ์ ์ด๋‹ค.

3.6 NVMe Checkpoint-Only Policy

NVMe๋Š” ํ•™์Šต ์ค‘ frequent offload ๋Œ€์ƒ์œผ๋กœ ์‚ฌ์šฉํ•˜์ง€ ์•Š๋Š”๋‹ค. HMT์—์„œ NVMe๋Š” ๋‹ค์Œ ์šฉ๋„๋กœ ์ œํ•œํ•œ๋‹ค.

  1. adapter ๋˜๋Š” low-rank optimizer checkpoint
  2. projection basis snapshot
  3. training state snapshot
  4. dataset cache

์ฆ‰, NVMe๋ฅผ "๋А๋ฆฐ VRAM"์œผ๋กœ ์“ฐ์ง€ ์•Š๊ณ , "์ €๋นˆ๋„ ์˜์† ์ €์žฅ์†Œ"๋กœ๋งŒ ์‚ฌ์šฉํ•œ๋‹ค.


4. HMT ์ „์ฒด ์•Œ๊ณ ๋ฆฌ์ฆ˜

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

5. ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜ ์ƒ์„ธ ์„ค๋ช…

5.1 Projection Basis Update Algorithm

๊ฐ€์žฅ ์ค‘์š”ํ•œ ๋ถ€๋ถ„์€ 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๋กœ ์ตœ์ ํ™”ํ•œ๋‹ค.

5.2 Low-Rank Optimizer Algorithm

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๋ฅผ ๊ตฌํ˜„ํ•˜๋Š” ๊ฒƒ์ด ์ข‹๋‹ค.

5.3 Activation Compression Algorithm

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ํ™”๊ฐ€ ํ•„์š”ํ•˜๋‹ค.


6. ์Šคํ…Œ์ด์ง€๋ณ„ ๊ฐœ๋ฐœ ๊ณ„ํš

์•„๋ž˜๋Š” ์ง์ ‘ ๊ฐœ๋ฐœ ๊ฐ€๋Šฅํ•œ ๋‹จ๊ณ„๋ณ„ ๊ณ„ํš์ด๋‹ค.

Stage 0. ์‹คํ—˜ ๊ธฐ์ค€์„  ๊ตฌ์ถ•

๋ชฉํ‘œ

๊ธฐ์กด ๋ฐฉ๋ฒ•๊ณผ ๋น„๊ตํ•  ๊ธฐ์ค€์„ ์„ ๋งŒ๋“ ๋‹ค.

๊ตฌํ˜„ ์–ธ์–ด

  • Python
  • PyTorch
  • Hugging Face Transformers
  • YAML

๊ตฌํ˜„ ๋‚ด์šฉ

  1. Llama ๊ณ„์—ด 1B~3B ๋ชจ๋ธ ๋กœ๋”ฉ
  2. BF16 ํ•™์Šต ๋ฃจํ”„ ์ž‘์„ฑ
  3. AdamW baseline
  4. QLoRA baseline
  5. VRAM, tokens/sec, loss curve ๊ธฐ๋ก

ํ•ต์‹ฌ ์•Œ๊ณ ๋ฆฌ์ฆ˜

Standard causal language modeling training loop

์‚ฐ์ถœ๋ฌผ

  • train_baseline.py
  • configs/baseline_adamw.yaml
  • configs/baseline_qlora.yaml
  • memory_profiler.py

์˜ˆ์‹œ ๊ตฌ์กฐ

hmt_train/
  train_baseline.py
  configs/
    baseline_adamw.yaml
    baseline_qlora.yaml
  hmt/
    data.py
    model_loader.py
    profiler.py

Stage 1. GaLore-style low-rank optimizer ๊ตฌํ˜„

๋ชฉํ‘œ

full gradient๋ฅผ ์ƒ์„ฑํ•˜๋˜ optimizer state๋Š” low-rank๋กœ ์ €์žฅํ•œ๋‹ค.

๊ตฌํ˜„ ์–ธ์–ด

  • Python
  • PyTorch

๊ตฌํ˜„ ์•Œ๊ณ ๋ฆฌ์ฆ˜

  • Dynamic Low-Rank Gradient Projection
  • Low-Rank AdamW

๊ฐœ๋ฐœ ์ˆœ์„œ

  1. nn.Linear layer๋งŒ ๋Œ€์ƒ์œผ๋กœ ์„ ์ •
  2. backward ํ›„ gradient hook์—์„œ G_l ํ™•๋ณด
  3. SVD ๊ธฐ๋ฐ˜ P_l, Q_l ๊ณ„์‚ฐ
  4. ฤœ_l = P_l^T G_l Q_l ๊ณ„์‚ฐ
  5. low-rank AdamW state ์ €์žฅ
  6. ฮ”W_l = P_l Update(ฤœ_l) Q_l^T๋กœ weight update
  7. full gradient ์ฆ‰์‹œ ์‚ญ์ œ

ํ•ต์‹ฌ ํŒŒ์ผ

  • hmt/optim/projector.py
  • hmt/optim/lowrank_adamw.py
  • hmt/trainer.py

์„ฑ๊ณต ๊ธฐ์ค€

  • AdamW ๋Œ€๋น„ GPU memory ๊ฐ์†Œ
  • loss๊ฐ€ ๋ฐœ์‚ฐํ•˜์ง€ ์•Š์Œ
  • ์ž‘์€ ๋ชจ๋ธ์—์„œ baseline๊ณผ ์œ ์‚ฌํ•œ validation loss

Stage 2. Dynamic Rank Selection ์ถ”๊ฐ€

๋ชฉํ‘œ

layer๋ณ„ gradient spectrum์— ๋”ฐ๋ผ rank๋ฅผ ์ž๋™ ์กฐ์ ˆํ•œ๋‹ค.

๊ตฌํ˜„ ์–ธ์–ด

  • Python
  • PyTorch

๊ตฌํ˜„ ์•Œ๊ณ ๋ฆฌ์ฆ˜

  • Gradient spectrum estimation
  • Energy-based rank selection
  • Rank scheduling

๊ฐœ๋ฐœ ์ˆœ์„œ

  1. rank ํ›„๋ณด๊ตฐ ์ •์˜: [32, 64, 128, 256]
  2. projection update interval K ์„ค์ •
  3. K step๋งˆ๋‹ค spectrum ์ถ”์ •
  4. energy threshold ฯ„ ๊ธฐ์ค€์œผ๋กœ rank ์„ ํƒ
  5. layer๋ณ„ rank log ์ €์žฅ

ํ•ต์‹ฌ ํŒŒ์ผ

  • hmt/optim/rank_scheduler.py
  • hmt/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 ํšจ์œจ ํ–ฅ์ƒ

Stage 3. Activation Compression prototype

๋ชฉํ‘œ

gradient checkpointing๋ณด๋‹ค ๋น ๋ฅด๊ฑฐ๋‚˜, ๋น„์Šทํ•œ ์†๋„์—์„œ ๋” ๋‚ฎ์€ activation memory๋ฅผ ๋‹ฌ์„ฑํ•œ๋‹ค.

๊ตฌํ˜„ ์–ธ์–ด

  • Python
  • PyTorch custom autograd

๊ตฌํ˜„ ์•Œ๊ณ ๋ฆฌ์ฆ˜

  • Blockwise INT8 activation compression
  • Hybrid activation policy

๊ฐœ๋ฐœ ์ˆœ์„œ

  1. MLP activation๋งŒ INT8 ์••์ถ•
  2. attention activation์€ recompute ์œ ์ง€
  3. residual๊ณผ layernorm์€ BF16 ์œ ์ง€
  4. custom autograd Function ์ž‘์„ฑ
  5. memory/์†๋„ ๋น„๊ต

ํ•ต์‹ฌ ํŒŒ์ผ

  • hmt/memory/activation_compress.py
  • hmt/autograd/compressed_linear.py
  • hmt/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 ํ—ˆ์šฉ ๋ฒ”์œ„ ๋‚ด ์œ ์ง€

Stage 4. CPU Basis Cache ๊ตฌํ˜„

๋ชฉํ‘œ

CPU RAM์„ weight offload๊ฐ€ ์•„๋‹ˆ๋ผ projection basis cache๋กœ ์‚ฌ์šฉํ•œ๋‹ค.

๊ตฌํ˜„ ์–ธ์–ด

  • Python
  • PyTorch

๊ตฌํ˜„ ์•Œ๊ณ ๋ฆฌ์ฆ˜

  • Basis lifecycle management
  • Asynchronous CPU-GPU transfer
  • Pinned memory staging

๊ฐœ๋ฐœ ์ˆœ์„œ

  1. ์˜ค๋ž˜๋œ P_l, Q_l์„ CPU pinned memory๋กœ ์ด๋™
  2. ํ˜„์žฌ step์— ํ•„์š”ํ•œ basis๋งŒ GPU์— ์œ ์ง€
  3. layer๋ณ„ basis cache hit/miss ๊ธฐ๋ก
  4. basis prefetch ๊ตฌํ˜„

ํ•ต์‹ฌ ํŒŒ์ผ

  • hmt/memory/cpu_basis_cache.py
  • hmt/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์˜ ์ž‘์€ ๋น„์œจ๋กœ ์œ ์ง€

Stage 5. Triton kernel ์ตœ์ ํ™”

๋ชฉํ‘œ

Python/PyTorch prototype์—์„œ ๋ณ‘๋ชฉ์ด ๋˜๋Š” ์—ฐ์‚ฐ์„ GPU kernel๋กœ ์ตœ์ ํ™”ํ•œ๋‹ค.

๊ตฌํ˜„ ์–ธ์–ด

  • Python
  • Triton

์ตœ์ ํ™” ๋Œ€์ƒ

  1. blockwise activation quantization
  2. blockwise activation dequantization
  3. low-rank projection matmul
  4. fused reconstruction + weight update
  5. optimizer state update

ํ•ต์‹ฌ ํŒŒ์ผ

  • hmt/kernels/quantize.py
  • hmt/kernels/dequantize.py
  • hmt/kernels/lowrank_update.py

์šฐ์„ ์ˆœ์œ„

  1. activation compress/decompress
  2. reconstructed update ์ ์šฉ
  3. low-rank gradient projection
  4. fused optimizer update

์„ฑ๊ณต ๊ธฐ์ค€

  • Python implementation ๋Œ€๋น„ tokens/sec ํ–ฅ์ƒ
  • GPU utilization ์ƒ์Šน
  • torch.profiler์—์„œ kernel launch overhead ๊ฐ์†Œ

Stage 6. APOLLO-style optimizer ์ถ”๊ฐ€

๋ชฉํ‘œ

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

๊ฐœ๋ฐœ ์ˆœ์„œ

  1. tensor-wise scaling ๋ฒ„์ „ ๊ตฌํ˜„
  2. channel-wise scaling ๋ฒ„์ „ ๊ตฌํ˜„
  3. rank-1 APOLLO-Mini ์Šคํƒ€์ผ ๊ตฌํ˜„
  4. GaLore-style optimizer์™€ ๋น„๊ต

ํ•ต์‹ฌ ํŒŒ์ผ

  • hmt/optim/apollo.py
  • hmt/optim/scaling.py

์„ฑ๊ณต ๊ธฐ์ค€

  • GaLore๋ณด๋‹ค optimizer state memory ๊ฐ์†Œ
  • AdamW ๋˜๋Š” GaLore์™€ ์œ ์‚ฌํ•œ loss curve
  • ํ•™์Šต ์•ˆ์ •์„ฑ ํ™•๋ณด

Stage 7. ํ†ตํ•ฉ HMT Trainer ๊ฐœ๋ฐœ

๋ชฉํ‘œ

๋ชจ๋“  ์ •์ฑ…์„ YAML config๋กœ ์ œ์–ดํ•  ์ˆ˜ ์žˆ๋Š” ํ•™์Šต ํ”„๋ ˆ์ž„์›Œํฌ๋ฅผ ๋งŒ๋“ ๋‹ค.

๊ตฌํ˜„ ์–ธ์–ด

  • Python
  • PyTorch
  • YAML
  • OmegaConf

๊ตฌํ˜„ ๋‚ด์šฉ

  1. model loading
  2. dataset packing
  3. dynamic rank optimizer
  4. activation compression
  5. CPU basis cache
  6. NVMe checkpoint-only save
  7. profiling
  8. 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

7. ์‹คํ—˜ ์„ค๊ณ„

7.1 ๋น„๊ต ๋Œ€์ƒ

  • Baseline A: BF16 AdamW
  • Baseline B: QLoRA
  • Baseline C: GaLore fixed-rank
  • Baseline D: APOLLO
  • Proposed: HMT dynamic-rank + hybrid activation compression

7.2 ์ธก์ • ์ง€ํ‘œ

  1. Peak GPU memory
  2. Average GPU memory
  3. tokens/sec
  4. step time
  5. validation loss
  6. CPU RAM usage
  7. PCIe traffic
  8. checkpoint size
  9. rank distribution
  10. activation compression error

7.3 ์‹คํ—˜ ๋ชจ๋ธ

์ดˆ๊ธฐ ์‹คํ—˜์€ ์ž‘์€ ๋ชจ๋ธ๋ถ€ํ„ฐ ์‹œ์ž‘ํ•ด์•ผ ํ•œ๋‹ค.

  • Phase 1: 125M ~ 350M
  • Phase 2: 1B
  • Phase 3: 3B
  • Phase 4: 7B
  • Phase 5: 13B ์ด์ƒ

์ฒ˜์Œ๋ถ€ํ„ฐ 7B ์ด์ƒ์œผ๋กœ ๊ฐ€๋ฉด ๋””๋ฒ„๊น… ๋น„์šฉ์ด ๋„ˆ๋ฌด ํฌ๋‹ค.


8. ๊ฐœ๋ฐœ ์Šคํƒ ์š”์•ฝ

8.1 ์–ธ์–ด ์„ ํƒ

์˜์—ญ ์–ธ์–ด/๋„๊ตฌ ์ด์œ 
ํ•™์Šต ๋ฃจํ”„ 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 ๋ณ‘๋ชฉ ๋ถ„์„

8.2 ๊ถŒ์žฅ ๊ตฌํ˜„ ์ˆœ์„œ

  1. Python/PyTorch๋กœ correctness ๊ฒ€์ฆ
  2. ์ž‘์€ ๋ชจ๋ธ์—์„œ loss curve ํ™•์ธ
  3. memory profiler๋กœ ์‹ค์ œ ์ ˆ๊ฐ ํ™•์ธ
  4. ๋ณ‘๋ชฉ kernel๋งŒ Triton์œผ๋กœ ์ด๋™
  5. APOLLO-style scaling ์ถ”๊ฐ€
  6. CPU basis cache ์ถ”๊ฐ€
  7. 3B/7B ๋ชจ๋ธ๋กœ ํ™•์žฅ
  8. ํ•„์š”ํ•  ๋•Œ๋งŒ CUDA extension ์ž‘์„ฑ

9. ์˜ˆ์ƒ ๊ธฐ์—ฌ์ 

๋ณธ ๋ฐฉ๋ฒ•์˜ ์˜ˆ์ƒ ๊ธฐ์—ฌ์ ์€ ๋‹ค์Œ๊ณผ ๊ฐ™๋‹ค.

์ฒซ์งธ, 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 ์ €์žฅ์†Œ๋กœ ์ œํ•œํ•œ๋‹ค.


10. ํ•œ๊ณ„์™€ ์œ„ํ—˜ ์š”์†Œ

10.1 ์ˆ˜์น˜ ์•ˆ์ •์„ฑ

activation compression๊ณผ low-rank gradient projection์€ ๋ชจ๋‘ ํ•™์Šต ํ’ˆ์งˆ์— ์˜ํ–ฅ์„ ์ค„ ์ˆ˜ ์žˆ๋‹ค. ํŠนํžˆ LayerNorm, residual stream, lm_head๋Š” ์••์ถ•์— ๋ฏผ๊ฐํ•  ์ˆ˜ ์žˆ์œผ๋ฏ€๋กœ BF16 ์œ ์ง€๊ฐ€ ์•ˆ์ „ํ•˜๋‹ค.

10.2 Projection overhead

basis ๊ฐฑ์‹ ์— SVD๋ฅผ ์‚ฌ์šฉํ•˜๋ฉด ๋งค์šฐ ๋А๋ฆด ์ˆ˜ ์žˆ๋‹ค. ๋”ฐ๋ผ์„œ ์‹ค์ œ ๊ตฌํ˜„์—์„œ๋Š” randomized SVD, power iteration, update interval์ด ํ•„์ˆ˜๋‹ค.

10.3 Reconstruction cost

low-rank update๋ฅผ ๋‹ค์‹œ full weight shape์œผ๋กœ ๋ณต์›ํ•˜๋Š” ๊ณผ์ •์ด ๋ณ‘๋ชฉ์ด ๋  ์ˆ˜ ์žˆ๋‹ค. ์ด ๋ถ€๋ถ„์€ Triton fused kernel๋กœ ์ตœ์ ํ™”ํ•ด์•ผ ํ•œ๋‹ค.

10.4 ๊ตฌํ˜„ ๋‚œ์ด๋„

Hugging Face Trainer๋ฅผ ๊ทธ๋Œ€๋กœ ์‚ฌ์šฉํ•˜๊ธฐ ์–ด๋ ต๋‹ค. ์ง์ ‘ training loop, optimizer step, gradient hook, activation hook์„ ์ œ์–ดํ•ด์•ผ ํ•œ๋‹ค.


11. ๊ฒฐ๋ก 

๋ณธ ๋…ผ๋ฌธ์€ ์ž‘์€ VRAM ํ™˜๊ฒฝ์—์„œ ๊ฑฐ๋Œ€ ์–ธ์–ด๋ชจ๋ธ์„ ํ•™์Šตํ•˜๊ธฐ ์œ„ํ•œ ๊ณ„์ธตํ˜• ๋ฉ”๋ชจ๋ฆฌ ํ•™์Šต ์•Œ๊ณ ๋ฆฌ์ฆ˜ HMT๋ฅผ ์ œ์•ˆํ•˜์˜€๋‹ค. ๊ธฐ์กด ๋ฐฉ์‹์€ ์ฃผ๋กœ ๋ชจ๋ธ ์ƒํƒœ๋ฅผ CPU๋‚˜ NVMe๋กœ ์ด๋™ํ•˜๊ฑฐ๋‚˜, ํ•™์Šต ๊ฐ€๋Šฅํ•œ ํŒŒ๋ผ๋ฏธํ„ฐ ์ˆ˜๋ฅผ ์ค„์ด๋Š” ๋ฐฉ์‹์— ์˜์กดํ•œ๋‹ค. ๋ฐ˜๋ฉด HMT๋Š” gradient, optimizer state, activation์˜ ํ‘œํ˜„ ์ž์ฒด๋ฅผ ์ค„์—ฌ GPU VRAM ์š”๊ตฌ๋Ÿ‰์„ ๋‚ฎ์ถ˜๋‹ค.

ํ•ต์‹ฌ ์•„์ด๋””์–ด๋Š” ๋‹ค์Œ๊ณผ ๊ฐ™๋‹ค.

  1. gradient๋Š” full matrix๋กœ ์˜ค๋ž˜ ๋ณด๊ด€ํ•˜์ง€ ์•Š๊ณ  low-rank space๋กœ projectํ•œ๋‹ค.
  2. optimizer state๋Š” full parameter shape์ด ์•„๋‹ˆ๋ผ low-rank shape์œผ๋กœ ์œ ์ง€ํ•œ๋‹ค.
  3. activation์€ layer type๋ณ„๋กœ keep, recompute, compress๋ฅผ ๋‹ค๋ฅด๊ฒŒ ์ ์šฉํ•œ๋‹ค.
  4. CPU RAM์€ weight offload๊ฐ€ ์•„๋‹ˆ๋ผ projection basis cache๋กœ ์‚ฌ์šฉํ•œ๋‹ค.
  5. NVMe๋Š” checkpoint-only ์ €์žฅ์†Œ๋กœ ์ œํ•œํ•œ๋‹ค.
  6. ์„ฑ๋Šฅ ๋ณ‘๋ชฉ์€ Triton kernel๋กœ ๋‹จ๊ณ„์ ์œผ๋กœ ์ตœ์ ํ™”ํ•œ๋‹ค.

์ง์ ‘ ๊ฐœ๋ฐœ ๊ด€์ ์—์„œ๋Š” Python + PyTorch๋กœ prototype์„ ๋งŒ๋“ค๊ณ , ๋ณ‘๋ชฉ์ด ํ™•์ธ๋œ compression/projection/update ์—ฐ์‚ฐ๋งŒ Triton์œผ๋กœ ์˜ฎ๊ธฐ๋Š” ๋ฐฉ์‹์ด ๊ฐ€์žฅ ํ˜„์‹ค์ ์ด๋‹ค. CUDA C++๋Š” ์ตœ์ข… ๋‹จ๊ณ„์—์„œ๋งŒ ํ•„์š”ํ•˜๋‹ค.


12. ์ตœ์ข… ๊ฐœ๋ฐœ ๋กœ๋“œ๋งต

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

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages