SANA LRM Architecture
System Overview
This implementation reuses the modular trainer architecture from flux:
scripts/train.pyorchestrates load/prepare/train/eval/save.accelerators/*manage distributed runtime, logging, and checkpoints.datasets/*build pairwise preference batches with timestep controls.models/*provide latent-space text-image feature extraction.criterions/*compute pairwise preference loss.tasks/*compute inference probabilities and accuracy metrics.
Data Flow
- Dataset loads pairwise examples from
pickapic-anonymous/pickapic_v1. - Each sample yields prompt tokens, image pair tensors, labels, and timestep pair.
- Criterion concatenates image pairs and passes text/image/timestep to the model.
- Model returns normalized text/image embeddings in shared projection space.
- Criterion computes pairwise logits and weighted preference loss.
- Task evaluation converts scores to probabilities and computes accuracy.
SANA Model Path
- Text branch:
- Load SANA tokenizer/text encoder(s) by checkpoint profile.
- Build prompt conditioning embeddings for transformer cross-attention.
- Build pooled text representation for reward embedding projection.
- Image branch:
- Preprocess RGB image to configured resolution.
- Encode image with SANA VAE to latents.
- Apply scheduler-based noise at selected timesteps.
- Run SANA transformer on noisy latents with text conditioning.
- Pool transformer latent outputs to image representation.
- Reward projection:
- Project text and image representations to common dimension.
- L2-normalize features.
- Compute scaled similarity using learnable
logit_scale.
Training Controls
- Mixed precision and gradient accumulation from accelerator config.
- Distributed feature gather for pairwise loss consistency.
- Validation/test evaluation at configured intervals.
- Periodic checkpointing and best-metric tracking.
Logging and Outputs
- Output root:
logs/lrm/{project_name}/{run_name}. - Saved artifacts:
- config snapshot
- train log
- periodic checkpoints
- final model state
Variant-Aware Profile Design
A model profile selects:
- pretrained checkpoint id
- default image size (512 or 1024)
- tokenizer/text path behavior
- optional special overrides for memory or precision
This keeps one codebase while enabling all four requested SANA checkpoints.