from __future__ import annotations import json import random from pathlib import Path import numpy as np import pandas as pd import torch import trackio from model import ( ConditionalGenerator, ProjectionCritic, TinyVisionJudge, parameter_count, ) from PIL import Image from safetensors.torch import load_file, save_file from torch.nn import functional as F PROJECT_DIR = Path(__file__).resolve().parent ROOT_DIR = PROJECT_DIR.parents[1] VISION_DIR = ROOT_DIR / "projects" / "tiny-vision-foundry" VISION_DATA = VISION_DIR / "data" / "train.parquet" JUDGE_WEIGHTS = ( VISION_DIR / "artifacts" / "tiny-student-scratch" / "model.safetensors" ) ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "pocket-wgan" DATA_DIR = PROJECT_DIR / "data" SEED = 2047 def seed_everything(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) def load_training_data() -> tuple[torch.Tensor, torch.Tensor]: frame = pd.read_parquet(VISION_DATA) pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16.0 labels = frame["label"].to_numpy(dtype=np.int64, copy=True) return torch.from_numpy(pixels), torch.from_numpy(labels) def gradient_penalty( critic: ProjectionCritic, real: torch.Tensor, fake: torch.Tensor, labels: torch.Tensor, ) -> torch.Tensor: alpha = torch.rand(len(real), 1) interpolated = alpha * real + (1 - alpha) * fake interpolated.requires_grad_(True) score, _ = critic(interpolated, labels) gradient = torch.autograd.grad( outputs=score.sum(), inputs=interpolated, create_graph=True, )[0] norm = gradient.flatten(1).norm(2, dim=1) return (norm - 1).square().mean() @torch.inference_mode() def quick_selection_score( generator: ConditionalGenerator, judge: TinyVisionJudge, real: torch.Tensor, real_labels: torch.Tensor, seed: int, ) -> dict: generator.eval() labels = torch.arange(10).repeat_interleave(40) generated = generator.generate(labels, seed=seed) predictions = judge(generated.reshape(-1, 1, 8, 8)).argmax(dim=1) fidelity = float((predictions == labels).float().mean()) ratios = [] for label in range(10): generated_variance = generated[labels == label].var(dim=0).mean() real_variance = real[real_labels == label].var(dim=0).mean() ratios.append(float(generated_variance / real_variance.clamp_min(1e-8))) diversity_ratio = float(np.mean(ratios)) return { "judge_fidelity": fidelity, "mean_diversity_ratio": diversity_ratio, "selection_score": fidelity + 0.15 * min(diversity_ratio, 1.0), } def nearest_neighbor_metrics( generated: torch.Tensor, labels: torch.Tensor, real: torch.Tensor, real_labels: torch.Tensor, ) -> dict: minimum_mse = [] unique_fractions = {} for label in range(10): fake_class = generated[labels == label] real_class = real[real_labels == label] distances = (fake_class[:, None, :] - real_class[None, :, :]).square().mean(2) minimum_mse.extend(distances.min(dim=1).values.tolist()) quantized = torch.round(fake_class * 16).to(torch.uint8).numpy() unique_fractions[str(label)] = float( len({sample.tobytes() for sample in quantized}) / len(quantized) ) minimum = np.asarray(minimum_mse) return { "mean_nearest_training_mse": float(minimum.mean()), "median_nearest_training_mse": float(np.median(minimum)), "exact_training_copy_fraction": float((minimum < 1e-8).mean()), "quantized_unique_fraction_by_class": unique_fractions, "mean_quantized_unique_fraction": float(np.mean(list(unique_fractions.values()))), } @torch.inference_mode() def evaluate( generator: ConditionalGenerator, critic: ProjectionCritic, judge: TinyVisionJudge, real: torch.Tensor, real_labels: torch.Tensor, ) -> tuple[dict, torch.Tensor, torch.Tensor, torch.Tensor]: generator.eval() critic.eval() judge.eval() labels = torch.arange(10).repeat_interleave(100) generated = generator.generate(labels, seed=SEED + 10_000) predictions = judge(generated.reshape(-1, 1, 8, 8)).argmax(dim=1) per_class_fidelity = {} per_class_variance = {} per_class_real_variance = {} per_class_diversity_ratio = {} for label in range(10): mask = labels == label real_mask = real_labels == label generated_variance = float(generated[mask].var(dim=0).mean()) real_variance = float(real[real_mask].var(dim=0).mean()) per_class_fidelity[str(label)] = float( (predictions[mask] == labels[mask]).float().mean() ) per_class_variance[str(label)] = generated_variance per_class_real_variance[str(label)] = real_variance per_class_diversity_ratio[str(label)] = generated_variance / real_variance memorization = nearest_neighbor_metrics( generated, labels, real, real_labels, ) report = { "judge_accuracy": float((predictions == labels).float().mean()), "judge_accuracy_by_class": per_class_fidelity, "mean_pixel_variance_by_class": per_class_variance, "real_mean_pixel_variance_by_class": per_class_real_variance, "diversity_ratio_by_class": per_class_diversity_ratio, "mean_diversity_ratio": float(np.mean(list(per_class_diversity_ratio.values()))), "samples": len(labels), "collapse_and_memorization_checks": memorization, } return report, generated, labels, predictions def save_grid(generated: torch.Tensor, labels: torch.Tensor, path: Path) -> None: images = torch.cat( [generated[labels == label][:10] for label in range(10)] ).reshape(10, 10, 8, 8) canvas = np.zeros((80, 80), dtype=np.uint8) for row in range(10): for column in range(10): canvas[row * 8 : (row + 1) * 8, column * 8 : (column + 1) * 8] = ( images[row, column].mul(255).clamp(0, 255).to(torch.uint8).numpy() ) Image.fromarray(canvas, mode="L").resize((800, 800), Image.Resampling.NEAREST).save( path ) def main() -> None: seed_everything(SEED) torch.set_num_threads(1) if not JUDGE_WEIGHTS.exists(): raise FileNotFoundError("Train Tiny Vision Foundry before Pocket WGAN.") real, real_labels = load_training_data() judge = TinyVisionJudge() judge.load_state_dict(load_file(JUDGE_WEIGHTS)) judge.eval() for parameter in judge.parameters(): parameter.requires_grad_(False) generator = ConditionalGenerator() critic = ProjectionCritic() generator_optimizer = torch.optim.Adam( generator.parameters(), lr=1.5e-4, betas=(0.0, 0.9) ) critic_optimizer = torch.optim.Adam( critic.parameters(), lr=1.5e-4, betas=(0.0, 0.9) ) batch_size = 128 generator_steps = 2_400 critic_steps = 3 rng = np.random.default_rng(SEED) history = [] best_score = -float("inf") best_step = 0 best_state = None trackio.init( project="pocket-wgan", name="projection-wgan-gp-v1", config={ "generator_parameters": parameter_count(generator), "critic_parameters": parameter_count(critic), "generator_steps": generator_steps, "critic_steps_per_generator": critic_steps, "gradient_penalty": 10.0, "training_examples": len(real), }, ) for step in range(1, generator_steps + 1): generator.train() critic.train() critic_loss_value = 0.0 gradient_penalty_value = 0.0 for _ in range(critic_steps): indexes = torch.from_numpy( rng.choice(len(real), batch_size, replace=False) ) real_batch = real[indexes] label_batch = real_labels[indexes] noise = torch.randn(batch_size, generator.noise_dimensions) fake_batch = generator(noise, label_batch).detach() real_score, real_logits = critic(real_batch, label_batch) fake_score, fake_logits = critic(fake_batch, label_batch) penalty = gradient_penalty(critic, real_batch, fake_batch, label_batch) auxiliary = F.cross_entropy(real_logits, label_batch) auxiliary = auxiliary + 0.25 * F.cross_entropy(fake_logits, label_batch) critic_loss = ( fake_score.mean() - real_score.mean() + 10.0 * penalty + 0.35 * auxiliary ) critic_optimizer.zero_grad(set_to_none=True) critic_loss.backward() critic_optimizer.step() critic_loss_value = float(critic_loss.detach()) gradient_penalty_value = float(penalty.detach()) labels = torch.from_numpy(rng.integers(0, 10, size=batch_size)).long() noise = torch.randn(batch_size, generator.noise_dimensions) generated = generator(noise, labels) score, logits = critic(generated, labels) generator_loss = -score.mean() + 0.75 * F.cross_entropy(logits, labels) generator_optimizer.zero_grad(set_to_none=True) generator_loss.backward() generator_optimizer.step() if step == 1 or step % 200 == 0: selection = quick_selection_score( generator, judge, real, real_labels, seed=SEED + step, ) record = { "generator_step": step, "generator_loss": float(generator_loss.detach()), "critic_loss": critic_loss_value, "gradient_penalty": gradient_penalty_value, **selection, } history.append(record) trackio.log(record) if selection["selection_score"] > best_score: best_score = selection["selection_score"] best_step = step best_state = { name: value.detach().cpu().clone() for name, value in generator.state_dict().items() } assert best_state is not None generator.load_state_dict(best_state) generation, generated, labels, predictions = evaluate( generator, critic, judge, real, real_labels, ) results = { "model": "Pocket WGAN-GP", "method": "Projection-conditioned WGAN-GP with auxiliary class supervision", "generator_parameters": parameter_count(generator), "critic_parameters": parameter_count(critic), "training_examples": len(real), "generator_steps": generator_steps, "critic_updates": generator_steps * critic_steps, "best_generator_step": best_step, "generation": generation, "judge": "Frozen Tiny Vision student, 98.52% real-image test accuracy", "training_history": history, } ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) DATA_DIR.mkdir(parents=True, exist_ok=True) save_file(generator.state_dict(), ARTIFACT_DIR / "generator.safetensors") save_file(critic.state_dict(), ARTIFACT_DIR / "critic.safetensors") save_grid(generated, labels, ARTIFACT_DIR / "samples.png") np.savez_compressed( ARTIFACT_DIR / "generated_samples.npz", pixels=generated.numpy(), labels=labels.numpy(), judge_predictions=predictions.numpy(), ) (ARTIFACT_DIR / "evaluation.json").write_text( json.dumps(results, indent=2), encoding="utf-8", ) pd.DataFrame( { "label": labels.numpy(), "judge_prediction": predictions.numpy(), "pixels": list(generated.numpy()), } ).to_parquet(DATA_DIR / "evaluation_samples.parquet", index=False) trackio.log( { "final_judge_accuracy": generation["judge_accuracy"], "final_diversity_ratio": generation["mean_diversity_ratio"], "final_unique_fraction": generation[ "collapse_and_memorization_checks" ]["mean_quantized_unique_fraction"], } ) trackio.finish() print(json.dumps(results, indent=2)) if __name__ == "__main__": main()