Skip to content

Latest commit

 

History

1 Commit

Folders and files

NameName
Last commit message
Last commit date
 
 

Repository files navigation

Measurement-Consistent Langevin Corrector (MCLC)

Project Page | arXiv

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.

Algorithm

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 z

The 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.

Usage

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.

TODO

  • Release the full example implementation based on PSLD.

Citation

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

About

[ICML'26] Official repository of "Measurement-Consistent Langevin Corrector for Stabilizing Latent Diffusion Inverse Problem Solvers"

Resources

Stars

2 stars

Watchers

1 watching

Forks

Releases

Packages

Contributors