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.