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_diffusersEfficient-Large-Model/Sana_1600M_512px_diffusersEfficient-Large-Model/Sana_Sprint_0.6B_1024px_diffusersEfficient-Large-Model/Sana_Sprint_1.6B_1024px_diffusers
Constraints
- Use
pickapic-anonymous/pickapic_v1dataset flow fromflux. - 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
fluxtrainer structure intolrm_sana/trainer. - Add top-level files:
lrm_sana/setup.py,lrm_sana/train_lrm_sana.sh. - Keep import paths local to
lrm_sanapackage.
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
- task:
- Add
step_sana_base.yamlwith 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.pyfromstep_flux_hf_dataset.pybaseline. - Implement
step_sana_task.pyfromstep_flux_task.pybaseline. - Implement
step_clip_criterion_sana.pyfrom flux criterion baseline. - Keep pairwise preference loss, distributed gather, and eval accuracy logic.
Phase 6: Runtime Script and Profiles
- Adapt
train_lrm_sana.shfromtrain_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
- quick smoke run (
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.