Official implementation of Measurement-Consistent Langevin Corrector for Stabilizing Latent Diffusion Inverse Problem Solvers (ICML 2026).
MCLC is a lightweight, plug-and-play corrector for latent diffusion inverse solvers. After a solver's measurement-consistency update, MCLC moves the latent toward the diffusion time marginal while restricting the update to the orthogonal complement of the measurement-consistency gradient. This preserves measurement consistency locally and helps stabilize the reverse process.
The example implementation code will be released later. The function below is the core MCLC update.
import torch
@torch.no_grad()
def mclc(
z: torch.Tensor,
score_fn,
measurement_grad: torch.Tensor,
num_steps: int = 3,
target_snr: float = 0.01,
eps: float = 1e-8,
) -> torch.Tensor:
"""Apply the Measurement-Consistent Langevin Corrector.
Args:
z: Latent after a solver's measurement-consistency update. This
reference implementation assumes a batch size of 1.
score_fn: Callable mapping the current latent to the diffusion score
at the current timestep. When using an epsilon-prediction model,
the score is "-epsilon_theta(z, t) / sigma_t".
measurement_grad: Gradient of the measurement residual with respect
to the latent, evaluated before the inner corrector iterations.
It is intentionally kept fixed throughout this function.
num_steps: Number of inner Langevin correction steps.
target_snr: Target signal-to-noise ratio used for adaptive step sizes.
eps: Numerical stability constant.
Returns:
The corrected latent, with the same shape as "z".
"""
if z.shape[0] != 1 or measurement_grad.shape[0] != 1:
raise ValueError("This reference implementation assumes batch size 1.")
# Unit vector pointing in the measurement-consistency direction.
# It is computed once and kept fixed throughout the inner MCLC steps.
g = measurement_grad.detach().reshape(-1)
g_hat = g / (g.norm() + eps)
def projection_onto_orthogonal_complement(v: torch.Tensor) -> torch.Tensor:
"""Remove the component of v parallel to the measurement gradient.
P_g^ortho(v) = v (I- gg^T)
"""
v_flat = v.reshape(-1)
parallel_component = torch.dot(v_flat, g_hat) * g_hat
return (v_flat - parallel_component).reshape_as(v)
for _ in range(num_steps):
score = score_fn(z)
noise = torch.randn_like(z)
# Neither update is allowed to move along the measurement-gradient
# direction.
proj_score = projection_onto_orthogonal_complement(score)
proj_noise = projection_onto_orthogonal_complement(noise)
score_norm = proj_score.norm()
noise_norm = proj_noise.norm()
step_size = (target_snr * noise_norm / (score_norm + eps)) ** 2
z = (
z
+ step_size * proj_score
+ torch.sqrt(2.0 * step_size) * proj_noise
)
return zThe projection direction must be the measurement-consistency gradient already computed by the base solver. Keeping it fixed during the inner corrector steps both avoids additional backward passes and matches the theoretical MCLC construction.
Apply MCLC immediately after the base solver's measurement-consistency update:
# Reuse the gradient computed by the inverse solver.
measurement_grad = torch.autograd.grad(measurement_error, z_t)[0]
z_prev = z_prev - measurement_grad
sigma_t = torch.sqrt(1.0 - alpha_bar_prev)
def score_fn(z):
epsilon = diffusion_model(z, t_prev, condition)
return -epsilon / sigma_t
z_prev = mclc(
z=z_prev,
score_fn=score_fn,
measurement_grad=measurement_grad,
num_steps=3,
target_snr=0.01,
)The appropriate corrector frequency, number of inner steps, and target SNR can depend on the base solver and task.
- Release the full example implementation based on PSLD.
If you find this work useful, please cite:
@inproceedings{
hyoseok2026measurementconsistent,
title={Measurement-Consistent Langevin Corrector for Stabilizing Latent Diffusion Inverse Problem Solvers},
author={Lee Hyoseok and Sohwi Lim and Eunju Cha and Tae-Hyun Oh},
booktitle={Forty-third International Conference on Machine Learning},
year={2026},
url={https://openreview.net/forum?id=QC7fOKv1jg}
}