| # 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. |
|
|