Index
2026-04-16 — Plan

EgoX2_exp1 Dual Expert 구현 계획

EgoX2_exp1 | Wan 2.2 A14B transformer + transformer_2 LoRA 동시 학습

TL;DR

~82 GB
예상 Peak VRAM
0.9
boundary_ratio
32ch
proj_out (RGB+PT)
128
LoRA rank

1 배경 / 목적

Wan 2.2 I2V A14B 체크포인트는 transformer/ (high-noise expert, denoising 초반 90%)와 transformer_2/ (low-noise refiner, 마지막 10%) 두 개의 WanTransformer3DModel을 갖고 있다. model_index.json"boundary_ratio": 0.9가 정의되어 있다.

현재 문제: sft_trainer.py:load_components()에서 transformer subfolder만 로드하고 transformer_2는 완전히 무시. 학습/추론 모두 high-noise expert 하나로만 동작 중 → low-noise refinement 단계 부재로 디테일 품질 저하 가능.

exp2에서 동일 이슈에 대한 구현이 완료되었으나, exp1은 코드 구조가 달라 별도 계획이 필요하다:

특성exp1 (이 계획)exp2 (기완료)
Transformer 클래스WanTransformer3DModel_GGAWanTransformer3DModel
Output32ch (RGB 16 + PtMap 16), expand_proj_out(32)16ch (RGB only)
Loss 함수_compute_loss_original + _compute_loss_with_ptmapcompute_loss HCPT 분기
Inference pipelineWanDROIDPipelineWanWidthConcatImageToVideoPipeline
LoRA rank128256

2 작업 내용 (수정 계획)

수정 파일 및 핵심 변경

1. core/finetune/schemas/args.py

boundary_ratio: float = 0.9 필드 추가 (Point Cloud Map 섹션 근처). --boundary_ratio argparse 추가.

2. core/finetune/schemas/components.py

Wan_Componentstransformer_2: Any = None 필드 추가. 기존 vars(self) 순회 로직이 device 관리를 자동 처리.

3. core/finetune/models/wan_i2v/sft_trainer.py (5곳)

  • (3A) load_components(): transformer_2 subfolder 존재 시 WanTransformer3DModel.from_pretrained 로드. use_pointmap=Truetransformer_2.expand_proj_out(32)도 적용.
  • (3B) _compute_loss_original(): timestep 기반 expert 선택. boundary_idx = int(num_train_timesteps * boundary_ratio). batch_size=1이므로 매 step 한 expert만 forward.
  • (3C) _compute_loss_with_ptmap(): 동일 로직, active_transformer(...)로 교체.
  • (3D) initialize_pipeline(): WanWidthConcatImageToVideoPipelinetransformer_2 + boundary_ratio 전달.
  • (3E) Pipeline __init__ + __call__: boundary_step = int(len(timesteps) * boundary_ratio), 루프 내 active_transformer 선택.

4. core/finetune/trainer.py (6곳)

  • (4A) prepare_trainable_parameters(): transformer_2에 동일 LoraConfig + add_adapter. gradient_checkpointing 활성화. use_pointmap이면 proj_out full fine-tune. ignore_list"transformer_2" 추가.
  • (4B) prepare_optimizer(): cast_training_params에 transformer_2 포함, trainable params 합산.
  • (4C) prepare_for_training(): accelerator.prepare()에 transformer_2 포함.
  • (4D) train() loop: transformer_2.train() 조건부 호출, models_to_accumulate에 추가, grad clipping에 포함.
  • (4E) __prepare_saving_loading_hooks(): transformer_2 LoRA를 pytorch_lora_weights_transformer_2.safetensors로 별도 저장.
  • (4F) _maybe_run_validation(): transformer_2.eval() / .train() 조건부 전환.

5. configs/droid_train.yaml

boundary_ratio: 0.9 추가 (1줄).

메모리 예상

14B params × 2 expert × bf16 = 56 GB (weights) LoRA rank 128 on both: ~2 GB (adapters) Optimizer (LoRA params AdamW m/v): ~4 GB Peak activations + grad: ~20 GB ───────────────────────────────────────────── 총 예상: ~82 GB B200 178 GB 대비: 충분 (여유 ~96 GB)
exp2 smoke test 참고: exp2 (LoRA rank 256, HCPT 5-stream)에서 peak 81.4 GB 실측. exp1은 rank 128이고 ptmap 32ch 구조지만 attention 토큰 수 차이로 비슷한 범위 예상.

3 결과 (현재 상태)

계획 수립 완료. 구현은 context.md 기준 "구현 완료, 학습 테스트 대기" 상태로 기재되어 있으나, 실제 smoke test는 미실행.

항목상태
args.py boundary_ratio 추가완료
components.py transformer_2 추가완료
sft_trainer.py 5곳 수정완료
trainer.py 6곳 수정완료
droid_train.yaml 업데이트완료
Smoke test (5-clip, both experts forward)대기
본 학습 (8 GPU DDP)대기
미확인 사항: find_unused_parameters=True로 변경 시 DDP 오버헤드 예상. transformer_2 None 체크가 모든 경로에 일관적으로 적용되었는지 실행 시 검증 필요.

4 Takeaway

exp1과 exp2의 dual expert 구현 차별점

exp2의 dual expert 구현(Steps 1-6)은 이미 완료되어 smoke test까지 통과했다. exp1은 WanTransformer3DModel_GGAptmap 32ch 구조로 인해 expand_proj_out(32)를 transformer_2에도 적용해야 하고, _compute_loss_with_ptmap_compute_loss_original 두 loss 경로 모두에 expert switching을 넣어야 한다는 점이 다르다.

코드 변경 규모는 exp2와 유사 (~35줄 sft_trainer + ~70줄 trainer). 하위 호환 보장: transformer_2 subfolder 없으면 components.transformer_2 = None으로 fallback.

5 Next Steps

Smoke test 실행

5-clip subset으로 login node GPU 1개 실행. 두 expert 모두 forward/backward 확인. 기존 checkpoint-5000에서 resume 시 transformer_2 LoRA 없음 → 처음부터 학습 필요할 수 있음.

CUDA_VISIBLE_DEVICES=<free_gpu> timeout 5400 python finetune.py \ --config configs/droid_train.yaml \ --train_steps=10 --validation_steps=5

본 학습 제출 (smoke 통과 후)

sbm "bash scripts/finetune_droid.sh" --gres=gpu:8 -c 192 --mem 1600GB --qos=core-extra