| 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() |
|
|
|
|