reverseStep method
One Langevin step of the reverse chain given a predicted noise
epsHat (same shape as xt).
μ_t = (1 / √α_t) · (x_t − β_t / √(1 - α̅_t) · ε̂) x_{t-1} = μ_t + √β̃_t · z (z ~ N(0, I), t > 0)
At t == 0 the added noise is zero and the mean is returned
directly, so this method safely terminates the chain.
Implementation
Tensor reverseStep(Tensor xt, int t, Tensor epsHat, {int? seed}) {
_checkT(t);
if (!_shapesEqual(xt.shape, epsHat.shape)) {
throw ArgumentError(
'reverseStep: epsHat shape ${epsHat.shape} != xt shape ${xt.shape}',
);
}
final invSqrtA = sqrtRecipAlphas[t];
final coefEps = betas[t] / sqrtOneMinusAlphaBars[t];
final mean = (xt - epsHat * coefEps) * invSqrtA;
if (t == 0) return mean;
final sigma = math.sqrt(posteriorVariance[t]);
if (sigma == 0.0) return mean;
final z = _gaussianTensor(xt.shape, seed: seed, device: xt.device);
return mean + z * sigma;
}