How sCM Distillation Enables One-Step Generation in SANA-Sprint
SANA-Sprint achieves one-step image generation by applying continuous-time consistency model (sCM) distillation, which trains the diffusion model to satisfy the ODE at any timestep, enabling a single deterministic jump from noise to image without iterative refinement.
SANA-Sprint is a high-speed variant of the Sana diffusion model developed by NVlabs that leverages advanced distillation techniques to reduce inference steps from hundreds to one. By implementing sCM (continuous-time consistency model) distillation, the model learns a time-consistent mapping that directly estimates clean latents from noisy inputs in approximately 0.1 seconds per 1024px image on an H100 GPU. This article examines the specific implementation details in the NVlabs/Sana repository that enable this one-step capability.
Training with Time-Continuous Consistency
The foundation of one-step generation is established during training in train_scripts/train_scm_ladd.py. The training pipeline optimizes a composite loss function that combines the classic L₂ reconstruction term with an sCM consistency loss and LADD (Large-scale Adversarial Diffusion Distillation) loss.
The sCM loss forces the network to satisfy the continuous consistency equation described in the sCM research paper. This trains the model to be time-agnostic, meaning it can predict the denoised latent x₀ from any noisy latent x_t at any continuous time t ∈ [0, T]. Once the model learns this continuous mapping, inference no longer requires the hundreds of small discrete steps characteristic of classic DDPMs.
The SCMScheduler and Trigonometric Flow
The SCMScheduler class in diffusion/scheduler/scm_scheduler.py (line 47) implements the trigonometric flow parameterization that enables the single-step jump. Unlike standard schedulers that require many timesteps, this scheduler configures a minimal two-point schedule when num_inference_steps is set to 2.
Configuring the Two-Step Schedule
The set_timesteps method constructs a schedule consisting of the maximum timestep and zero, with an optional intermediate point. For true one-step generation, the configuration uses:
from diffusion.scheduler.scm_scheduler import SCMScheduler
scheduler = SCMScheduler(prediction_type="trigflow")
scheduler.set_timesteps(
num_inference_steps=2, # Results in one execution of the sampling loop
max_timesteps=1.57080, # ≈ π/2, default from sCM paper
intermediate_timesteps=1.0, # Optional intermediate point
)
Internally, lines 108-110 of scm_scheduler.py build the timestep tensor [max_timesteps, intermediate_timesteps, 0]. When num_inference_steps equals 2, the scheduler prepares a trajectory where the model predicts the denoised output directly.
The Single-Step Update Mechanism
During the forward pass, the scheduler computes the denoised estimate pred_x0 using the trig-flow formula:
pred_x0 = cos(s) * sample - sin(s) * model_output
When only two timesteps are present, the scheduler omits the noise term and sets prev_sample = pred_x0. This yields a single deterministic jump from the initial noisy latent directly to the final denoised latent, bypassing iterative refinement entirely.
One-Step Inference Implementation
The inference pipeline in scripts/inference_sana_sprint.py activates one-step mode by selecting the "scm" sampler. The script defines a step dictionary that automatically configures the minimal step count:
# scripts/inference_sana_sprint.py – sampler selection
sample_steps_dict = {"scm": 2}
sample_steps = args.step if args.step != -1 else sample_steps_dict[args.sampling_algo]
The Sampling Loop Execution
The sampling loop (lines 75-88) executes only once when configured for sCM inference. The model prediction is scaled by sigma_data (a learned parameter from the sCM training) before being passed to the scheduler:
# Single-step sampling loop from inference_sana_sprint.py
for i, t in enumerate(timesteps[:-1]): # Executes once when len(timesteps)==2
model_pred = sigma_data * model(
latents / sigma_data,
timestep,
caption_embs,
**model_kwargs
)
latents, denoised = scheduler.step(
model_pred, i, t, latents, return_dict=False
)
The resulting denoised tensor is immediately passed to the VAE via vae_decode(), producing the final image without additional denoising iterations.
Why One-Step Generation Works
The effectiveness of SANA-Sprint's one-step generation stems from three technical components:
- Consistency Distillation: Training enforces that the model satisfies the diffusion ODE at any continuous time, removing the need for many small discrete steps.
- TrigFlow Parameterization: The sinusoidal time scaling provides a closed-form solution for the ODE, allowing an analytic update (
pred_x0) that is exact for the learned continuous dynamics. - Deterministic Schedule: The two-point schedule (
max_timesteps→ 0) directly evaluates this solution in a single transformation.
Practical Code Examples
Initializing the Scheduler for One-Step Inference
from diffusion.scheduler.scm_scheduler import SCMScheduler
# Initialize with trigflow prediction type used during sCM training
scheduler = SCMScheduler(prediction_type="trigflow")
# Configure for single-step generation
scheduler.set_timesteps(
num_inference_steps=2,
max_timesteps=1.57080,
intermediate_timesteps=1.0,
)
Running One-Step Generation
import torch
from diffusion.model.builder import build_model, get_vae, vae_decode
# Load model and VAE
model = build_model(...).to(device).eval()
vae = get_vae(...).to(device)
# Initialize latent with sigma_data scaling
latents = torch.randn(1, model.latent_dim, h, w, device=device) * scheduler.sigma_data
# Encode text prompts
caption_embs, emb_masks = ... # From tokenizer/text encoder
# Single model evaluation at max timestep
timestep = scheduler.timesteps[0].expand(1).to(device)
model_pred = scheduler.sigma_data * model(
latents / scheduler.sigma_data,
timestep,
caption_embs,
data_info={"cfg_scale": torch.tensor([cfg_scale])},
mask=emb_masks,
)
# Execute one-step update
latents, denoised = scheduler.step(
model_pred, timeindex=0, timestep=timestep, sample=latents, return_dict=False
)
# Decode to image
image = vae_decode(vae, denoised / scheduler.sigma_data)
Command-Line Inference
python scripts/inference_sana_sprint.py \
--config configs/sana_sprint_config/1024ms/SanaSprint_1600M_1024px_allqknorm_bf16_scm_ladd.yaml \
--model_path hf://Efficient-Large-Model/Sana_Sprint_1.6B_1024px/checkpoints/Sana_Sprint_1.6B_1024px.pth \
--sampling_algo scm
Summary
- sCM distillation trains the model to predict denoised latents at any continuous time
t, making it time-agnostic. - The SCMScheduler uses
prediction_type="trigflow"and a two-timestep schedule to perform a single deterministic jump from noise to image. - Inference requires only one step (configured via
num_inference_steps=2), executing in approximately 0.1 seconds on an H100 GPU. - Key implementation files include
train_scripts/train_scm_ladd.pyfor training andscripts/inference_sana_sprint.pyfor inference.
Frequently Asked Questions
What is the difference between sCM distillation and standard diffusion training?
Standard diffusion models learn to predict noise or velocity at specific discrete timesteps, requiring many iterative steps during inference. sCM distillation trains the model to satisfy the consistency property across continuous time, meaning the model can directly predict the final denoised output from any intermediate noisy state. This eliminates the need for iterative refinement.
How many inference steps does SANA-Sprint actually require?
SANA-Sprint requires one inference step for generation, configured by setting num_inference_steps=2 in the SCMScheduler. The schedule creates two timesteps (max and zero), but the sampling loop executes only once, performing a single jump from the initial noise to the final denoised latent.
What role does the trigflow prediction type play in the scheduler?
The prediction_type="trigflow" parameter in SCMScheduler enables the trigonometric flow parameterization used by sCM models. This parameterization calculates the denoised estimate using the formula pred_x0 = cos(s) * sample - sin(s) * model_output, which provides a closed-form solution to the diffusion ODE under continuous time. This allows the scheduler to compute the final latent deterministically without iterative noise removal.
Where is the sCM consistency loss implemented in the training code?
The sCM consistency loss is implemented in train_scripts/train_scm_ladd.py. The training script combines this loss with reconstruction and LADD losses to optimize the model. Specifically, the sCM term enforces that the model's output at a coarse timestep matches the output of a finer-grained step, training the network to maintain consistency across the continuous time domain.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →