a12s12's picture
add level-drop (25% uniform-over-levels)
adc76ac verified
|
Raw
History Blame Contribute Delete
4.33 kB
# Phase-2 ์ „ํ™˜ ๋…ธํŠธ (DiT unfreeze) โ€” 2026-07-16 (์ตœ์ข…๋ณธ, B ๋ฐฉ์‹)
## ๋ฌด์—‡์ด ๋ฐ”๋€Œ์—ˆ๋‚˜
Phase-1: **DiT trunk(453M) ์™„์ „ freeze**, encoder(ViT)+attn-pool+cond๋งŒ ํ•™์Šต.
Phase-2: **DiT unfreeze**, encoder์™€ **๋™์ผ lr**๋กœ ํ•จ๊ป˜ ํ•™์Šต (dit_lr_scale ์—†์Œ).
## LR ์ฒ˜๋ฆฌ (์ค‘์š” โ€” B ๋ฐฉ์‹์œผ๋กœ ํ™•์ •)
Phase-1์€ step380000์—์„œ **100-epoch warmup์˜ 38% ์ง€์ (lrโ‰ˆ9.5e-6)**์ด์—ˆ์Œ.
Phase-2๋Š” ์ด ์ง€์ ๋ถ€ํ„ฐ **์›๋ž˜ warmup ์Šค์ผ€์ค„์„ ๊ทธ๋Œ€๋กœ ์ด์–ด๊ฐ**:
- ์‹œ์ž‘ lr โ‰ˆ **9e-6** (step380000 ๊ฐ’๊ณผ ๋™์ผ), warmup์ด ๋๋‚˜๋ฉฐ(step 1,000,800) **2.5e-5๋กœ ์ ์  ๋„๋‹ฌ**.
- โŒ fresh 100ep warmup(0๋ถ€ํ„ฐ, full๊นŒ์ง€ 48h) ์•„๋‹˜ โ€” ๊ทธ๊ฑด ๋‚ญ๋น„๋ผ ํ๊ธฐ.
- โŒ 3ep ๋‹จ์ถ• warmup(0โ†’2.5e-5 ๊ธ‰์ƒ์Šน) ์•„๋‹˜ โ€” DiT init์— ๊ธ‰๊ฒฉ.
- โœ… step380000 LR ์ง€์ ์—์„œ ์ž์—ฐ์Šค๋Ÿฝ๊ฒŒ ์ด์–ด๋ฐ›์•„ ์›๋ž˜ ์Šค์ผ€์ค„๋Œ€๋กœ ์ƒ์Šน.
## ์–ด๋–ป๊ฒŒ (๊ตฌํ˜„)
1. **๊ฐ€์ค‘์น˜**: `init_from`์ด accelerate ckpt ๋””๋ ‰ํ† ๋ฆฌ๋„ ๋ฐ›๋„๋ก ํ™•์žฅ โ†’ step380000์˜
model.safetensors(encoder+pool+DiT+cond) ๋กœ๋“œ. (loss ์—ฐ์†์„ฑ์œผ๋กœ ๊ฒ€์ฆ: ์‹œ์ž‘ diff/repa๊ฐ€
Phase-1 ์ข…๋ฃŒ๊ฐ’๊ณผ ๋™์ผ)
2. **LR ์Šค์ผ€์ค„ ์œ„์น˜**: trainer์— `PHASE2_RESUME_STEP` env ์ถ”๊ฐ€ โ†’ `self.steps=380000`์œผ๋กœ
์„ค์ •. loaded_steps=-1์ด๋ผ **๋ฐ์ดํ„ฐ fast-forward ์—†์ด**(38 epoch ์žฌ์ฝ๊ธฐ ํšŒํ”ผ) ์Šค์ผ€์ค„๋Ÿฌ๊ฐ€
`step_update(380000)`์œผ๋กœ lr=9e-6 ์œ„์น˜. warmup_epochs=100 ์œ ์ง€.
3. **optimizer**: fresh (frozenโ†’unfrozen์ด๋ฉด param group 134Mโ†’599M ๋ณ€๊ฒฝ โ†’ accelerate ์ „์ฒด
load_state ๋ถˆ๊ฐ€). momentum๋งŒ ์ƒˆ๋กœ, LR ์œ„์น˜๋Š” ์ •ํ™•.
4. **ckpt ๋ฒˆํ˜ธ**: step380000๋ถ€ํ„ฐ ์ด์–ด๊ฐ€ **step390000, 400000...** ์—ฐ์† ์ €์žฅ (์žฌ์‹œ์ž‘ ์•„๋‹˜).
5. trainable 134M(P1) โ†’ **599M(P2, DiT ํฌํ•จ)**. FID eval์€ step400000(50k ๋ฐฐ์ˆ˜)์—์„œ ์ฒซ ์‹œ์ž‘.
## ํŒŒ์ผ (์ „๋ถ€ ์ด repo code/ ์— ์—…๋กœ๋“œ๋จ)
- `spatial_diffuse_slot.py` โ€” ๋ชจ๋ธ(SpatialAttnPool + DiTSpatial + init_from ๋””๋ ‰ํ† ๋ฆฌ ์ง€์›)
- `spatial_mask.py` โ€” spatial-align mask ๋นŒ๋”
- `diffusion_trainer_PHASE2_PATCH.txt` โ€” trainer์˜ PHASE2_RESUME_STEP ํŒจ์น˜ (์ ์šฉ ์œ„์น˜ ๋ช…์‹œ)
- `trainer_utils_create_optimizer.py` โ€” dit_lr_scale ์ง€์› create_optimizer (B์—์„  ๋ฏธ์‚ฌ์šฉ, ์ฐธ๊ณ )
- `tokenizer_l_spatial.yaml` (Phase-1) / `tokenizer_l_spatial_phase2.yaml` (Phase-2)
- `train_spatial_l.sh` (P1) / `train_spatial_l_phase2.sh` (P2, PHASE2_RESUME_STEP env)
- `hf_ckpt_watcher.py` (P1) / `hf_ckpt_watcher_phase2.py` (P2) / `make_loss_plot*.py`
## HF ๋„ค์ž„์ŠคํŽ˜์ด์Šค (์ด repo)
- Phase-1(frozen, ๋ณด์กด): `step360000~380000/`, `samples_all_steps/`, `loss_curves.png`
โ†’ **step380000/์€ Phase-2 init ์†Œ์Šค๋ผ ๋ฐ˜๋“œ์‹œ ์œ ์ง€** (resume์— ํ•„์š”)
- Phase-2(unfrozen): `phase2_step390000+/`, `samples_all_steps_phase2/`, `loss_curves_phase2.png`, `logs_phase2/`
- watcher๊ฐ€ ์ตœ์‹  2๊ฐœ ์œ ์ง€ + squash(์šฉ๋Ÿ‰ ํšŒ์ˆ˜). repo๋Š” **public**(private ์šฉ๋Ÿ‰ํ•œ๋„ ํšŒํ”ผ).
## resume (์ƒˆ ์„œ๋ฒ„)
1. `code/`์˜ ํŒŒ์ผ๋“ค + `semanticist` repo์— ๋ฐฐ์น˜, `diffusion_trainer.py`์— PHASE2 ํŒจ์น˜ ์ ์šฉ.
2. Phase-2 ์ด์–ด๊ฐ€๊ธฐ: ์ตœ์‹  `phase2_stepN/`(HF)์„ ๋กœ์ปฌ์— ๋‘๊ณ  `configs/tokenizer_l_spatial_phase2.yaml`์˜
`init_from`์„ ๊ทธ ๊ฒฝ๋กœ๋กœ, `PHASE2_RESUME_STEP=N`์œผ๋กœ `train_spatial_l_phase2.sh` ์‹คํ–‰.
3. Phase-2 ์ฒ˜์Œ๋ถ€ํ„ฐ: `init_from=step380000` + `PHASE2_RESUME_STEP=380000` (ํ˜„์žฌ ์„ค์ •).
## Level drop ์ถ”๊ฐ€ (2026-07-16, step390000๋ถ€ํ„ฐ)
์šฐ๋ฆฌ multi-res ๋ฐฉ์‹์˜ ํ•ต์‹ฌ์ธ **level drop(nested)** ๋ฅผ Phase-2์— ์ถ”๊ฐ€ (step390000์—์„œ resume).
- **๋ฐฉ์‹**: whole-level, coarse-first. keep_levels ~ uniform(1,4) โ†’ {1x1},{1x1,2x2},{1x1,2x2,4x4},{all} ๊ฐ **25%**.
โ†’ ๋ ˆ๋ฒจ๋ณ„ ์œ ์ง€ํ™•๋ฅ  1x1=100%, 2x2=75%, 4x4=50%, 8x8=25%.
- **์™œ 25%(๋ ˆ๋ฒจ๊ท ๋“ฑ)**: ํ† ํฐ๋น„์œจ(uniform-over-token) ๋ฐฉ์‹์€ 8x8์ด 64/85๋ผ coarse-only๊ฐ€ ๊ฑฐ์˜ ํ•™์Šต ์•ˆ ๋จ.
๋ ˆ๋ฒจ๊ท ๋“ฑ์€ ๋ชจ๋“  granularity๋ฅผ ๊ณ ๋ฅด๊ฒŒ ํ•™์Šต โ†’ ๊ฐ ๋ ˆ๋ฒจ์ด ๋…๋ฆฝ์ ์œผ๋กœ ์“ธ๋ชจ์žˆ์Œ (multi-res reasoning ๋ชฉ์ ).
- **๋“œ๋กญ ์ˆœ์„œ**: 8x8(finest, DiT ์ด๋ฏธ์ง€ํ† ํฐ๊ณผ 1:1)๋ถ€ํ„ฐ ๋“œ๋กญ, 1x1(global)์€ ํ•ญ์ƒ ์œ ์ง€.
์ „๋ถ€ ๋“œ๋กญ(uncond)์€ CFG(uncond_drop_prob)๊ฐ€ ๋”ฐ๋กœ ๋‹ด๋‹น.
- **๊ตฌํ˜„**: `LevelNestedSampler` (spatial_diffuse_slot.py), config `enable_nest: True`.
- inference: `inference_with_n_slots`(ํ† ํฐ budget)๋ฅผ whole-level prefix๋กœ ๋งคํ•‘ (85โ†’all, 21โ†’8x8 drop, 5โ†’[2x2,1x1], 1โ†’[1x1]).