Architectures & Paradigms

How a model is built, and the algorithms it learns by.

7 papers

Written by Junkun Yuan.

Click here to go back to main contents.


Papers are displayed in reverse chronological order. High-impact or inspiring works are highlighted in red.

Generative Modeling Paradigms

Improved Mean Flows: On the Challenges of Fastforward Generative Models

Zhengyang Geng (2,871), Yiyang Lu (154), Zongze Wu (2,933), Eli Shechtman (74,268), J. Zico Kolter (74,167), Kaiming He (871,053)

CMU · MIT · Adobe · THU

Conference on Computer Vision and Pattern Recognition (CVPR), 2026

Dec 01, 2025      iMF (90)      code (349)

imageflow


Recast MeanFlow's self-written target as plain velocity regression, and turn its frozen guidance scale into one more test-time condition.

Figure 1. Where the extra input enters (Fig. 1). Both panels compute one loss on the instantaneous velocity v from a network uθ that predicts the average velocity u across [r,t] — the two fields the MeanFlow identity ties together. In (a) the gray compound function reads two arrows, the noisy latent z and its conditional velocity e−x, the second of which no legitimate regressor of z may use; in (b) the JVP's tangent is the network's own v, so only z goes in and the box becomes a function of the noisy sample alone, which is the paper's first contribution stated in one picture.
  • What MeanFlow regresses (§3, Eq. 4–7). MeanFlow identity u=v−(t−r)du/dt (Eq. 4), with du/dt=v∂zu+∂tu≜JVP(u;v) (Eq. 5), makes u trainable: target utgt=(e−x)−(t−r)JVP(uθ;e−x) (Eq. 6) is fitted by 𝔼‖uθ−sg(utgt)‖2 (Eq. 7) — e−x standing for v in both places.

  • The identity read backwards (§4.1, Eq. 8–10). Solve for v(zt)=u(zt)+(t−r)du/dt (Eq. 8), make it as Vθ≜uθ+(t−r)JVPsg(uθ;e−x) (Eq. 9) and minimize 𝔼‖Vθ−(e−x)‖2 (Eq. 10) — fully equivalent to Eq. 6–7, and now a v-loss re-parameterized by u-pred.

  • The input that should not be there (Eq. 11, Fig. 1a). Written out, that compound function is Vθ(zt,e−x), reading the conditional velocity as a second argument — illegitimate for a regressor of zt, and traceable to Eq. 6, which put the conditional where Eq. 5 asks for the marginal.

  • Predict the tangent instead (Eq. 12). Vθ(zt)≜uθ(zt)+(t−r)JVPsg(uθ;vθ) (Eq. 12) — both arguments are now network outputs of zt alone. Stop-gradient survives, but inside the prediction rather than on the target, and is kept because dropping it brings second-order gradients in θ.

  • Two cheap vθ (§4.1, Tab. 1a). The boundary condition vθ(zt,t)≡uθ(zt,t,t) is free — FID 32.69 → 29.42 unguided, 3.43 → 2.99 on XL/2. An auxiliary v-head over the last 8 layers, on its own flow-matching loss, free at inference, does worse unguided, 30.76, and better guided, 5.68.

  • Why the conditional tangent hurts (§4.1, Eq. 2, Fig. 3). Feeding e−x to the loss leaks nothing, since the true regression target is the marginal v(zt)≜𝔼[e−x∣zt]; but as the JVP's tangent its variance is magnified by the Jacobian it multiplies, which is what the two loss curves separate.

Figure 2. The loss that goes the wrong way (Fig. 3). Across all 50K steps the original MeanFlow objective drifts upward and swings hard, while iMF's falls once, flattens near its floor and stays there for good.
  • Guidance, fixed and unfixed (§4.2, Eq. 13–15). MeanFlow bakes vcfg(zt∣c)=ωv(zt∣c)+(1−ω)v(zt) (Eq. 13) in at one ω chosen before training. iMF makes it an input, Vθ(·∣c,ω)≜uθ(zt∣c,ω)+(t−r)JVPsg (Eq. 15), with ω~p(ω)∝ω−β drawn on [1,8], skewed toward the small ω.

  • The target it fits (App. A, Eq. 17). vtgt=(e−x)+(1−1/ω)(uθ(zt∣t,t,c)−uθ(zt∣t,t,∅)) (Eq. 17) is MeanFlow's guided target with ω an input, so ω=1 means no guidance. FID 5.68 → 5.52, small only because MeanFlow's fixed ω already suited the B/2 model.

  • The interval is a condition too (§4.2, Tab. 1b). Pass Ω={ω,tmin,tmax}, sampling tmin~U[0,0.5] and tmax~U[0.5,1] and switching guidance off (ω=1) when t falls outside — 5.52 → 4.57, while the ω=1 column itself improves 30.76 → 20.95, so training across scales generalizes better.

  • In-context conditioning (§4.3, Eq. 16, Tab. 1c). uθ=uθ(zt∣r,t,c,Ω) (Eq. 16) is too heterogeneous for one adaLN-zero to sum; replicate each into tokens — 8 for the class, 4 for the rest — concatenated with the image patches: 4.57 → 4.09, and 133M → 89M once adaLN is gone.

  • Experiments (§5). ImageNet-256 in the VAE latent, from scratch, 1-NFE FID-50K on MeanFlow's code and B/2 baseline: the changes step 6.17 → 5.68 → 4.57 → 4.09, SwiGLU, RMSNorm and RoPE 3.82, 640 epochs 3.39 at 89M; iMF-XL/2 takes 1.72 at 1-NFE and 1.54 at 2-NFE.

  • The size labels do not calibrate (§5.2, Tab. 2). Without adaLN-zero, parameters go to depth: iMF-XL/2 has 48 layers, 610M and 174.6 Gflops to MF-XL/2's 28, 676M and 119.0, so 1.72 against 3.43 is not at matched FLOPs; iMF-M/2's 2.27 at 49.9 Gflops over MF's 5.01 at 54.0 is.

The training and sampling algorithm and its code. Follow Algorithm 2 of paper.

def train():
  # 1. Sample timesteps r and t from lognorm(-0.4, 1.0), sorted so r <= t; 50% of rows get r = t
  r, t = sample_time_steps(bs)

  # 2. Sample the guidance scale and interval; rows with r = t get the full interval [0, 1]
  omega = sample_cfg_scale(bs)  # omega ~ omega^-beta on [1, 8]; assume beta=1
  t_min, t_max = sample_cfg_interval(bs, r == t)  # t_min ~ U[0, 0.5], t_max ~ U[0.5, 1]

  # 3. Build the noisy latent z_t and its velocity v_t
  z_t = (1 - t) * z_0 + t * noise  # linear path
  v_t = noise - z_0  # conditional velocity

  # 4. Predict the average velocity u and the instantaneous velocity v with a class dropout
  dropped = rand(bs) < 0.1  # 10% class dropout
  y_in = where(dropped, null, y)
  u, v = model(z_t, r, t, omega, t_min, t_max, y=y_in)  # Eq. (16); omega and the interval are inputs

  # 5. Predict cond and uncond instantaneous velocity at t with the v head; omega = 1 outside the interval
  omega_t = where((t >= t_min) & (t <= t_max), omega, 1)
  with no_grad():
      v_c, v_u = model(cat([z_t, z_t]), cat([t, t]), cat([t, t]), cat([omega_t, ones]), cat([zeros, zeros]),
                       cat([ones, ones]), y=cat([y, null]), predict_u=False).chunk(2)

  # 6. Build the guided velocity target
  v_target = v_t + (1 - 1 / omega_t) * (v_c - v_u)  # Eq. (17)
  v_target = where(dropped, v_t, v_target)

  # 7. Build the compound prediction V, with the predicted v_c as the JVP tangent
  fn = lambda z, r_, t_: model(z, r_, t_, omega, t_min, t_max, y=y_in, predict_v=False)
  _, dudt = jvp(fn, (z_t, r, t), (v_c, zeros_like(r), ones_like(t)))  # no_grad, so dudt is sg(du/dt)
  V = u + (t - r) * dudt  # Eq. (15)

  # 8. Calculate loss of the u head and the auxiliary v head
  loss = adaptive_loss(V - v_target) + adaptive_loss(v - v_target)  # sum of squared error / (sg(.) + 0.01)^1
  loss = loss.mean()


def sample():
  time_steps = linspace(1, 0, num_steps + 1)
  for i in range(num_steps):
      t = full((bs,), time_steps[i])
      r = full((bs,), time_steps[i + 1])
      u, _ = model(z, r, t, omega, t_min, t_max, y=y, predict_v=False)  # guidance is a condition: one forward
      z = z - (t - r) * u

Mean Flows for One-step Generative Modeling

Zhengyang Geng (2,871), Mingyang Deng (1,603), Xingjian Bai (777), J. Zico Kolter (74,167), Kaiming He (871,053)

CMU · MIT

Advances in Neural Information Processing Systems (NeurIPS), 2025

May 19, 2025      MeanFlow (538)      code (609)

imageflow


Model the average velocity over an interval instead of the instantaneous one — a one-step generator trained from scratch, with guidance in the field.

  • Average velocity (§4.1, Eq. 3). u(zt,r,t)≜(1/(t−r))∫rtv(zτ,τ)dτ — displacement across [r,t], a field induced by v itself and by no network; limr→tu=v, and integral additivity already gives (t−r)u(zt,r,t)=(s−r)u(zs,r,s)+(t−s)u(zt,s,t), consistency with nothing imposed.

  • The MeanFlow identity (Eq. 6). Multiply the definition by t−r and differentiate in t with r fixed; the product rule and the fundamental theorem give u+(t−r)du/dt=v, that is u(zt,r,t)=v(zt,t)−(t−r)du/dt — the intractable integral traded for a local relation between the fields.

  • One JVP computes it (Eq. 8, Tab. 1b). du/dt=(dzt/dt)∂zu+(dr/dt)∂ru+(dt/dt)∂tu=v∂zu+∂tu, since dzt/dt=v, dr/dt=0: a Jacobian–vector product with tangent (v,0,1), one extra backward pass. Corrupt it and one-step collapses — 61.06 FID against 137.96–329.22.

  • The objective (Eq. 9–11). ℒ(θ)=𝔼‖uθ(zt,r,t)−sg(utgt)‖22 with utgt=vt−(t−r)(vt∂zuθ+∂tuθ): the network's own partials stand in for those of the true u, the conditional vt for the marginal, and stop-gradient removes double backpropagation. Zero loss recovers the identity, hence Eq. 3.

  • Sampling (§4.1, Eq. 12). zr=zt−(t−r)uθ(zt,r,t) — the integral is the prediction, so z0=z1−uθ(z1,0,1) is one call, few steps chaining it.

  • Guidance inside the field (§4.2, Eq. 19). Define vcfg≜ωv(zt,t∣c)+(1−ω)v(zt,t) and average it; since vcfg(zt,t)=v(zt,t)=ucfg(zt,t,t), the unconditional half is the network's own diagonal, v~t=ωvt+(1−ω)uθ(zt,t,t), so guidance stays a single call (ω=3 gives 15.53).

  • Guidance, improved (App. B.1, Eq. 21, Tab. 5). A mixing scale κ lets the class-conditional diagonal into the target as well: v~t≜ω(ϵ−x)+κuθ(zt,t,t∣c)+(1−ω−κ)uθ(zt,t,t); solving it gives Eq. 13's form vcfg(zt,t∣c)=ω′v(zt,t∣c)+(1−ω′)v(zt,t) at the effective scale ω′=ω/(1−κ). At a fixed ω′=2, κ from 0 to 0.9 takes the 1-NFE FID 20.15 → 18.63; the released XL/2 ships ω=1, κ=0.5.

  • Adaptive loss weighting (§4.3, App. B.2, Tab. 1e). For the error Δ=uθ−utgt, the powered loss ‖Δ‖22γ equals the squared L2 loss under the weight w=1/(‖Δ‖22+c)p, p=1−γ, c=10−3, applied as sg(w)·‖Δ‖22: p=1 gives 61.06, Pseudo-Huber-like p=0.5 63.98, plain L2 79.75.

  • Two clocks (§4.3, Tab. 1a/c/d). (r,t) come from lognorm(−0.4,1.0), and t>r, and only 25% of pairs keep r≠t — at 0% the correction vanishes, the loss is exactly flow matching and 1-NFE fails (328.91 against 61.06). The network is fed (t,t−r), best of four embeddings.

Figure 1. Instantaneous against average. v is the path's tangent; u(z,r,t) aligns instead with the displacement across [r,t] — a field conditioned on both endpoints, shown here for three settings of t, and it is exactly the object that a MeanFlow network regresses.
  • Experiments (§5). ImageNet-256 in a VAE latent, FID-50K at 1-NFE, plus CIFAR-10: B/2 to XL/2 reads 6.17 / 5.01 / 3.84 / 3.43 at 240 epochs (Shortcut 10.60), 2.20 at 2-NFE trained longer, DiT-XL/2's 2.27 at 250×2; CIFAR-10 2.92 without EDM preconditioning, against iCT's 2.83.

The training and sampling algorithm and its code. Follow Algorithm 1 and 2 of paper.

def train():
  # 1. Sample timesteps r and t from lognorm(-0.4, 1.0), sorted so r <= t
  r, t = sample_time_steps(bs)

  # 2. Build the noisy latent z_t and its velocity v_t
  z_t = (1 - t) * z_0 + t * noise  # linear path
  v_t = noise - z_0  # conditional velocity

  # 3. Predict the average velocity u with a class dropout
  dropped = rand(bs) < 0.1  # 10% class dropout
  y = where(dropped, null, y)
  u = model(z_t, r, t, y)

  # 4. Build CFG-guided velocity
  with no_grad():
      u_c, u_u = model(cat([z_t, z_t]), cat([t, t]), cat([t, t]), y=cat([y, null])).chunk(2)
  v_tilde = 0.2 * v_t + 0.92 * u_c + (1 - 0.2 - 0.92) * u_u  # Eq. 21; assume omega=0.2, kappa=0.92
  v_hat = where((t <= 0.8) & ~dropped, v_tilde, v_t)

  # 5. Calculate average velocity target
  _, dudt = jvp(lambda z, r_, t_: model(z, r_, t_, y), (z_t, r, t), (v_hat, zeros_like(r), ones_like(t)))  # Eq. 8
  u_target = v_hat - (t - r) * dudt  # Eq. 10

  # 6. Calculate loss
  loss = (u - u_target.detach()).square().flatten(1).sum(1)  # Eq. 9
  loss = (loss / (loss.detach() + 0.01)).mean()  # p=1, c=0.01 as the official code; paper: c=1e-3


def sample():
  time_steps = linspace(1, 0, num_steps + 1)
  for i in range(num_steps):
      t = full((bs,), time_steps[i])
      r = full((bs,), time_steps[i + 1])
      z = z - (t - r) * model(z, r, t, y=y)

One Step Diffusion via Shortcut Models

Kevin Frans (1,979), Danijar Hafner (17,880), Sergey Levine (277,035), Pieter Abbeel (284,265)

UC Berkeley

International Conference on Learning Representations (ICLR), 2025

Oct 16, 2024      Shortcut Models (379)      code (768)

imageflow


One network, conditioned on step size, jumps any distance along the flow — trained in one run, with no teacher, schedule, or second stage.

Figure 1. Two objectives, one run. At d≈0, regress empirical velocities like any flow-matching model (top); for larger d, chain two half-size shortcuts into the training target (bottom) — generation capability propagates stepwise from many steps down toward a single one.
Figure 2. Training (Alg. 1). First k: flow-matching targets x1−x0 at d=0; the rest: two chained small steps, then stopgradded — no teacher anywhere.
Figure 3. Five panels of failure (Fig. 2). Left: overlapping noise–data pairings force flow matching onto the mean direction 𝔼[vt∣xt]. Right: fewer steps drift further toward the dataset mean; at one step (red circles) multi-modality is gone — the objective's own optimum fails here, not undertraining.
Figure 4. Sampling (Alg. 2). M Euler steps of size 1/M each.
  • Why naive one-step fails (§2, Fig. 2). Flow matching learns the average direction over all pairings passing through xt; at t=0 that average points at the dataset mean, so one big Euler step lands there — multimodal data is unreachable in a single step even at the optimum.

  • Condition on the step size (§3, Eq. 3). The shortcut is the normalized jump to the correct next point, xt+d′=xt+s(xt,t,d)d; at d→0 it is the flow itself, so shortcut models generalize flow matching — one model serves every budget from 128 steps to one.

  • Self-consistency instead of a teacher (§3, Eq. 4–5). One step equals two chained halves, s(xt,t,2d)=s(xt,t,d)/2+s(xt+d′,t+d,d)/2: bootstrap targets come from the model itself, while a batch fraction at d=0 regresses flow matching. One run, ~16% extra compute.

  • A binary grid of step sizes (§3.1). 128 base units give eight shortcut lengths d∈(1/128,…,1), each trained against two of the next size down, keeping bootstrap paths short; many-step quality matches the flow baseline, few- and one-step matches two-stage distillation methods.

  • Experiments (§5). CelebA-HQ-256 and ImageNet-256, every objective retrained on one DiT-B codebase at matched compute, FID-50K at 128 / 4 / 1 steps: 6.9 / 13.8 / 20.5 and 15.5 / 28.3 / 40.3 — past every end-to-end and two-stage method but progressive distillation's one step.

  • Bootstrapping still scales, and drives robots (§5.3–5.5). Unlike Q-learning, one-step FID keeps falling with size (DiT-XL: 10.6, 3.8 at 128 steps), and as a one-step policy on Push-T and Transport it scores 0.87 / 0.80 where a one-step diffusion policy collapses to 0.12 / 0.00.

The method straight from the repo, cut to the lines that carry the idea:

# One batch, two target kinds (targets_shortcut.py); 1/bootstrap_every of it bootstraps.
dt_base = ladder_index  # the network is conditioned on the log2 exponent, not on dt itself
dt = 1 / 2**dt_base  # the binary ladder [1, 1/2, ..., 1/128]
t = randint(0, 2**dt_base) / 2**dt_base  # t lands only on multiples of dt
x_t = (1 - (1 - 1e-5) * t) * x_0 + t * x_1  # a 1e-5 noise floor, not an exact lerp

# Bootstrap targets: two half-steps of the model itself (EMA and CFG both optional).
v1 = model(x_t, t, dt_base + 1)  # +1 = one rung down the ladder, step dt/2
x_mid = clip(x_t + (dt/2) * v1, -4, 4)  # clipped midpoint - stability the paper never mentions
v2 = model(x_mid, t + dt/2, dt_base + 1)
v_target = clip((v1 + v2) / 2, -4, 4)  # no grad flows through targets

# Flow-matching targets for the rest of the batch, on the finest grid.
t = randint(0, 128) / 128
v_target = x_1 - (1 - 1e-5) * x_0  # empirical velocity; dt_base = log2(128) encodes d = 0

# One loss serves both target kinds (train.py).
loss = mean((model(x_t, t, dt_base) - v_target) ** 2)

Simplifying, Stabilizing and Scaling Continuous-Time Consistency Models

Cheng Lu (15,386), Yang Song (73,221)

OpenAI

International Conference on Learning Representations (ICLR), 2025

Oct 14, 2024      sCM (276)

It made continuous-time consistency training work at scale — TrigFlow, tangent normalization and a JVP recipe that MeanFlow, rCM and the newer few-step families build on.

imagediffusiondistillationunread


Simplify consistency models, stabilize the continuous-time limit by taming its tangent, scale to 1.5B: two steps within 10% of the teacher.

  • The continuous-time limit (§2.2, Eq. 2). A consistency model fθ(xt,t) maps any trajectory point to its start, fθ(x,0)≡x. Discrete-time CMs match adjacent times, 𝔼[w(t)d(fθ(xt,t),fθ−(xt−Δt,t−Δt))], xt−Δt from an ODE solver: solver error plus a hand-tuned Δt schedule. With ℓ2, as Δt→0 the gradient is ∇θ𝔼[w(t)fθ⊤dfθ−/dt] with the tangent dfθ−/dt=∇xtfθ−dxt/dt+∂tfθ−: solver-free, and until now unstable.
Figure 1. The limit, drawn. Top and middle: a discrete-time CM's second point comes from an ODE solver step and lands just off the trajectory by 𝒪(Δt); bottom: as Δt→0 the pair collapses to one point and its tangent dxt/dt, which a continuous-time CM follows directly — no solver, no error.
Figure 2. Alg. 1. The whole recipe in one loop; each line has a bullet below — TrigFlow noising, t=arctan(eτ/σd), the rearranged tangent with warmup r, its normalization by ‖g‖+c, the adaptive weight wϕ. Training and distillation differ in one line: where dxt/dt comes from.
  • TrigFlow (§3, Eq. 3–5). With noise z~𝒩(0,σd2I) at the data's scale, xt=cos(t)x0+sin(t)z, t∈[0,π/2]: EDM's coefficients become cskip=cos(t), cout=−σdsin(t), cin≡1/σd; the PF-ODE is dxt/dt=σdFθ(xt/σd,cnoise(t)), trained on ‖σdFθ−vt‖22 with vt=cos(t)z−sin(t)x0; the CM is a DDIM step, fθ(xt,t)=cos(t)xt−sin(t)σdFθ(xt/σd,cnoise(t)). EDM enters by t=arctan(σ/σd); flow matching is a special case.

  • Where the tangent blows up (§4.1, Eq. 6–7). dfθ−/dt=−cos(t)(σdFθ−−dxt/dt)−sin(t)(xt+σddFθ−/dt): F, the ODE and xt are tame and ∇xtF·dxt/dt well-conditioned, so the culprit is sin(t)∂tFθ−=sin(t)∂tcnoise·∂emb/∂cnoise·∂Fθ−/∂emb — three factors, one fix each. Fig. 4: EDM with Fourier scale 16 sends the norm past 150, positional EDM still explodes at t→π/2, TrigFlow stays under 4, on CIFAR-10 models.

  • Three architectural fixes (§4.1). (1) cnoise(t)=t: EDM's log(σdtant) gives sin(t)∂tcnoise=1/cos(t)→∞. (2) Positional time embeddings: the derivative of emb(c)=sin(s·2πωc+ϕ) scales with s; positional means s≈0.02 — iCT's scale, explained. (3) Adaptive double normalization, y=norm(x)⊙pnorm(s(t))+pnorm(b(t)): raw AdaGN made CMs diverge; pixel-normalizing scale and shift keeps its expressivity.

  • Tangent normalization (§4.2, Fig. 5a). Most gradient variance rides on dfθ−/dt, so replace it by g/(‖g‖+c), c=0.1, or clip it to [−1,1]: on ImageNet 512 either pulls 1-step FID well below the raw tangent's, and the two land together — applied before the weighting below.

  • Adaptive weighting (§4.2, Eq. 8, App. D). The identity ∇θ𝔼[Fθ⊤y]=12∇θ𝔼‖Fθ−Fθ−+y‖22 turns Eq. 2 into an MSE, so EDM2's learned weight applies: ℒsCM=𝔼[ewϕ(t)/D‖Fθ−Fθ−−cos(t)dfθ−/dt‖22−wϕ(t)], the prior weight 1/(σdtant) turning sin(t) into cos(t); at the optimum every time's weighted loss is 1. Times come from t=arctan(eτ/σd), τ~𝒩(Pmean,Pstd2) with Pmean≤−0.8: near clean data, where the real signal enters.

  • Diffusion init and tangent warmup (§4.2, App. G). Init is the teacher's EMA; the sin(t) term is scaled by r=min(1,iters/10k), optional, against gradient spikes. Distillation costs ~2× a diffusion step and reaches the teacher's 2-step quality within ~20k iterations, under 20% of its compute.

  • Discrete never catches up (§4.2, App. E, Fig. 5c). Every fix ports to discrete time through Δθ−(xt,t,t′)=(fθ−(xt,t)−fθ−(xt′,t′))/sin(t−t′), xt′ one DDIM step away, which tends to dfθ−/dt as t′→t (Eq. 37) and replaces the JVP in Eq. 8 unchanged; on N EDM-spaced levels (ρ=7), FID improves with N up to 1024, worsens past it from numerical precision, and on ImageNet 512 the continuous model sits below every N.

  • JVP rearrangement (§5.1). The tangent dFθ−/dt=∇xtFθ−·dxt/dt+∂tFθ− overflows in FP16 near t=0 and π/2. The loss needs only cos(t)df/dt∝cos(t)sin(t)dF/dt, so the JVP is fed the tangent (cos(t)sin(t)dxt/dt, cos(t)sin(t)σd) at (xt/σd,t) — the scale folded in upstream.

  • Flash-attention JVP (§5.1, App. F). For y=softmax(x)V=pV the output tangent is ty=ptV+(p⊙tx)V−(ptx⊤)y; the two new terms need only a running vector g=(ex−m⊙tx)V and scalar μ=∑iexi−mtx,i, merged block by block like Flash Attention's own m,ℓ,f — g rescaled by ema−m, μ summed the same way — and read out as tpV=g/ℓ−(μ/ℓ)y, so attention and its JVP come out of one pass with no attention matrix stored.

  • Sampling and guidance (§5.2, App. G). One step from tmax=arctan(80/σd), two steps re-noising to t=1.1. On ImageNet 512 CFG is distilled in: the student takes the guidance scale s~𝒰[1,2] as an input embedded like t, the teacher at that scale; sCT, teacher-free, gets none.

Figure 3. Continuous beats every grid (Fig. 5c). 1-step FID of EDM2-distilled CMs on ImageNet 512 across training, all techniques applied to both sides: discrete time improves as N grows to 1024, degrades at 2048 and 4096 from numerical precision, and never reaches the continuous-time curve, which leads from the first checkpoint on.
Figure 4. Recall against guidance (Fig. 7b). EDM2-M on ImageNet 512: VSD's recall falls with the guidance scale like over-guided diffusion, alone or with sCD; 2-step sCD stays with the teacher — the diversity claim in a curve.
  • Experiments (§5.2, Tab. 1–2). CIFAR-10, ImageNet 64 and 512, FID at 1 / 2 steps for sCT and sCD: two steps already read 2.06 / 1.48 / 1.88 against teachers 2.01 / 1.33 / 1.73; 1-step sCD-XXL 2.28 beats StyleGAN-XL's 2.41 and VAR's 2.63 at under 10% of their compute.

  • sCD scales like its teacher; sCT does not in latent space (Fig. 6, Tab. 2). The sCD/teacher FID ratio holds at ~1.3 (one step) and ~1.09 (two) from S to XXL, so the gap shrinks; sCT wins at 64px but reads 10.13 against sCD's 3.07 at 512-S, variance blamed on the VAE latent.

  • VSD buys precision with recall; sCD keeps both (Fig. 7, Tab. 7). From guidance 1 to 2 VSD's recall collapses as under heavy guidance, as 2-step sCD tracks the teacher; VSD on top of sCD moves 1-step FID only 2.75 → 2.67 but FD-DINOv2 83.78 → 54.81: the metrics disagree.

Improved Techniques for Training Consistency Models

Yang Song (73,221), Prafulla Dhariwal (211,189)

OpenAI

International Conference on Learning Representations (ICLR), 2024

Oct 22, 2023      Improved Consistency Models (476)

imagediffusion


Consistency training grown up — drop the teacher-side EMA, swap LPIPS for Pseudo-Huber, grow the discretization, and beat distillation.

  • Why leave CD and LPIPS (§1). Distillation caps a student at its teacher; LPIPS leaks ImageNet features into FID, whose Inception judge saw the same data. iCT drops both and wins anyway: one-step FID 2.51 on CIFAR-10 and 3.25 on ImageNet-64, past CD at one and two steps.

  • No EMA on the teacher network (§3.2). Of the two arguments tying CT to consistency matching, the gradient one holds only when θ−=θ; with an EMA teacher the N→∞ limit is not even a valid objective — a Dirac toy example breaks it. So μ=0, on the teacher side only.

  • Pseudo-Huber replaces LPIPS (§3.3). d=‖x−y‖22+c2−c with c=0.00054d sits between ℓ2 and ℓ1: it beats ℓ2 outright and catches LPIPS once N(k) grows large — no feature network left to backpropagate through, and the heuristic scales c with data dimensionality.

  • Schedules and small knobs (§3.1, §3.4–3.5). N(k) doubles from 10 up to 1280 across training; noise levels draw from a discretized lognormal (Pmean=−1.1, Pstd=2); weighting λ=1/(σi+1−σi); plus a smaller Fourier scale (0.02) and genuinely large dropout (0.3 on CIFAR-10).

Figure 1. The whole paper in one table. Every design change from the original consistency training, side by side — the four headline fixes the bullets unpack, plus the hyperparameters that carry them — the originals on the left, their replacements on the right.
  • Experiments (§4, Tab. 2–3). CIFAR-10 and ImageNet-64 by FID, IS, precision and recall, GANs with pretrained discriminators excluded: iCT reads 2.83 / 2.46 at 1 / 2 steps on CIFAR-10 and 4.02 / 3.20 on ImageNet, iCT-deep 2.51 / 2.24 and 3.25 / 2.77 — past CD and TRACT.

  • Recall, and thin margins (§4). ImageNet-64 recall rises to 0.63 from CT's 0.47 and BigGAN-deep's 0.48; one-step iCT edges StyleGAN-ADA on CIFAR-10 (2.83 vs 2.92) and BigGAN-deep on ImageNet (4.02 vs 4.06); two-step iCT-deep's 2.24 ties Score SDE's 2.20 at 2000 steps.

Latent Consistency Models: Synthesizing High-Resolution Images with Few-Step Inference

Simian Luo (1,888), Yiqin Tan (1,251), Longbo Huang (8,544), Jian Li (13,034), Hang Zhao (41,987)

Tsinghua University

arXiv, 2023

Oct 06, 2023      LCM (931)      code (4.6K)

imagediffusiondistillation


Consistency distillation moves into Stable Diffusion's latent space, swallows classifier-free guidance in one stage, and lands 768px in 2–4 steps.

  • Latent consistency distillation (§4.1). The consistency function moves to SD's latent space, fθ(zt,c,t)↦z0, parameterized on the teacher's own noise prediction and initialized from it; the ODE solver Ψ — DDIM or DPM-Solver — appears only in training, never at inference.

  • Guidance folded in, one stage (§4.2). The augmented ODE runs on the guided noise (1+ω)ϵc−ωϵ∅ — its solver estimate is two solver calls combined — with the student conditioned on ω~U[2,14]. Guided-Distill took two stages and 45 A100-days; LCM, 32 A100-hours.

  • Skipping-step (§4.3). Adjacent times on SD's 1000-step schedule sit so close that the consistency loss nearly vanishes and convergence crawls; aligning tn+k with tn at k=20 cuts the effective schedule from a thousand points to dozens; convergence lands in 2,000 iterations.

  • Latent consistency fine-tuning (§4.4). A trained LCM adapts to custom image datasets without giving up its few-step inference.

  • Experiments (§5). LAION-Aesthetics-6+ (12M, 512px) and 6.5+ (650K, 768px) from SD2.1, FID and CLIP on 30K images at ω=8: 768px LCM reads 34.22 / 16.32 / 13.53 FID at 1 / 2 / 4 steps against Guided-Distill's 120.28 / 30.70 / 16.70 and DPM++'s 188.91 / 67.14 / 20.08.

  • Guided-Distill was reproduced small (§5). Its numbers are a re-implementation at batch 72, 100K iterations, not the paper's batch 512 on 32 A100s — an admitted handicap: the claim is faster convergence at equal compute, no ceiling; k=50 suits DPM/DPM++, breaks DDIM.

Consistency Models

Yang Song (73,221), Prafulla Dhariwal (211,189), Mark Chen (134,179), Ilya Sutskever (847,818)

OpenAI

International Conference on Machine Learning (ICML), 2023

Mar 02, 2023      Consistency Models (2,160)      code (6.5K)

It opened the consistency route to one-step generation — self-consistency as the training signal, teacher optional — seeding a large family of few-step models.

imagediffusiondistillation


One evaluation maps noise straight to data — learn the map from any ODE point to its origin, distilled from a teacher or trained from scratch.

Figure 1. The object being learned. Every point of a probability-flow ODE trajectory maps to its origin; self-consistency — one trajectory, one output — is the training signal, and the name.
  • The consistency function (§3, Fig. 1). f(xt,t)↦xϵ sends every point of the probability-flow ODE to its origin, so sampling is a single network evaluation at xT; self-consistency — every point of one trajectory shares an output — is the property both training modes below enforce.

  • The boundary condition (§3, Eq. 5). f(·,ϵ) must be the identity, or f≡0 solves everything; it comes almost for free from the skip form cskip(t)x+cout(t)Fθ(x,t) with cskip(ϵ)=1, cout(ϵ)=0 — the EDM-style shape that lets existing diffusion architectures drop in unchanged.

  • Consistency distillation (§4, Alg. 2). Noise an image to xtn+1, take one teacher step back, and pull the two outputs together — the target side from an EMA copy, stopgradded. With LPIPS and N=18: CIFAR-10 3.55, ImageNet-64 6.20 one-step, past PD everywhere but one point.

  • Consistency training (§5, Alg. 3). Swap the teacher's score for the unbiased estimate −(xt−x)/t2: the pair becomes x+tnz and x+tn+1z with one shared z, no diffusion model anywhere — a standalone family that still matches PD's one step, with N and μ grown on schedules.

  • Multistep sampling (Alg. 1). Alternate the model with re-noising to a lower level, time points picked greedily by ternary search; extra compute buys quality back, and mapping any xt to xϵ gives zero-shot inpainting, colorization, super-resolution and stroke-guided editing.

Figure 2. Sampling as alternation (Alg. 1). Line one is already a sample — one evaluation at x^T. Each pass re-noises to a lower τn and maps again: N evaluations buying N−1 corrections.
Figure 3. The two trainers, side by side (Alg. 2, 3). Left: Distillation buys its trajectory pairs with one teacher ODE step. Right: Training builds them from the same data and noise sample alone.
Figure 4. The two trainers, redrawn in Alg. 2–3's notation. Gray runs identically in both rows — noising to xtn+1, student fθ, EMA target fθ−, the weighted metric; the yellow box is the entire fork: the pair's second point is one teacher ODE step Φ (top) or the same x,z re-noised at tn (bottom).
  • Experiments (§6). CIFAR-10, ImageNet-64, LSUN Bedroom / Cat by FID, IS, precision, recall, from EDM teachers: CD beats PD at 1–2 steps everywhere but one point and the synthetic-data distillers; CT, teacher-free, gets 8.70 / 5.83 on CIFAR-10, past every VAE and flow.

  • The metric is half the margin (§6.2). PD retrained with LPIPS improves uniformly on its ℓ2 (Fig. 4), so part of CD's lead is the metric, not the objective; Heun and N=18 come from EDM unchanged, and CT samples share structure with EDM's from one noise — no mode collapse.

Both trainers in one routine, straight from the repo, cut to the lines that carry the idea:

# The parameterization (karras_diffusion.py, denoise): f(x, s_min) = x by construction.
c_skip = sd**2 / ((sigma - s_min)**2 + sd**2)  # EDM's scalings with sigma -> sigma - s_min
c_out = (sigma - s_min) * sd / (sigma**2 + sd**2)**0.5  # at s_min: c_skip = 1, c_out = 0
f = c_skip * x_t + c_out * unet(c_in * x_t, 250 * log(sigma))  # nothing left to learn there

# One routine serves both modes (consistency_losses): adjacent Karras levels t > t2.
x_t = x + t * z  # noise a real image
pred = f(x_t, t, theta)  # the trainable side; dropout RNG state saved here
with no_grad():  # one solver step t -> t2 builds the trajectory pair
    denoiser = teacher_denoise(x_t, t) if distill else x  # CT: the raw image is the score
    d = (x_t - denoiser) / t  # CD upgrades this to Heun with a second evaluation at t2
    x_t2 = x_t + (t2 - t) * d
    target = f(x_t2, t2, theta_ema).detach()  # EMA copy, rerun under the student's dropout mask
loss = lpips(pred, target)  # or plain l2; LPIPS inputs are resized up to 224px first

# Sampling is f(s_max * z, s_max); Alg. 1 re-noises lower and applies f again for quality.

Last updated on September 17, 2026  ·  Citations from Semantic Scholar & Google Scholar; stars from GitHub, September 2026