Buckets:
| #!/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 | |
| 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 | |
| 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.