Method / Clean-latent oracle
How clean-latent oracle guidance is calculated
The oracle encodes the complete correct video, measures the error of the DiT’s current clean-latent prediction, and uses its gradient to correct the noisy sampling state. This page describes the operations performed by the code for the 5-ball, 49-frame comparison.
Settings used in the comparison
| Video | 49 frames at 128 × 128 pixels |
|---|---|
| Conditioning | First 2 latent frames, corresponding to 5 video frames |
| Oracle target | VAE encoding of the complete ground-truth video |
| Guidance strength | \(\lambda=12\) |
| RMS floor | \(\varepsilon=10^{-6}\) |
| Guidance window | Every reverse step: start 0, end 1 |
| Reverse steps | \(K\in\{10,20,50,100,200\}\) |
| Sampler | Flow matching, shift 5, denoising strength 1 |
| Model parameters | DiT and VAE frozen; empty prompts and classifier-free guidance scale 1 |
Each latent tensor has shape \(B\times16\times13\times16\times16\): batch, channels, time, height, width. Let \(F\) contain all entries in latent time indices 2 through 12, and let \(C\) contain indices 0 and 1. Thus \(|F|=16\cdot11\cdot16\cdot16=45{,}056\) entries per video. \(P_F\) keeps the future entries and zeros the conditioning entries.
1. Encode the target and conditioning
The evaluator supplies the complete held-out video as oracle_target_frames. Its RGB values are scaled to \([-1,1]\). The sampler encodes it once with its VAE, including the VAE’s channel-wise latent scaling:
Here \(x^\star\) is the ground-truth video and \(z^\star\) is its clean target latent. The target is detached from the gradient graph and kept fixed throughout sampling. The encoder runs with tiling disabled.
The conditioning path also VAE-encodes the supplied video and keeps its first two latent frames. These are the fixed conditioning latents. Sampling starts from seeded Gaussian noise, with those two latent frames overwritten by the conditioning latents.
2. Evaluate the DiT at the current noise level
Index reverse steps by \(j=0,\ldots,K-1\). With shift 5, the scheduler constructs its noise levels and model timesteps as
The terminal noise level is \(\sigma_K=0\). At the start of each step, the code restores the conditioning entries in \(z_j\), detaches the state from the preceding step, and enables gradients with respect to this state. It then evaluates the frozen DiT once:
\(v_j\) is the model’s velocity prediction. \(\hat z_{0,j}\) is the clean-latent estimate used to calculate the oracle loss at this step.
3. Compute the future-latent MSE
For each video \(b\), subtract the target from the clean prediction on the future entries, square the differences, and average over channels, future latent frames, and spatial positions. Then average over the batch:
The subtraction and squared-error reduction use float32. In code, the future slice is [:, :, num_condition_frames:]. The oracle’s optimization loss is this latent MSE; Rollout-5 is calculated separately when evaluating the generated video.
4. Differentiate through the clean estimate
The implementation calculates the gradient in two operations. First it differentiates the loss with respect to the clean estimate and detaches that result. Then it propagates this gradient through the clean estimate to the current noisy state:
clean_gradient = torch.autograd.grad(
oracle_loss, clean_prediction,
only_inputs=True, retain_graph=True,
)[0].detach()
state_gradient = torch.autograd.grad(
(clean_prediction * clean_gradient).sum(),
latents, only_inputs=True,
)[0]
The gradient includes the dependence of \(v_\theta(z_j,t_j)\) on \(z_j\). Model weights stay fixed. The gradient graph covers the current DiT evaluation; each completed sampler update is detached before the next step.
5. Normalize the future-state gradient
Take the future slice of \(g_j\). For each video independently, calculate its root mean square in float32 and clamp it from below at \(10^{-6}\). Divide that video’s future gradient by the resulting scalar:
The RMS is cast back to the gradient’s dtype before division. The correction tensor starts as zeros, and only its future slice is filled. Its conditioning entries remain zero.
6. Apply the correction and decode the result
The code forms the ordinary flow-scheduler proposal using the same \(v_j\), subtracts the scaled normalized gradient, and restores the conditioning latents:
The gradient comes from the current state \(z_j\); the code subtracts its correction from the scheduler proposal \(\tilde z_{j+1}\). It applies one correction at each reverse step. After the final step, the VAE decodes the resulting latent into the generated video, with tiling disabled.
The code also records the final future-latent MSE, calculated between the final latent and the fixed target. The videos and Rollout-5 scores on the project page are evaluated from these final decoded outputs.
Where strength 12 comes from
The scale-selection code groups calibration results by guidance strength and chooses the candidate with the lowest mean original Rollout-5 error over seeds 0, 1, and 2. The saved selection for the 5-ball, 49-frame model is 12. The holdout run manifests use that same strength for all five reverse-step counts.
Implementation references
This description follows research-repository revision 109de50 and the saved oracle run settings.
src/sshv2/diffsynth/pipelines/wan_video_guided.py:_encode_oracle_target,_prepare_inputs,_oracle_call,_normalised_correction, and_restore_condition.src/sshv2/guidance/oracle.py:OracleLatentGuide.per_sample_lossand its batch mean.src/sshv2/diffsynth/models/wan_video_vae.py: latent scaling, encoding, decoding, and video/latent frame counts.lib/diffsynth/diffsynth/pipelines/wan_video_new.pyandlib/diffsynth/diffsynth/schedulers/flow_match.py: scheduler construction, noise levels, and Euler update.scripts/sshv2/eval_guidance.py: input video, conditioning, and oracle-target arguments.scripts/sshv2/select_clean_latent_oracle_seriality_scales.py: calibration selection rule.out/guidance/clean-latent-oracle-hf-seriality-subset-v1/CALIBRATION_SELECTIONS.json: selected strength.out/guidance/missing-ball-aware-rollout5-5ball49f-d-oracle-v1/oracle/5ball-49f/holdout/: settings for the displayed oracle runs.