reverseStep method

Tensor reverseStep(
  1. Tensor xt,
  2. int t,
  3. Tensor epsHat, {
  4. int? seed,
})

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;
}