How do features emerge during diffusion?¶

Here we ask: how does digit identity emerge as MNIST images are denoised? At low SNR, several classes are compatible with an observation. As SNR increases, the digit's identity becomes resolved. Feature information dynamics locates this transition along the diffusion noise scale.

Feature information dynamics turns this question into a curve. In this quickstart, you'll see the effect, estimate the curve, and try another digit.

Run all cells to use the provided MNIST denoiser (~8 MB). GPU recommended; CPU also works. Saved figures let you follow along immediately.

In [1]:
from pathlib import Path
import json
import numpy as np
import matplotlib.pyplot as plt
import torch
from feature_information_dynamics.mnist_data import load_mnist
from feature_information_dynamics.mnist_calibrated import SharedDenoiser
from feature_information_dynamics.noise import stable_noise
from feature_information_dynamics.bundle import sha256

ROOT = Path.cwd().resolve()
if ROOT.name == "notebooks": ROOT = ROOT.parent
ASSETS = ROOT / "examples" / "mnist_class"
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.set_num_threads(4)
torch.backends.cuda.matmul.allow_tf32 = False
torch.backends.cudnn.allow_tf32 = False
plt.rcParams.update({"figure.dpi": 110, "font.size": 11})
asset = json.loads((ASSETS / "checkpoint.json").read_text())
checkpoint = ASSETS / "selected.pt"
if sha256(checkpoint) != asset["sha256"]: raise ValueError("Tutorial checkpoint checksum mismatch")
state = torch.load(checkpoint, map_location="cpu", weights_only=True)
model = SharedDenoiser(state["config"]["channels"]).to(device)
model.load_state_dict(state["model"])
model.eval()
arrays, _ = load_mnist(ROOT / "data" / "mnist", download=True)
images = arrays["t10k-images-idx3-ubyte.gz"]
labels = arrays["t10k-labels-idx1-ubyte.gz"]

1. Watch digit identity emerge¶

Read left to right, from noise toward a clear image. The unconditional denoiser averages over plausible digits; the class-conditional denoiser averages only within one class. Their reconstructions reveal where class ambiguity remains and where the observation resolves it. These are clean-image predictions at fixed noise levels, rather than samples from a reverse diffusion trajectory.

The images below share the same clean digit and noise across ten SNRs. See the paper's Diffusion and flow-matching parameterization for this noise coordinate.

In [2]:
@torch.no_grad()
def prepare_digit(digit=7, example=0, levels=None):
    position = np.flatnonzero(labels == digit)[example]
    x = torch.from_numpy(images[position:position+1]).to(device)
    y = torch.tensor([digit], device=device)
    eps = torch.randn((1, 784), generator=torch.Generator().manual_seed(20261006)).to(device)
    frames = []
    for log_snr in visual_grid if levels is None else levels:
        t = torch.full((1, 1), 1/(1 + 10**(-log_snr/2)), device=device)
        noisy = t*x + (1-t)*eps
        frames.append(torch.cat([noisy, model(noisy,t,y,False), model(noisy,t,y,True)]).cpu().numpy().reshape(3,28,28))
    return {"digit": digit, "example": example, "clean": x.cpu().numpy().reshape(28,28), "frames": np.asarray(frames)}

def show_digit(digit=7, example=0):
    sample = prepare_digit(digit, example)
    fig, axes = plt.subplots(3, 10, figsize=(16,5))
    for col, level in enumerate(visual_grid):
        for row in range(3):
            axes[row,col].imshow(sample["frames"][col,row], cmap="gray", vmin=-1,vmax=1)
            axes[row,col].set_xticks([]); axes[row,col].set_yticks([])
        axes[0,col].set_title(f"{level:.2f}", fontsize=10)
    for row, name in enumerate(["Noisy", "Unconditional", f"Class {digit}"]): axes[row,0].set_ylabel(name)
    fig.suptitle(f"Digit {digit}: noise → resolved identity   |   log10 SNR", fontsize=14)
    fig.tight_layout(); fig.savefig(preview_dir / f"contact_sheet_{digit}.png", bbox_inches="tight"); plt.show()
    return sample

visual_grid = np.linspace(-2, 2, 10)
preview_dir = ROOT / "docs" / "evidence" / "mnist_quickstart"
preview_dir.mkdir(parents=True, exist_ok=True)
samples = [show_digit(digit) for digit in (7, 3, 8)]
No description has been provided for this image
No description has been provided for this image
No description has been provided for this image

2. Measure class ambiguity across noise levels¶

For clean images $X$ and labels $Y$, the paper uses the Gaussian channel $X_\gamma=\sqrt\gamma X+N$. We feed the model its equivalent bounded form $X_t=tX+(1-t)N$, where $\gamma=(t/(1-t))^2$.

At each SNR, measure the reconstruction error without and with the label. MMSE is the smallest possible average squared reconstruction error; its optimal denoiser is the conditional mean. Comparing the two MMSEs isolates the uncertainty associated with class identity, beyond variation within a class:

$$\widehat\Delta_Y(\gamma)=\widehat m_\varnothing(\gamma)-\widehat m_Y(\gamma).$$

The workflow is add noise → denoise twice → measure squared errors → subtract. Both modes see identical images and noise. The cell below runs this on the full 10,000-image test split with two noise draws. Our trained networks approximate these optimal denoisers. See Gaussian channel and MMSE and Practical estimators in the paper.

In [3]:
grid = np.linspace(-5, 3, 33)             # log10 SNR
seeds = [20261004, 20261005]
batch_size = 200
errors = np.empty((2, len(grid), len(seeds), len(images)), dtype=np.float32)
ids = [f"test:{i}" for i in range(len(images))]

with torch.no_grad():
    for repeat, seed in enumerate(seeds):
        noise = stable_noise(ids, 784, seed)
        for j, log_snr in enumerate(grid):
            for start in range(0, len(images), batch_size):
                stop = min(start + batch_size, len(images))
                x = torch.from_numpy(images[start:stop]).to(device)
                y = torch.from_numpy(labels[start:stop]).to(device)
                eps = torch.from_numpy(noise[start:stop]).to(device)
                t = torch.full((len(x), 1), 1/(1 + 10**(-log_snr/2)), device=device)
                noisy = t*x + (1-t)*eps
                without_label = model(noisy, t, y, False)
                with_label = model(noisy, t, y, True)
                errors[0, j, repeat, start:stop] = (without_label-x).double().square().sum(1).cpu().numpy()
                errors[1, j, repeat, start:stop] = (with_label-x).double().square().sum(1).cpu().numpy()
        print(f"Noise draw {repeat + 1}/{len(seeds)} complete")

risks = errors.astype(np.float64).mean(axis=(2, 3))
gap = risks[0] - risks[1]
Noise draw 1/2 complete
Noise draw 2/2 complete
In [4]:
fig, axes = plt.subplots(1, 2, figsize=(10, 3.2))
axes[0].plot(grid, risks[0], label="Without label")
axes[0].plot(grid, risks[1], label="With label")
axes[0].fill_between(grid, risks[1], risks[0], alpha=.15)
axes[0].set(ylabel="Reconstruction error", title="Two denoising risks")
axes[0].legend()
axes[1].plot(grid, gap)
axes[1].axhline(0, color="gray", linewidth=.7)
axes[1].set(ylabel="MMSE gap estimate", title="Class ambiguity across SNR")
for ax in axes:
    ax.set_xlabel("log10 SNR"); ax.grid(alpha=.2)
fig.tight_layout(); plt.show()
No description has been provided for this image

3. Locate where class information emerges¶

The I-MMSE relation converts that error reduction into the paper's feature information density on the log-SNR axis:

$$D_Y^{(\log\gamma)}=\frac{\gamma}{2}\bigl(m_\varnothing-m_Y\bigr).$$

This density is the rate of increase of $I(Y;X_\gamma)$ with respect to $\ln\gamma$: how quickly the noisy image reveals digit identity. A peak marks a noise scale where class information accumulates most rapidly. The MMSE gap alone is not this rate on a logarithmic axis; the factor $\gamma/2$ comes from I-MMSE and the change of coordinates.

For the derivation, see the paper's Feature Information Density section, particularly Gaussian channel and MMSE and I-MMSE relation and feature information density. The implementation's time conversion follows Diffusion and flow-matching parameterization.

In [5]:
density = .5 * 10.0**grid * gap

fig, ax = plt.subplots(figsize=(8, 3.2))
ax.plot(grid, density, "o-", color="tab:purple")
ax.fill_between(grid, 0, density, color="tab:purple", alpha=.12)
ax.axhline(0, color="gray", linewidth=.7)
ax.set(xlabel="log10 SNR", ylabel="Density per ln SNR", title="Where class information appears")
ax.grid(alpha=.2); fig.tight_layout(); plt.show()
No description has been provided for this image

Explore the transition¶

The density tells us where class information enters fastest. Integrating it summarizes class generation progress. An exact density is nonnegative; finite neural estimates can have a slightly negative tail. For the progress display only, we integrate the nonnegative part of the estimated density on our grid:

$$P(\ell)=\frac{\int_{-5}^{\ell}[\widehat D_Y^{(\ln\gamma)}(u)]_+\,du} {\int_{-5}^{3}[\widehat D_Y^{(\ln\gamma)}(u)]_+\,du}, \qquad \ell=\log_{10}\gamma,\quad [d]_+=\max(d,0).$$

The factor $\ln(10)$ cancels in this ratio. This display is monotone from 0 to 100% within the measured SNR range. It is not classifier accuracy or an individual digit's probability. The original signed losses, gaps, density, and integral are preserved; this convention does not improve estimation accuracy.

The next previews show ten stages, from 5% to 95%. Frames are evaluated at those SNRs, not interpolated images. The interactive page links them to all four curves.

In [6]:
import runpy
from IPython.display import FileLink, display
cumulative = np.r_[0., np.cumsum(.5*(density[1:]+density[:-1])*np.diff(grid)*np.log(10))]
raw_progress = cumulative / cumulative[-1]
progress_density = np.maximum(density, 0.)
positive_cumulative = np.r_[0., np.cumsum(.5*(progress_density[1:]+progress_density[:-1])*np.diff(grid)*np.log(10))]
if positive_cumulative[-1] <= 0: raise ValueError("No positive density to normalize")
progress = positive_cumulative / positive_cumulative[-1]
stage_targets = np.arange(.05, 1., .10)
stage_levels = []
for target in stage_targets:
    j = np.flatnonzero((progress[:-1] < target) & (progress[1:] >= target))[0]
    stage_levels.append(np.interp(target, progress[j:j+2], grid[j:j+2]))
stage_levels = np.asarray(stage_levels)
interactive_levels = np.unique(np.r_[grid, stage_levels])
interactive_samples = [prepare_digit(s["digit"], s["example"], interactive_levels) for s in samples]
stage_indices = np.array([np.argmin(abs(interactive_levels-level)) for level in stage_levels])
np.savez_compressed(preview_dir / "preview_data.npz", visual_grid=visual_grid,
    grid=grid, risks=risks, gap=gap, density=density, digits=np.array([s["digit"] for s in samples]),
    clean=np.stack([s["clean"] for s in samples]),
    frames=np.stack([s["frames"] for s in samples]), progress=progress,
    raw_progress=raw_progress, progress_density=progress_density,
    stage_targets=stage_targets, stage_levels=stage_levels,
    interactive_levels=interactive_levels,
    interactive_frames=np.stack([s["frames"] for s in interactive_samples]),
    stage_frames=np.stack([s["frames"][stage_indices] for s in interactive_samples]))
runpy.run_path(str(ROOT / "tools" / "render_mnist_preview.py"),
               init_globals={"ROOT": ROOT}, run_name="__main__")
display(FileLink("../docs/preview/index.html", result_html_prefix="Open interactive preview: "))
No description has been provided for this image
No description has been provided for this image
No description has been provided for this image
No description has been provided for this image
No description has been provided for this image
No description has been provided for this image
Offline interactive preview and six progress-layout figures saved.
Open interactive preview: ../docs/preview/index.html

A quick check¶

The area under the curve estimates class information. Since the label is determined by the clean digit, its entropy provides a reference (about $\ln10$ nats). Integration on a log10 axis needs a factor of $\ln10$.

Neural denoisers approximate the optimal MMSE, so the estimated curve can be biased and occasionally negative. We keep those values; the check below reports the raw finite-range integral.

In [7]:
area_weights = .5 * np.diff(grid) * np.log(10)
integral = float(np.sum((density[1:] + density[:-1]) * area_weights))
p = np.bincount(labels, minlength=10) / len(labels)
entropy = float(-np.sum(p * np.log(p)))
relative_error = abs(integral - entropy) / entropy
print(f"Estimated class information: {integral:.3f} nats")
print(f"Class entropy reference:     {entropy:.3f} nats")
print(f"Relative integral error:     {relative_error:.1%}")

result = {"estimate_nats": integral, "reference_nats": entropy,
          "relative_error": relative_error, "passed": bool(relative_error <= .2),
          "sample_count": len(images), "noise_repeat_count": len(seeds),
          "checkpoint_sha256": sha256(checkpoint), "log10_snr": grid.tolist()}
output = ROOT / "runs" / "mnist_quickstart"
output.mkdir(parents=True, exist_ok=True)
(output / "result.json").write_text(json.dumps(result, indent=2) + "\n")
Estimated class information: 2.626 nats
Class entropy reference:     2.301 nats
Relative integral error:     14.1%
Out[7]:
636

Try it yourself¶

Change the digit and example below. Where does its identity become visually resolved? Compare this individual example with the population density curve.

In [8]:
# Change these two values, then run this cell.
DIGIT = 3
EXAMPLE = 1
RUN_EXPERIMENT = False
if RUN_EXPERIMENT:
    show_digit(DIGIT, EXAMPLE)

Optional: train your own denoiser¶

To build this estimator from scratch, train with and without class labels on paired noisy images. We provide a shared U-Net trainer using 55,000 training images and 5,000 held-out validation images. The validation risk selects the checkpoint.

Set RUN_TRAINING=True to run the fixed 100,000-update budget in a new directory. Then load that run's selected.pt and repeat the measurement above. See the paper's Practical estimators subsection for the connection to MMSE estimation.

In [9]:
RUN_TRAINING = False
if RUN_TRAINING:
    from feature_information_dynamics.mnist_calibrated import train
    train(ROOT / "configs" / "mnist_concepts_full.json",
          ROOT / "data" / "mnist", ROOT / "runs" / "my_mnist_denoiser", str(device))

Beyond one feature: chained information decomposition¶

Class identity is one feature. To track class, object mask, and Canny edges together, use cumulative conditions:

$$\varnothing\ \longrightarrow\ Y_1\ \longrightarrow\ (Y_1,Y_2) \ \longrightarrow\ (Y_1,Y_2,Y_3).$$

Let $m_k$ be the MMSE with the first $k$ features supplied. Each adjacent pair isolates an additional information component:

$$D_{k\mid k-1}^{(\log\gamma)} =\frac{\gamma}{2}(m_{k-1}-m_k),\qquad \sum_{k=1}^{K}D_{k\mid k-1}^{(\log\gamma)} =\frac{\gamma}{2}(m_0-m_K).$$

This is a chain of increments: mask contributes beyond class, and edges contribute beyond class and mask. Shared information is allocated along the chain rather than counted repeatedly. Changing the order changes the allocation. The same paired-noise measurement applies to each adjacent pair; this notebook measures only its class step. See the paper's Chained Information Decomposition section.