NEPA-DiT
Multi-embedding prediction.

Embedding Prediction Helps Image Generation

Preprint, 2026

Sihan Xu1*, Ji Xie2*, Zilin Wang1, Hui Shen1, Stella X. Yu1

1University of Michigan   2Carnegie Mellon University   *Equal contribution

Overview

In diffusion transformers, a class label or a text prompt is embedded once, and the same condition is reused at every denoising step. We ask whether predicted embeddings can serve as this condition instead.


  • We propose Embedding Conditioned Generation, which conditions a DiT generator on the clean embeddings predicted by a NEPA model from the condition and the noisy image, recomputed at every denoising step.
  • We train the NEPA model with Multi-Embedding Prediction, which extends NEPA to predict the embeddings of the whole clean image at once.
  • On class-conditional ImageNet 256Γ—256, we study the condition of the generator, the design of Multi-Embedding Prediction, and the scaling of both models; combined with REPA, NEPA-DiT-XL reaches an FID of 1.32.

Multi-embedding prediction.
Multi-embedding prediction. From the condition and the noisy image, the model predicts the next embeddings in the sequence, those of the clean image. Supervision is applied in embedding space; the predicted embeddings then condition a diffusion generator.

Method

Our method has two stages. First, we train a NEPA model with Multi-Embedding Prediction to predict clean image embeddings from condition tokens and noisy image embeddings. Second, we train a DiT generator with Embedding Conditioned Generation, where the frozen NEPA model produces predictive embeddings at each denoising step and the DiT uses them as its condition.

Multi-Embedding Prediction. Multi-Embedding Prediction (MEP) extends NEPA to predict the next K embeddings simultaneously from the same context.

𝐩n+1:n+K=hΞΈ(𝐳≀n),β„’MEP=1Kβˆ‘k=1Kπ’Ÿ(𝐳n+k,𝐩n+k).(3)

NEPA is the special case K=1. Like multi-token prediction, MEP predicts several items at once; it differs in predicting in embedding space rather than token space.

Multi-token vs. multi-embedding prediction.
Multi-token vs. multi-embedding prediction. Both predict several items at once. (a) Tokens are discrete, so hidden states are decoded into tokens by output heads. (b) Embeddings are continuous, so the outputs of the model are the predictions themselves. In both, the same context predicts all targets in one forward pass.

NEPA for Generation. In this sequence, the clean embeddings are the next N embeddings after the condition and the noisy image, and we train a NEPA model to predict them with MEP, taking K=N. The sequence is ordered by generation state rather than by patch position, so the embedding that follows noisy patch i is clean patch i. MEP predicts all N embeddings of the clean image in one forward pass, each from the full context of the condition and the noisy image.

𝐩=hΞΈ([c,⟨boi⟩],𝐳t),β„’MEP=π’Ÿ(𝐳0,𝐩).(4)
MEP for generation.
MEP for generation. The clean image embeddings are the next embeddings after the condition and the noisy image; the output at each noisy patch predicts the embedding that follows it.

Attention mask. The condition tokens and ⟨boi⟩ form a causal prefix that never attends to image tokens.

Embedding prediction loss. As illustrated in Figure 5, the positives lie on the diagonal, pairing each 𝐩i with 𝐳0,i, while the other patches of the same image serve as negatives:

π’Ÿ(𝐳0,𝐩)=βˆ’1Nβˆ‘i=1Nlogexp(𝐩i⊀𝐳0,i)βˆ‘j=1Nexp(𝐩i⊀𝐳0,j).(5)
Attention mask.
Attention mask. Rows are queries and columns are keys; dark cells are masked.
Embedding prediction loss.
Embedding prediction loss.

Embedding Conditioned Generation. At each denoising step, the current noisy input 𝐱t is first fed into the NEPA model, which outputs predictive embeddings 𝐩t=hθ([c,⟨boi⟩],f(𝐱t)).

The generator receives the condition only through 𝐩t and has no separate class embedding; the timestep enters through adaLN modulation, as in DiT.

β„’ECG=𝔼𝐱0,π›œ,t,c[w(t)‖𝐯tβˆ’GΟ•(𝐱t,t,𝐩t)β€–2].(6)

Condition caching. Because the condition prefix never attends to image tokens, its key-value cache in the NEPA model is computed once before the denoising loop and reused at every step; only the image tokens are re-encoded as 𝐱t changes.

Classifier-free guidance. During training, we randomly replace the condition c with a learned null token. At inference, the unconditional branch feeds the null token to the NEPA model, which yields 𝐩tβˆ…, and guidance is applied to the velocities predicted by the generator:

𝐯^cond=GΟ•(𝐱t,t,𝐩t),𝐯^uncond=GΟ•(𝐱t,t,𝐩tβˆ…),𝐯^cfg=𝐯^uncond+sβ‹…(𝐯^condβˆ’π―^uncond),(7)

Ablation Studies

Following previous work (Ma et al., 2024; Wang et al., 2026b), we run ablations at the B scale with NEPA-B and patch size 2, one factor at a time (defaults in gray); Table 6 then selects patch size 4.

Loss. InfoNCE gives the best FID, 31.36 against 35.23 for MSE and 37.82 for cosine similarity (Table 1). ImageNet-1K accuracy of the fine-tuned NEPA model follows the opposite order.

Loss. Acc. is ImageNet-1K top-1 accuracy of the fine-tuned NEPA model.
LossFID↓Acc. (%)↑
Cos. sim.37.8283.1
MSE35.2382.6
InfoNCE31.3681.7
Stop-grad.
DetachFID↓
w/n.c.
w/o31.36

Stop-gradient. Detaching the targets 𝐳0 from the gradient prevents training from converging (n.c. in Table 2), so gradients flow through the targets.

Target normalization. Normalizing the target embeddings raises FID from 31.36 to 34.76 (Table 3), so we use raw targets.

Timestep input. Feeding the diffusion timestep to the NEPA model changes FID from 31.36 to 31.40 (Table 4), presumably because the noise level is already evident from the noisy patches; we therefore keep the NEPA model timestep-free, while the generator still receives the timestep through adaLN.

Target norm.
TargetFID↓
Norm.34.76
Raw31.36
Timestep input.
TimestepFID↓
w/o31.36
w/31.40

Vocabulary. The InfoNCE loss scores each prediction against a bank of candidate targets. We compare banks built from the patches of the same image (instance), of all images in the per-device batch (batch), and of all images across devices (full). The instance, batch, and full banks give 31.36, 32.13, and 31.21 (Table 5). We use the instance bank, which requires no communication across devices; its negatives are the other patches of the same image, as in Equation 5.

Patch size. The patch size of the NEPA model sets the number of predicted embeddings (Table 6). Patch size 4 gives 31.00 against 31.36 for patch size 2 with a quarter of the tokens, which makes the query at every denoising step correspondingly cheaper, whereas patch size 8 is too coarse (33.20). We therefore use patch size 4 in all remaining experiments.

Vocabulary.
BankFID↓
Instance31.36
Batch32.13
Full31.21
Patch size.
PatchTokensFID↓
225631.36
46431.00
81633.20

Condition of the Generator

At the same small scale, Table 7 keeps the generator and its training recipe fixed and varies only what it is conditioned on. The class embedding gives 36.39, and adding the patch embeddings of 𝐱t to the class token gives 37.21. The last three rows give the generator the same extra network as ours, of the NEPA-XL architecture, so that their parameters and FLOPs per step match those of our model; what differs is how that network is trained. Trained end to end with the generator as a single network, it gives 29.52. Pretrained as a flow-matching model for 240 epochs, the same budget as MEP, and then frozen, its features give 30.86. Trained with MEP and frozen, its predicted embeddings give 25.04.

Condition of the generator. All rows use the same generator and training recipe (400K iterations, best over sampling steps) and differ only in its condition; the last three rows also carry the extra network of our model.
Condition of the generatorExtra networkTrainingObjectiveFID↓
Class embedding–––36.39
Class token + noisy patches–––37.21
Single network, end to endNEPA architecturejointflow matching29.52
Pretrained flow modelNEPA architecture240 epochs, frozenflow matching30.86
NEPA modelNEPA240 epochs, frozenMEP25.04

Refresh rate. Table 14 queries the NEPA model less often than every denoising step, reusing the last prediction in between. Querying it every 8 steps raises FID from 1.57 to 1.82, and querying it only once, at the first step, gives 242.63.

Refresh rate. NEPA-DiT-XL + REPA, 96-step ODE sampling with the guidance of Table 9. The generator runs at every step; between queries, it reuses the last predicted embeddings.
NEPA queryFID↓
Every step1.57
Every 8 steps1.82
Once, at the first step242.63

Scaling Behavior

Beyond the small-scale setting, Table 8 scales both models, training all nine combinations of NEPA-B/L/XL and DiT-B/L/XL, each for 400K iterations and evaluated with 64 ODE steps without guidance. FID decreases monotonically along both axes. Scaling the generator from DiT-B to DiT-XL improves FID by about 14 points for every NEPA model, and scaling the NEPA model from NEPA-B to NEPA-XL improves FID by about 5 points for every generator.

Scaling. FID-50K after 400K iterations; Acc. is ImageNet-1K top-1 accuracy (%) of the fine-tuned NEPA model.
Generator (FID↓)
NEPAAcc.↑DiT-BDiT-LDiT-XL
NEPA-B81.331.0017.6916.54
NEPA-L83.427.4614.3513.74
NEPA-XL85.225.0412.8611.27
Training progress of NEPA-DiT-B/L/XL.
Training progress of NEPA-DiT-B/L/XL, FID-50K without guidance.

Training progress. Figure 6 follows NEPA-DiT-B/L/XL over training; all three use the same frozen NEPA-XL model. Each point is FID-50K with 64 ODE steps and no guidance, evaluated every 100K iterations. Larger generators are better at every checkpoint, and the gap opens early: at 200K iterations, NEPA-DiT-L (17.55) and NEPA-DiT-XL (15.23) are already well ahead of NEPA-DiT-B at 400K (25.04).

Comparison with Previous Methods

Our final model combines Embedding Conditioned Generation, which sets the condition of the generator, with REPA (Yu et al., 2025), which aligns its features.

Among the methods of Table 9 in the SD-VAE latent space, NEPA-DiT-XL + REPA reaches an FID of 1.32 with the 250-step SDE sampler and 1.57 with the 96-step ODE sampler. Most of its training compute goes into the NEPA model, trained once: with it included, NEPA-DiT-XL, 1.39B parameters in all, is trained with 3.1Γ—1020 FLOPs, about a third of what SiT-XL/2 with REPA uses over its 800 epochs to reach 1.42.

Class-conditional generation on ImageNet 256Γ—256 with guidance. Top (gray): methods with modified or alternative tokenizers. Bottom: methods in the standard SD-VAE latent space. Metrics and training epochs on ImageNet-1K are taken from each paper; for ours, epochs are those of the NEPA model + the generator. β€œβ€“β€ means not reported. ‑AutoGuidance.
MethodTokenizerEpochsFID↓sFID↓IS↑Pre.↑Rec.↑
Other tokenizers
REPA + EQ-VAE (Kouzelis et al., 2025a)EQ-VAE2001.705.13283.00.790.62
LightningDiT-XL/1 (Yao et al., 2025)VA-VAE8001.354.15295.30.790.65
LightningDiT + IG (Zhou et al., 2026)VA-VAE6801.194.11269.00.790.66
DiT-XL + CMuon (Chen et al., 2026)VA-VAE2001.18––––
REPA-E (Leng et al., 2025)E2E-VAE8001.124.09302.90.790.66
Send-VAE + REPA (Page et al., 2026)Send-VAE8001.214.10315.10.790.66
SFD-XL‑ (Pan et al., 2026b)SemVAE8001.063.89267.00.780.67
RAE DiTDH-XL‑ (Zheng et al., 2026)RAE8001.13–262.60.780.67
MixFlow + RAE (Li et al., 2026)RAE800+2001.104.40259.70.780.67
RAEv2 (Singh et al., 2026)RAE801.06–255.3––
SD-VAE
DiT-XL/2 (Peebles & Xie, 2023)SD-VAE14002.274.60278.20.830.57
SiT-XL/2 (Ma et al., 2024)SD-VAE14002.064.49277.50.830.59
REPA (Yu et al., 2025)SD-VAE8001.424.70305.70.800.65
TREAD (Krause et al., 2025)SD-VAE7401.694.73292.70.810.63
DDT-XL/2 (Wang et al., 2026b)SD-VAE4001.26–310.60.790.65
REG (Wu et al., 2025)SD-VAE8001.364.25299.40.770.66
U-REPA (Tian et al., 2026)SD-VAE4001.41––––
SRA (Jiang et al., 2026)SD-VAE8001.584.65311.40.800.63
LSEP (Yun et al., 2025)SD-VAE8001.464.94296.80.800.64
SPRINT + REPA (Park et al., 2026)SD-VAE4001.59––0.800.64
SiT-XL/2 + IG (Zhou et al., 2026)SD-VAE8001.464.79265.70.800.64
SiT-XL/2 + Sparse Guid. (Krause et al., 2026)SD-VAE4001.584.45249.70.800.63
SRA 2 (Wang et al., 2026a)SD-VAE8001.524.63316.20.820.62
UDT+-XL/2 + REPA (Yun et al., 2026)SD-VAE3201.384.38306.30.790.66
RecFM-XL (Huang et al., 2026)SD-VAE1602.49––––
NEPA-DiT-XL + REPA (ours), ODE-96SD-VAE240+801.574.82298.00.790.63
NEPA-DiT-XL + REPA (ours), SDE-250SD-VAE240+801.324.32311.00.790.68

Guidance. Restricting guidance to an interval (KynkÀÀnniemi et al., 2024) admits larger scales: scale 3.6 on t∈[0.4,1] gives 1.57 and 1.32, the best with both samplers, selected by FID-50K.

Query rate. Querying at every step also sets the sampling cost of the NEPA model, 92 GFLOPs per step next to 309 for the generator; with 96 ODE steps, one image takes 62 TFLOPs in all, below the 91 TFLOPs of SiT-XL/2 with REPA with 250 SDE steps.

Guidance. Applied only for t in the interval. Shaded: setting of Table 9.
ScaleIntervalODE-96SDE-250
1.0[0, 1]8.777.59
1.4[0, 1]2.582.32
2.4[0.3, 1]1.721.46
3.6[0.4, 1]1.571.32
Parameters and sampling FLOPs. GFLOPs per sampling step for one image; the NEPA model runs on the 64 image tokens with the condition cached. Totals include guidance on t∈[0.4,1] and VAE decoding. SiT-XL/2 is counted in the same way.
ParamsGFLOPs / step
ModelGeneratorNEPAGeneratorNEPATFLOPs / image
NEPA-DiT-B135M711M5992–
NEPA-DiT-L458M711M21092–
NEPA-DiT-XL, ODE-96683M711M3099262
NEPA-DiT-XL, SDE-250683M711M30992160
SiT-XL/2 + REPA, SDE-250675M–229–91
Qualitative results of NEPA-DiT-XL on ImageNet 256Γ—256.
Qualitative results of NEPA-DiT-XL on ImageNet 256Γ—256.

Additional Qualitative Results

Figures 8–10 show more samples of NEPA-DiT-XL on ImageNet 256Γ—256, selected from 16 samples per class.

Goldfish (1)
Goldfish (1)
Bald eagle (22)
Bald eagle (22)
Macaw (88)
Macaw (88)
Sulphur-crested cockatoo (89)
Sulphur-crested cockatoo (89)
King penguin (145)
King penguin (145)
Golden retriever (207)
Golden retriever (207)
White wolf (270)
White wolf (270)
Arctic fox (279)
Arctic fox (279)
Lion (291)
Lion (291)
Tiger (292)
Tiger (292)
Otter (360)
Otter (360)
Lesser panda (387)
Lesser panda (387)
Giant panda (388)
Giant panda (388)
Balloon (417)
Balloon (417)
Space shuttle (812)
Space shuttle (812)
Cheeseburger (933)
Cheeseburger (933)
Mushroom (947)
Mushroom (947)
Pizza (963)
Pizza (963)
Alp (970)
Alp (970)
Cliff (972)
Cliff (972)
Coral reef (973)
Coral reef (973)
Valley (979)
Valley (979)
Volcano (980)
Volcano (980)
Daisy (985)
Daisy (985)

Conclusion and Future Work

We have shown that predictive embeddings can serve as the condition of an image generator. In Embedding Conditioned Generation, a DiT generator is conditioned at every denoising step on the clean-image embeddings that a NEPA model, trained with Multi-Embedding Prediction, predicts from the condition and the current noisy image. On ImageNet 256Γ—256, combined with REPA, NEPA-DiT-XL reaches an FID of 1.32 with the 250-step SDE sampler and 1.57 with the 96-step ODE sampler, trained with about a third of REPA's compute even though it trains and runs two networks.

Text conditions. Text-to-image models either encode the prompt once with a separate text encoder, which cannot adapt to the evolving noisy image and often fixes the length of the condition, or process the prompt together with the image in a single unified model, where the condition and the image share one backbone. A NEPA model could instead read the whole prompt before the noisy image and predict the clean image embeddings from both, as it does for class labels in this paper.

Unified models. Since the NEPA model is an autoregressive Transformer, a single NEPA model could in principle serve both as a language model with an output head and as the conditioning model of a DiT generator. Studying text conditions, other latent spaces, and higher resolutions are directions for future work.