aryadomain's picture
Add files using upload-large-folder tool
ef8f3ad verified
|
Raw
History Blame Contribute Delete
3.45 kB

SANA LRM Implementation Plan

Goal

Build a working latent-space reward model (LRM) under lrm_sana that follows the proven training path used in flux and aligns with the LRM objective from https://arxiv.org/abs/2502.01051.

Scope

  • Keep code inside lrm_sana.
  • Keep all markdown docs inside lrm_sana/docs.
  • Reuse the available dataset path/protocol from flux.
  • Preserve flux-like run facilities: launcher profiles, logs structure, checkpoint behavior.
  • Support these SANA checkpoints via model profiles:
    • Efficient-Large-Model/Sana_600M_512px_diffusers
    • Efficient-Large-Model/Sana_1600M_512px_diffusers
    • Efficient-Large-Model/Sana_Sprint_0.6B_1024px_diffusers
    • Efficient-Large-Model/Sana_Sprint_1.6B_1024px_diffusers

Constraints

  • Use pickapic-anonymous/pickapic_v1 dataset flow from flux.
  • Keep pseudo-preference CSV filtering support (vqa_aes_clip_score_mp.csv).
  • Follow distributed DeepSpeed launch conventions from flux/train_flux.sh.
  • Add variant-aware SANA text path so single/dual encoder variants can be handled safely.

Phases

Phase 1: Docs-First Baseline

  • Create planning and architecture docs in lrm_sana/docs.
  • Define implementation checkpoints and verification criteria before code rewrites.

Phase 2: Project Scaffolding

  • Mirror flux trainer structure into lrm_sana/trainer.
  • Add top-level files: lrm_sana/setup.py, lrm_sana/train_lrm_sana.sh.
  • Keep import paths local to lrm_sana package.

Phase 3: Config and Registration Wiring

  • Add SANA config groups and names:
    • task: step_sana
    • dataset: step_sana
    • model: step_sana_base
    • criterion: step_clip_sana
  • Add step_sana_base.yaml with flux-like accelerator/dataset defaults.

Phase 4: SANA Preference Model

  • Implement trainer/models/sana_preference_model.py:
    • SANA tokenizer/text encoder loading.
    • SANA VAE latent encode + scaling.
    • Scheduler timestep/noise mixing.
    • SANA transformer forward.
    • Text/image projection heads and logit_scale.
  • Keep save/load behavior compatible with current trainer checkpoint flow.

Phase 5: Dataset/Task/Criterion Integration

  • Implement step_sana_hf_dataset.py from step_flux_hf_dataset.py baseline.
  • Implement step_sana_task.py from step_flux_task.py baseline.
  • Implement step_clip_criterion_sana.py from flux criterion baseline.
  • Keep pairwise preference loss, distributed gather, and eval accuracy logic.

Phase 6: Runtime Script and Profiles

  • Adapt train_lrm_sana.sh from train_flux.sh:
    • RUN_PROFILE=main|quick
    • offline cache env defaults
    • pseudo-preference CSV fallback
    • DeepSpeed sharded launch
  • Add model profile selection variable for the four SANA checkpoints.

Phase 7: Verification

  • Static checks:
    • module imports
    • hydra config composition
    • dataclass/config registration validity
  • Runtime checks:
    • quick smoke run (max_steps=1)
    • output/log/checkpoint tree validation

Deliverables

  • Working SANA LRM code under lrm_sana.
  • Documentation set in lrm_sana/docs.
  • Flux-like launcher for H200 node testing.

Success Criteria

  • Training script launches and completes a quick run end-to-end.
  • Dataloader, model forward, criterion loss, and eval metric path work without interface mismatch.
  • Logs/checkpoints/config snapshots are written to expected locations.
  • Model profile switch works across the four requested SANA checkpoints.