SabaPivot's picture
download
raw
22.1 kB
#!/usr/bin/env python3
"""Scaled UNet/DiT split-consistency experiment for ICML 2026 paper #2449.
The paper uses FFHQ and 50k-step EDM training. This bounded-cost reproduction
uses FashionMNIST, compact genuine convolutional UNet and patch-transformer DiT
backbones, two disjoint splits, two dataset sizes, and deterministic VE/EDM
probability-flow sampling. Key results are printed as well as written out so
that immutable Hugging Face Job logs retain the evidence if artifact upload
fails.
"""
from __future__ import annotations
import argparse
import copy
import csv
import json
import math
import os
import random
import time
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchvision.datasets import FashionMNIST
from torchvision.transforms import Compose, Normalize, ToTensor
from torchvision.utils import make_grid, save_image
def seed_all(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
class SigmaEmbedding(nn.Module):
def __init__(self, dim: int):
super().__init__()
half = dim // 2
self.register_buffer("freq", torch.exp(torch.linspace(math.log(1.0), math.log(1000.0), half)))
self.mlp = nn.Sequential(nn.Linear(dim, dim), nn.SiLU(), nn.Linear(dim, dim))
def forward(self, sigma: torch.Tensor) -> torch.Tensor:
phase = torch.log(sigma.clamp_min(1e-5))[:, None] * self.freq[None, :]
return self.mlp(torch.cat([phase.sin(), phase.cos()], dim=1))
class TinyUNet(nn.Module):
def __init__(self, base: int = 32, emb_dim: int = 64):
super().__init__()
self.emb = SigmaEmbedding(emb_dim)
self.in_conv = nn.Conv2d(1, base, 3, padding=1)
self.down = nn.Conv2d(base, base * 2, 4, stride=2, padding=1)
self.mid1 = nn.Conv2d(base * 2, base * 2, 3, padding=1)
self.mid2 = nn.Conv2d(base * 2, base * 2, 3, padding=1)
self.up = nn.Conv2d(base * 3, base, 3, padding=1)
self.out = nn.Conv2d(base, 1, 3, padding=1)
self.e1 = nn.Linear(emb_dim, base)
self.e2 = nn.Linear(emb_dim, base * 2)
self.e3 = nn.Linear(emb_dim, base * 2)
self.n1 = nn.GroupNorm(8, base)
self.n2 = nn.GroupNorm(8, base * 2)
self.n3 = nn.GroupNorm(8, base * 2)
self.n4 = nn.GroupNorm(8, base)
def forward(self, x: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor:
e = self.emb(sigma)
h1 = F.silu(self.n1(self.in_conv(x) + self.e1(e)[:, :, None, None]))
h2 = F.silu(self.n2(self.down(h1) + self.e2(e)[:, :, None, None]))
h3 = F.silu(self.n3(self.mid1(h2) + self.e3(e)[:, :, None, None]))
h3 = F.silu(self.n3(self.mid2(h3) + h2))
h3 = F.interpolate(h3, size=h1.shape[-2:], mode="nearest")
h4 = F.silu(self.n4(self.up(torch.cat([h3, h1], dim=1))))
return self.out(h4)
class TinyDiT(nn.Module):
def __init__(self, patch: int = 4, hidden: int = 96, depth: int = 3, heads: int = 4):
super().__init__()
self.patch = patch
self.grid = 28 // patch
self.patch_embed = nn.Conv2d(1, hidden, patch, stride=patch)
self.pos = nn.Parameter(torch.randn(1, self.grid * self.grid, hidden) * 0.02)
self.emb = SigmaEmbedding(hidden)
layer = nn.TransformerEncoderLayer(
hidden, heads, hidden * 4, dropout=0.0, activation="gelu",
batch_first=True, norm_first=True,
)
self.blocks = nn.TransformerEncoder(layer, depth)
self.norm = nn.LayerNorm(hidden)
self.out = nn.Linear(hidden, patch * patch)
def forward(self, x: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor:
tokens = self.patch_embed(x).flatten(2).transpose(1, 2)
tokens = tokens + self.pos + self.emb(sigma)[:, None, :]
patches = self.out(self.norm(self.blocks(tokens)))
b = x.shape[0]
patches = patches.view(b, self.grid, self.grid, self.patch, self.patch)
return patches.permute(0, 1, 3, 2, 4).reshape(b, 1, 28, 28)
class EDMPrecond(nn.Module):
def __init__(self, backbone: nn.Module, sigma_data: float = 0.5):
super().__init__()
self.backbone = backbone
self.sigma_data = sigma_data
def forward(self, x: torch.Tensor, sigma: torch.Tensor) -> torch.Tensor:
s = sigma[:, None, None, None]
sd = self.sigma_data
c_skip = sd * sd / (s * s + sd * sd)
c_out = s * sd / torch.sqrt(s * s + sd * sd)
c_in = 1.0 / torch.sqrt(s * s + sd * sd)
return c_skip * x + c_out * self.backbone(c_in * x, sigma)
def make_model(arch: str) -> EDMPrecond:
if arch == "unet":
return EDMPrecond(TinyUNet())
if arch == "dit":
return EDMPrecond(TinyDiT())
raise ValueError(arch)
def train_model(arch: str, train_images: torch.Tensor, steps: int, batch: int,
model_seed: int, device: torch.device) -> dict[str, torch.Tensor]:
seed_all(model_seed)
model = make_model(arch).to(device)
ema = copy.deepcopy(model).eval()
for p in ema.parameters():
p.requires_grad_(False)
optimizer = torch.optim.AdamW(model.parameters(), lr=2e-4, weight_decay=1e-4)
scaler = torch.amp.GradScaler("cuda", enabled=device.type == "cuda")
generator = torch.Generator(device="cpu").manual_seed(model_seed + 991)
losses = []
model.train()
started = time.perf_counter()
for step in range(steps):
idx = torch.randint(0, train_images.shape[0], (batch,), generator=generator)
clean = train_images[idx].to(device, non_blocking=True)
sigma = torch.exp(-1.2 + 1.2 * torch.randn(batch, device=device)).clamp(0.02, 3.0)
noisy = clean + sigma[:, None, None, None] * torch.randn_like(clean)
weight = (sigma * sigma + 0.25) / (sigma * 0.5) ** 2
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=device.type == "cuda"):
pred = model(noisy, sigma)
loss = (weight[:, None, None, None] * (pred - clean).square()).mean()
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
with torch.no_grad():
decay = 0.995
for ep, p in zip(ema.parameters(), model.parameters()):
ep.lerp_(p, 1.0 - decay)
losses.append(float(loss.detach()))
if (step + 1) % 200 == 0 or step == 0:
print(json.dumps({
"event": "train", "arch": arch, "seed": model_seed,
"step": step + 1, "steps": steps,
"loss_mean_100": float(np.mean(losses[-100:])),
"elapsed_s": time.perf_counter() - started,
}), flush=True)
state = {k: v.detach().cpu() for k, v in ema.state_dict().items()}
del model, ema, optimizer, scaler
torch.cuda.empty_cache()
return state
@torch.no_grad()
def sample_model(state: dict[str, torch.Tensor], arch: str, initial: torch.Tensor,
steps: int, device: torch.device, sigma_max: float = 3.0,
sigma_min: float = 0.01) -> torch.Tensor:
model = make_model(arch).to(device).eval()
model.load_state_dict(state)
rho = 5.0
ramp = torch.linspace(0, 1, steps, device=device)
sigmas = (sigma_max ** (1 / rho) + ramp * (sigma_min ** (1 / rho) - sigma_max ** (1 / rho))) ** rho
sigmas = torch.cat([sigmas, torch.zeros(1, device=device)])
x = initial.to(device) * sigma_max
for i in range(steps):
s = sigmas[i]
sn = sigmas[i + 1]
sb = torch.full((x.shape[0],), s, device=device)
den = model(x, sb)
derivative = (x - den) / s
proposal = x + (sn - s) * derivative
if sn > 0:
den_next = model(proposal, torch.full((x.shape[0],), sn, device=device))
derivative_next = (proposal - den_next) / sn
x = x + (sn - s) * 0.5 * (derivative + derivative_next)
else:
x = proposal
output = x.clamp(-1.0, 1.0).cpu()
del model
torch.cuda.empty_cache()
return output
@torch.no_grad()
def denoise_model(state: dict[str, torch.Tensor], arch: str, noisy: torch.Tensor,
sigma: float, device: torch.device) -> torch.Tensor:
model = make_model(arch).to(device).eval()
model.load_state_dict(state)
out = model(noisy.to(device), torch.full((noisy.shape[0],), sigma, device=device)).cpu()
del model
torch.cuda.empty_cache()
return out
def rankdata(values: np.ndarray) -> np.ndarray:
order = np.argsort(values, kind="mergesort")
ranks = np.empty_like(order, dtype=np.float64)
ranks[order] = np.arange(values.size, dtype=np.float64)
return ranks
def spearman(x: np.ndarray, y: np.ndarray) -> float:
return float(np.corrcoef(rankdata(x), rankdata(y))[0, 1])
def solve_kappa(z: float, eig: np.ndarray, gamma: float) -> float:
def residual(k: float) -> float:
return k - z - gamma * k * np.mean(eig / (eig + k))
lo = max(z, 1e-10)
hi = z + gamma * float(np.mean(eig)) + 10.0
for _ in range(100):
mid = 0.5 * (lo + hi)
if residual(mid) > 0:
hi = mid
else:
lo = mid
return 0.5 * (lo + hi)
def covariance_stats(images: torch.Tensor) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
x = images.flatten(1).numpy().astype(np.float64)
mean = x.mean(0)
centered = x - mean
cov = centered.T @ centered / x.shape[0]
eig, vec = np.linalg.eigh(cov)
eig = np.clip(eig, 1e-10, None)
return mean, eig, vec
def linear_sample(mean: np.ndarray, eig: np.ndarray, vec: np.ndarray,
noise: np.ndarray, sigma_max: float = 3.0) -> np.ndarray:
scale = np.sqrt(eig / (eig + sigma_max * sigma_max))
projected = (noise * sigma_max) @ vec
return mean + projected @ (vec * scale[None, :]).T
def nearest_mse(generated: np.ndarray, train: np.ndarray, chunk: int = 1000) -> float:
values = []
for sample in generated:
best = float("inf")
for start in range(0, train.shape[0], chunk):
diff = train[start:start + chunk] - sample
best = min(best, float(np.mean(diff * diff, axis=1).min()))
values.append(best)
return float(np.mean(values))
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--small-n", type=int, default=1500)
parser.add_argument("--large-n", type=int, default=6000)
parser.add_argument("--steps-small", type=int, default=700)
parser.add_argument("--steps-large", type=int, default=1000)
parser.add_argument("--batch", type=int, default=128)
parser.add_argument("--sample-count", type=int, default=64)
parser.add_argument("--sample-steps", type=int, default=18)
parser.add_argument("--seed", type=int, default=2449)
parser.add_argument("--resume", action="store_true", help="Reuse checkpoints already present under --output")
args = parser.parse_args()
args.output.mkdir(parents=True, exist_ok=True)
seed_all(args.seed)
torch.set_float32_matmul_precision("high")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(json.dumps({
"event": "start", "device": str(device),
"gpu": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
"torch": torch.__version__, "config": vars(args) | {"output": str(args.output)},
}), flush=True)
transform = Compose([ToTensor(), Normalize((0.5,), (0.5,))])
dataset = FashionMNIST(root=str(args.output / "data"), train=True, download=True, transform=transform)
all_images = torch.stack([dataset[i][0] for i in range(len(dataset))])
permutation = torch.randperm(len(dataset), generator=torch.Generator().manual_seed(args.seed))
pools = {1: permutation[:30000], 2: permutation[30000:60000]}
sizes = {"small": args.small_n, "large": args.large_n}
steps_by_size = {"small": args.steps_small, "large": args.steps_large}
states: dict[str, dict[str, torch.Tensor]] = {}
train_seconds = 0.0
for arch in ["unet", "dit"]:
for size_name, n in sizes.items():
for split in [1, 2]:
key = f"{arch}_{size_name}_split{split}"
subset = all_images[pools[split][:n]]
checkpoint = args.output / f"{key}.pt"
if args.resume and checkpoint.exists():
states[key] = torch.load(checkpoint, map_location="cpu", weights_only=True)
print(json.dumps({"event": "checkpoint_reused", "key": key, "path": str(checkpoint)}), flush=True)
else:
started = time.perf_counter()
states[key] = train_model(
arch, subset, steps_by_size[size_name], args.batch,
args.seed + 1000 * (arch == "dit") + 100 * (size_name == "large") + split,
device,
)
train_seconds += time.perf_counter() - started
torch.save(states[key], checkpoint)
print(json.dumps({"event": "checkpoint", "key": key, "path": str(checkpoint)}), flush=True)
if args.resume and (args.output / "metrics.json").exists():
previous_metrics = json.loads((args.output / "metrics.json").read_text(encoding="utf-8"))
train_seconds = float(previous_metrics.get("train_seconds", train_seconds))
initial = torch.randn(args.sample_count, 1, 28, 28, generator=torch.Generator().manual_seed(args.seed + 77))
generated: dict[str, torch.Tensor] = {}
for key, state in states.items():
arch = key.split("_")[0]
generated[key] = sample_model(state, arch, initial, args.sample_steps, device)
torch.save(generated[key], args.output / f"samples_{key}.pt")
print(json.dumps({"event": "sampled", "key": key}), flush=True)
# Population eigensystem is estimated from the two disjoint large pools.
combined_large = torch.cat([all_images[pools[1][:args.large_n]], all_images[pools[2][:args.large_n]]], dim=0)
pop_mean, pop_eig, pop_vec = covariance_stats(combined_large)
noise_flat = initial.flatten(1).numpy().astype(np.float64)
linear_outputs = {}
for split in [1, 2]:
train_tensor = all_images[pools[split][:args.large_n]]
mean, eig, vec = covariance_stats(train_tensor)
linear_outputs[split] = linear_sample(mean, eig, vec, noise_flat)
sigma_probe = 0.5
heldout_idx = permutation[2 * args.large_n:2 * args.large_n + 96]
probe_clean = all_images[heldout_idx]
probe_noise = torch.randn(probe_clean.shape, generator=torch.Generator().manual_seed(args.seed + 88))
probe_noisy = probe_clean + sigma_probe * probe_noise
rows = []
gain_rows = []
metrics = {
"device": str(device),
"gpu": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
"small_n": args.small_n,
"large_n": args.large_n,
"train_seconds": train_seconds,
"sample_count": args.sample_count,
"sample_steps": args.sample_steps,
"linear_large_cross_split_mse": float(np.mean((linear_outputs[1] - linear_outputs[2]) ** 2)),
}
pop_centered_initial = noise_flat @ pop_vec
for arch in ["unet", "dit"]:
for size_name, n in sizes.items():
a = generated[f"{arch}_{size_name}_split1"].flatten(1).numpy()
b = generated[f"{arch}_{size_name}_split2"].flatten(1).numpy()
unrelated = np.roll(b, 1, axis=0)
same_seed_mse = np.mean((a - b) ** 2, axis=1)
unrelated_mse = np.mean((a - unrelated) ** 2, axis=1)
train1 = all_images[pools[1][:n]].flatten(1).numpy()
train2 = all_images[pools[2][:n]].flatten(1).numpy()
nn_mse = 0.5 * (nearest_mse(a, train1) + nearest_mse(b, train2))
den1 = denoise_model(states[f"{arch}_{size_name}_split1"], arch, probe_noisy, sigma_probe, device)
den2 = denoise_model(states[f"{arch}_{size_name}_split2"], arch, probe_noisy, sigma_probe, device)
delta = (den1 - den2).flatten(1).numpy().astype(np.float64)
projected_delta = delta @ pop_vec
mode_mse = np.mean(projected_delta ** 2, axis=0)
kappa = solve_kappa(sigma_probe ** 2, pop_eig, gamma=pop_eig.size / n)
chi = pop_eig / (pop_eig + kappa) ** 2
mode_r = spearman(mode_mse, chi)
input_pc = (probe_noisy.flatten(1).numpy().astype(np.float64) - pop_mean) @ pop_vec
mean_denoised = 0.5 * (den1 + den2)
output_pc = (mean_denoised.flatten(1).numpy().astype(np.float64) - pop_mean) @ pop_vec
net_gain = np.sum(input_pc * output_pc, axis=0) / np.maximum(np.sum(input_pc ** 2, axis=0), 1e-12)
population_gain = pop_eig / (pop_eig + sigma_probe ** 2)
rmt_gain = pop_eig / (pop_eig + kappa)
lower_band = slice(pop_eig.size // 5, pop_eig.size // 2)
upper_band = slice(3 * pop_eig.size // 4, pop_eig.size)
location_weight = pop_eig / (pop_eig + kappa) ** 2
location_predictor = np.sum(location_weight[None, :] * pop_centered_initial ** 2, axis=1)
location_r = spearman(location_predictor, same_seed_mse)
centered_a = (a - pop_mean) @ pop_vec
centered_b = (b - pop_mean) @ pop_vec
gen_var = 0.5 * (centered_a.var(0) + centered_b.var(0))
gain = np.sqrt(np.maximum(gen_var, 1e-12) / pop_eig)
quart = pop_eig.size // 4
low_gain = float(np.median(gain[:quart]))
high_gain = float(np.median(gain[-quart:]))
prefix = f"{arch}_{size_name}"
metrics[prefix] = {
"cross_split_mse": float(same_seed_mse.mean()),
"unrelated_seed_mse": float(unrelated_mse.mean()),
"nearest_train_mse": nn_mse,
"same_vs_unrelated_ratio": float(same_seed_mse.mean() / unrelated_mse.mean()),
"same_vs_nearest_ratio": float(same_seed_mse.mean() / nn_mse),
"eigenmode_spearman_r": mode_r,
"location_spearman_r": location_r,
"low_eigenband_gain": low_gain,
"high_eigenband_gain": high_gain,
"kappa": kappa,
"denoiser_gain_rmse_to_rmt": float(np.sqrt(np.mean((net_gain - rmt_gain) ** 2))),
"denoiser_gain_rmse_to_population": float(np.sqrt(np.mean((net_gain - population_gain) ** 2))),
"lower_band_network_gain": float(np.median(net_gain[lower_band])),
"lower_band_rmt_gain": float(np.median(rmt_gain[lower_band])),
"lower_band_population_gain": float(np.median(population_gain[lower_band])),
"lower_band_overshrink_vs_population": float(np.median(population_gain[lower_band] - net_gain[lower_band])),
"upper_band_network_gain": float(np.median(net_gain[upper_band])),
}
for seed_idx in range(args.sample_count):
rows.append({
"kind": "seed", "arch": arch, "size": size_name, "n": n,
"index": seed_idx, "x": float(location_predictor[seed_idx]),
"y": float(same_seed_mse[seed_idx]),
})
for mode_idx in range(pop_eig.size):
rows.append({
"kind": "mode", "arch": arch, "size": size_name, "n": n,
"index": mode_idx, "x": float(chi[mode_idx]), "y": float(mode_mse[mode_idx]),
})
gain_rows.append({
"arch": arch, "size": size_name, "n": n, "index": mode_idx,
"eigenvalue": float(pop_eig[mode_idx]), "network_gain": float(net_gain[mode_idx]),
"rmt_gain": float(rmt_gain[mode_idx]), "population_gain": float(population_gain[mode_idx]),
})
metrics[f"{arch}_consistency_decay_small_to_large"] = (
metrics[f"{arch}_small"]["cross_split_mse"] / metrics[f"{arch}_large"]["cross_split_mse"]
)
metrics[f"{arch}_low_band_gain_small_over_large"] = (
metrics[f"{arch}_small"]["low_eigenband_gain"] / metrics[f"{arch}_large"]["low_eigenband_gain"]
)
metrics[f"{arch}_high_band_gain_small_over_large"] = (
metrics[f"{arch}_small"]["high_eigenband_gain"] / metrics[f"{arch}_large"]["high_eigenband_gain"]
)
with (args.output / "raw.csv").open("w", newline="", encoding="utf-8") as handle:
writer = csv.DictWriter(handle, fieldnames=["kind", "arch", "size", "n", "index", "x", "y"])
writer.writeheader()
writer.writerows(rows)
with (args.output / "gains.csv").open("w", newline="", encoding="utf-8") as handle:
writer = csv.DictWriter(handle, fieldnames=[
"arch", "size", "n", "index", "eigenvalue", "network_gain", "rmt_gain", "population_gain",
])
writer.writeheader()
writer.writerows(gain_rows)
(args.output / "metrics.json").write_text(json.dumps(metrics, indent=2, sort_keys=True), encoding="utf-8")
grid_rows = []
labels = []
for arch in ["unet", "dit"]:
for size_name in ["small", "large"]:
for split in [1, 2]:
key = f"{arch}_{size_name}_split{split}"
grid_rows.append(generated[key][:8])
labels.append(key)
image_grid = make_grid(torch.cat(grid_rows), nrow=8, normalize=True, value_range=(-1, 1), padding=2)
save_image(image_grid, args.output / "sample_grid.png")
(args.output / "grid_rows.json").write_text(json.dumps(labels, indent=2), encoding="utf-8")
print("FINAL_METRICS=" + json.dumps(metrics, sort_keys=True), flush=True)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
22.1 kB
·
Xet hash:
300d7770759948241bb5768e7d054b976ec33a7c0d02fa11fc52dc12879d1b25

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.