An image diffusion model denoises one picture. A video model denoises a whole clip at once, so every frame stays consistent with every other. That one change, from a 2D canvas to a 3D block of space and time, turns a model that fits comfortably on one GPU into a workload measured in GPU-hours per minute of output. To understand video generation as an engineer, you need to follow the data: how many tokens a clip becomes, what attention does with them, and where memory and time go.
This page builds that picture from first principles, using the openly published Wan2.1 configuration as a concrete reference. The diffusion maths itself is in flow matching and the maths of Stable Diffusion; here the focus is the architecture and the GPU consequences.
From image diffusion to video diffusion
Early video diffusion models extended image U-Nets by factorising: 2D spatial convolutions and attention within each frame, plus 1D temporal attention across frames at each pixel position. Factorisation is cheap, because no operation ever looks at all of space-time together, and it lets you start from pretrained image weights. Its weakness is coherence over long motion, since information has to hop between the two axes layer by layer.
Current open models such as HunyuanVideo, CogVideoX and Wan follow the path that Sora's technical report popularised: compress the video with a 3D autoencoder, cut the latent into spacetime patches, and run a transformer whose self-attention spans every patch in the clip. Full 3D attention models motion directly and scales cleanly with parameters and data. Its price is a sequence length that grows with frames times height times width, and an attention cost that grows with the square of that.
Step 1: compress with a causal 3D VAE
Running a transformer on raw pixels would be hopeless, so the clip is first encoded into a latent. Wan2.1's VAE uses a stride of 4 in time and 8 in each spatial axis, with 16 latent channels. It is causal in time: each latent frame depends only on current and past frames, and the first frame is encoded on its own. That is why Wan expects 4n+1 frames: 81 frames become 1 + 80/4 = 21 latent frames. Causality also lets the same encoder handle a single image as a one-frame video, so image and video data can be mixed in training.
The VAE is not free. Decoding a long high-resolution clip is memory-hungry because of the 3D convolutions, which is why implementations decode in temporal chunks with cached boundary features, or in spatial tiles with overlap blending. Tile seams and chunk-boundary flicker are classic artefacts of getting this wrong.
Step 2: patchify into spacetime tokens
The latent is cut into patches of 1 latent frame by 2 by 2 latent pixels, each projected to the model width of 5,120 for the 14B model. Count the tokens for a 5-second, 16 fps, 832 by 480 clip: 21 latent frames, a 60 by 104 latent grid, 30 by 52 patches per frame, so 21 times 1,560 = 32,760 tokens. A 1280 by 720 clip gives 21 times 45 times 80 = 75,600 tokens. For comparison, a 1024 by 1024 image in a typical latent DiT is around 4,096 tokens.
def video_tokens(frames, height, width, vae_stride=(4, 8, 8), patch=(1, 2, 2)):
# Causal VAE: the first frame is encoded alone, then every 4 frames become one latent frame.
t_lat = 1 + (frames - 1) // vae_stride[0]
h_lat, w_lat = height // vae_stride[1], width // vae_stride[2]
return (t_lat // patch[0]) * (h_lat // patch[1]) * (w_lat // patch[2])
def forward_flops(tokens, params=14e9, layers=40, dim=5120):
dense = 2 * params * tokens # every weight used once per token
attn = 4 * tokens * tokens * dim * layers # QK^T and AV, all layers
return dense, attn
for name, (f, h, w) in {"480p": (81, 480, 832), "720p": (81, 720, 1280)}.items():
n = video_tokens(f, h, w)
dense, attn = forward_flops(n)
total = 100 * (dense + attn) # 50 steps, CFG doubles the forwards
print(f"{name}: {n:,} tokens, attention share {attn / (dense + attn):.0%}, "
f"{total:.1e} FLOPs per clip")
# 480p: 32,760 tokens, attention share 49%, 1.8e+17 FLOPs per clip
# 720p: 75,600 tokens, attention share 69%, 6.8e+17 FLOPs per clipThe FLOP figures are order-of-magnitude estimates: they count only the dense weights and the two attention matmuls and ignore the VAE, the text encoder and cross-attention. Even so the shape is clear. At 480p attention is already about half the work; at 720p it dominates. At an assumed sustained 400 TFLOPS of BF16 on one large data-centre GPU, a 720p clip is roughly half an hour of compute, which is why multi-GPU inference is the norm for the large models.
Step 3: the diffusion transformer
Each of the 40 blocks does three things. Self-attention over all 32,760 tokens lets any patch see any other patch in space and time; 3D rotary position embeddings tell it where each token sits. Cross-attention lets the video tokens read the text embeddings from umT5-XXL, a multilingual encoder that runs once per prompt. A feed-forward network, 13,824 wide in the 14B model, does per-token computation. The diffusion timestep conditions every block through adaptive layer norm, which scales and shifts the normalised activations.
Materialising the attention matrix is impossible: 32,760 squared scores in BF16 is about 2 GB per head per layer, and there are 40 heads. FlashAttention-style kernels are mandatory, computing attention in tiles that never leave on-chip memory; see FlashAttention on GPUs. With them, memory grows linearly with sequence length and the remaining cost is pure arithmetic.
Training: flow matching on long sequences
Wan and several peers train with flow matching. Sample noise, pick a time between 0 and 1, interpolate on a straight line between noise and the clean latent, and train the network to predict the velocity along that line. Captions are dropped some fraction of the time so the same model can run classifier-free guidance later.
# Flow-matching training step for a video DiT (schematic, PyTorch-style).
for batch in loader: # videos and captions
with torch.no_grad():
x1 = vae.encode(batch.video) # [B, 16, T', H', W'] latents
ctx = text_encoder(batch.caption) # frozen; often precomputed offline
x0 = torch.randn_like(x1) # noise
t = torch.sigmoid(torch.randn(x1.shape[0], device=x1.device)) # timestep density, a design choice
xt = (1 - t.view(-1, 1, 1, 1, 1)) * x0 + t.view(-1, 1, 1, 1, 1) * x1
target = x1 - x0 # velocity along the straight path
if random.random() < 0.1:
ctx = null_ctx # caption dropout enables CFG at inference
with torch.autocast("cuda", dtype=torch.bfloat16):
v = dit(xt, t, ctx) # sequence-parallel inside the attention layers
loss = F.mse_loss(v.float(), target.float())
loss.backward()
clip_grad_norm_(dit.parameters(), 1.0)
opt.step(); opt.zero_grad(set_to_none=True)The GPU problems are about sequence length, not parameter count. Activations for one sample at 75,600 tokens and width 5,120 are about 770 MB per tensor in BF16, and each block keeps several, so activation checkpointing is standard. Weights, gradients and optimiser state for 14B parameters do not fit on one device, so they are sharded with FSDP or ZeRO. And one sample may not fit on one GPU at all, so attention itself is split. Ulysses-style sequence parallelism uses an all-to-all to swap from splitting tokens to splitting heads for the attention computation; ring attention keeps tokens split and passes key and value blocks around a ring, as described in ring attention. Ulysses is limited by the head count and needs fast all-to-all inside a node; ring attention scales further but needs compute to hide its communication.
Data shapes the curriculum. Teams typically pretrain on images and low-resolution short clips, where tokens are cheap, then progressively raise resolution and duration. Batches mix clip lengths, so bucketing by token count keeps GPUs from idling on padding.
Inference: where the time goes and how to cut it
Sampling multiplies the forward pass. Wan's text-to-video default is 50 steps, and classifier-free guidance at its default scale of 5 runs a conditional and an unconditional pass each step, so a clip costs about 100 transformer forwards plus one VAE decode. The levers, from cheapest to most invasive:
- Fit first. The Wan2.1 README quotes 8.19 GB of VRAM for the 1.3B text-to-video model; offloading trades PCIe traffic for memory, so it is slow but works on consumer cards.
- Parallelise one clip. FSDP shards weights; Ulysses or ring attention splits the sequence across GPUs, cutting latency for a single request.
- Fewer steps. Better solvers and step-distilled models reduce 50 steps to a handful, at some cost in detail and motion quality.
- Cheaper guidance. Batch the conditional and unconditional passes together, apply guidance only on some steps, or use guidance-distilled weights.
- Reuse work across steps. Adjacent denoising steps change features slowly, so caching and reusing block outputs skips computation; quality must be checked per model.
- Sparse or windowed attention. Restricting attention to local spacetime windows attacks the quadratic term directly, but it changes the model and needs fine-tuning.
# Single GPU, 1.3B model at 480p: offload weights and the text encoder to fit a consumer card
python generate.py --task t2v-1.3B --size 832*480 --ckpt_dir ./Wan2.1-T2V-1.3B \
--offload_model True --t5_cpu --prompt "a red fox running through fresh snow"
# Eight GPUs, 14B model: shard weights with FSDP and split attention with Ulysses
torchrun --nproc_per_node=8 generate.py --task t2v-14B --size 1280*720 \
--ckpt_dir ./Wan2.1-T2V-14B --dit_fsdp --t5_fsdp --ulysses_size 8 \
--prompt "a red fox running through fresh snow"
Worked example: planning a 720p service
Suppose a product needs 5-second 720p clips from the 14B model. From the estimate above, one clip is about 6.8 times 10 to the 17 FLOPs. At an assumed 40 percent utilisation of a GPU with roughly 1 PFLOPS of dense BF16, that is about 400 TFLOPS sustained and close to 1,700 GPU-seconds per clip, around 3.5 minutes on 8 GPUs if sequence parallelism scales well. A pool of 64 GPUs then serves roughly 140 clips an hour. If the product needs ten times that, step distillation to 8 steps without CFG cuts forwards from 100 to 8, a 12-fold reduction, and the decision becomes a quality evaluation rather than a hardware purchase. These are planning numbers; measure your own throughput before committing.
Failure modes
- Out of memory at decode. The transformer fits but the VAE decode does not. Use chunked or tiled decoding.
- Frame count not 4n+1. Causal VAEs silently drop or pad frames; generate the lengths the model was trained on.
- Temporal flicker and morphing. Often a sign of too few steps, too low guidance, or resolutions outside the training buckets.
- Poor multi-GPU scaling. Ulysses across nodes without fast all-to-all stalls on communication; keep it inside an NVLink domain and use ring attention or data parallelism across nodes.
- Precision drift. Accumulating attention or the sampler state in low precision causes colour shifts; keep softmax and solver arithmetic in FP32.
- Training stalls on data loading. Decoding video is CPU-heavy; precompute latents and text embeddings offline.
Trade-offs to decide explicitly
Full 3D attention gives the best coherence and the worst scaling; factorised or windowed attention is cheaper and weaker on long motion. Higher VAE compression shortens sequences and speeds everything, at the cost of fine detail and harder reconstruction. More sampling steps raise quality linearly in cost; distillation buys speed with a quality risk that must be measured on your prompts. Longer clips are usually better produced by extending or chaining shorter generations than by training ever-longer contexts. For how video understanding models treat the same tokens, see video LLMs on GPUs.
What to do next
- Compute tokens per clip for the resolutions and durations you need using the function above.
- Run the 1.3B model on one GPU with offloading and measure seconds per step and peak memory.
- Profile one forward pass and confirm attention dominates at your target resolution.
- Try multi-GPU inference with FSDP plus Ulysses inside one node and record scaling efficiency.
- Evaluate one step-reduction method on a fixed prompt set before adopting it.
- If training, precompute latents and text embeddings, bucket by token count, and enable activation checkpointing from the start.