all_code_base / lrm /lrm_sana /docs /architecture.md
aryadomain's picture
Add files using upload-large-folder tool
ef8f3ad verified
|
Raw
History Blame Contribute Delete
2.39 kB

SANA LRM Architecture

System Overview

This implementation reuses the modular trainer architecture from flux:

  • scripts/train.py orchestrates 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

  1. Dataset loads pairwise examples from pickapic-anonymous/pickapic_v1.
  2. Each sample yields prompt tokens, image pair tensors, labels, and timestep pair.
  3. Criterion concatenates image pairs and passes text/image/timestep to the model.
  4. Model returns normalized text/image embeddings in shared projection space.
  5. Criterion computes pairwise logits and weighted preference loss.
  6. Task evaluation converts scores to probabilities and computes accuracy.

SANA Model Path

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