Spaces:
Running
Running
| 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 ConditionalVAE, TinyVisionJudge, parameter_count | |
| from PIL import Image | |
| from safetensors.torch import load_file, save_file | |
| from torch.nn import functional as F | |
| from torch.utils.data import DataLoader, TensorDataset | |
| PROJECT_DIR = Path(__file__).resolve().parent | |
| ROOT_DIR = PROJECT_DIR.parents[1] | |
| VISION_DIR = ROOT_DIR / "projects" / "tiny-vision-foundry" | |
| DATA_DIR = VISION_DIR / "data" | |
| JUDGE_WEIGHTS = VISION_DIR / "artifacts" / "tiny-student-scratch" / "model.safetensors" | |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "glyph-forge-cvae" | |
| def seed_everything(seed: int) -> None: | |
| random.seed(seed) | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| def load_split(name: str, *, shuffle: bool, batch_size: int) -> DataLoader: | |
| frame = pd.read_parquet(DATA_DIR / f"{name}.parquet") | |
| pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16.0 | |
| labels = frame["label"].to_numpy(dtype=np.int64, copy=True) | |
| return DataLoader( | |
| TensorDataset(torch.from_numpy(pixels), torch.from_numpy(labels)), | |
| batch_size=batch_size, | |
| shuffle=shuffle, | |
| generator=torch.Generator().manual_seed(2031), | |
| ) | |
| def losses( | |
| reconstruction: torch.Tensor, | |
| pixels: torch.Tensor, | |
| mean: torch.Tensor, | |
| log_variance: torch.Tensor, | |
| beta: float, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: | |
| reconstruction_loss = F.binary_cross_entropy(reconstruction, pixels) | |
| kl = -0.5 * torch.mean(1 + log_variance - mean.square() - log_variance.exp()) | |
| return reconstruction_loss + beta * kl, reconstruction_loss, kl | |
| def evaluate_reconstruction(model: ConditionalVAE, loader: DataLoader) -> dict: | |
| model.eval() | |
| squared_error = 0.0 | |
| kl_total = 0.0 | |
| examples = 0 | |
| for pixels, labels in loader: | |
| reconstruction, mean, log_variance = model(pixels, labels) | |
| squared_error += F.mse_loss( | |
| reconstruction, | |
| pixels, | |
| reduction="sum", | |
| ).item() | |
| kl = -0.5 * torch.mean( | |
| 1 + log_variance - mean.square() - log_variance.exp(), | |
| dim=1, | |
| ) | |
| kl_total += kl.sum().item() | |
| examples += len(labels) | |
| return { | |
| "reconstruction_mse": squared_error / (examples * 64), | |
| "mean_kl": kl_total / examples, | |
| "examples": examples, | |
| } | |
| def generation_metrics( | |
| model: ConditionalVAE, | |
| judge: TinyVisionJudge, | |
| samples_per_class: int = 100, | |
| ) -> tuple[dict, torch.Tensor, torch.Tensor]: | |
| model.eval() | |
| judge.eval() | |
| labels = torch.arange(10).repeat_interleave(samples_per_class) | |
| latent = torch.randn(len(labels), model.latent_dimensions) | |
| generated = model.decode(latent, labels) | |
| predictions = judge(generated.reshape(-1, 1, 8, 8)).argmax(dim=1) | |
| per_class = { | |
| str(label): float( | |
| (predictions[labels == label] == labels[labels == label]).float().mean() | |
| ) | |
| for label in range(10) | |
| } | |
| diversity = { | |
| str(label): float(generated[labels == label].var(dim=0).mean()) | |
| for label in range(10) | |
| } | |
| return ( | |
| { | |
| "judge_accuracy": float((predictions == labels).float().mean()), | |
| "judge_accuracy_by_class": per_class, | |
| "mean_pixel_variance_by_class": diversity, | |
| "samples": len(labels), | |
| }, | |
| generated, | |
| labels, | |
| ) | |
| def save_grid(generated: torch.Tensor, labels: torch.Tensor, path: Path) -> None: | |
| selected = [] | |
| for label in range(10): | |
| selected.append(generated[labels == label][:10]) | |
| images = torch.cat(selected).reshape(10, 10, 8, 8).cpu().numpy() | |
| canvas = np.zeros((10 * 8, 10 * 8), dtype=np.uint8) | |
| for row in range(10): | |
| for column in range(10): | |
| canvas[ | |
| row * 8 : (row + 1) * 8, | |
| column * 8 : (column + 1) * 8, | |
| ] = np.clip(images[row, column] * 255, 0, 255).astype(np.uint8) | |
| Image.fromarray(canvas, mode="L").resize((800, 800), Image.Resampling.NEAREST).save( | |
| path | |
| ) | |
| def main() -> None: | |
| seed_everything(2031) | |
| if not JUDGE_WEIGHTS.exists(): | |
| raise FileNotFoundError("Train Tiny Vision Foundry before GlyphForge.") | |
| train_loader = load_split("train", shuffle=True, batch_size=96) | |
| validation_loader = load_split("validation", shuffle=False, batch_size=256) | |
| test_loader = load_split("test", shuffle=False, batch_size=256) | |
| model = ConditionalVAE() | |
| judge = TinyVisionJudge() | |
| judge.load_state_dict(load_file(JUDGE_WEIGHTS)) | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=0.002, weight_decay=0.001) | |
| epochs = 160 | |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) | |
| best_validation_mse = float("inf") | |
| best_epoch = 0 | |
| best_state = None | |
| trackio.init( | |
| project="glyph-forge-cvae", | |
| name="conditional-vae-8d-v1", | |
| config={ | |
| "parameters": parameter_count(model), | |
| "latent_dimensions": model.latent_dimensions, | |
| "epochs": epochs, | |
| "beta": 0.04, | |
| }, | |
| ) | |
| for epoch in range(1, epochs + 1): | |
| model.train() | |
| running_total = 0.0 | |
| running_reconstruction = 0.0 | |
| running_kl = 0.0 | |
| examples = 0 | |
| for pixels, labels in train_loader: | |
| reconstruction, mean, log_variance = model(pixels, labels) | |
| total, reconstruction_loss, kl = losses( | |
| reconstruction, | |
| pixels, | |
| mean, | |
| log_variance, | |
| beta=0.04, | |
| ) | |
| optimizer.zero_grad(set_to_none=True) | |
| total.backward() | |
| optimizer.step() | |
| running_total += total.item() * len(labels) | |
| running_reconstruction += reconstruction_loss.item() * len(labels) | |
| running_kl += kl.item() * len(labels) | |
| examples += len(labels) | |
| scheduler.step() | |
| validation = evaluate_reconstruction(model, validation_loader) | |
| if validation["reconstruction_mse"] < best_validation_mse: | |
| best_validation_mse = validation["reconstruction_mse"] | |
| best_epoch = epoch | |
| best_state = { | |
| key: value.detach().cpu().clone() | |
| for key, value in model.state_dict().items() | |
| } | |
| if epoch == 1 or epoch % 10 == 0: | |
| trackio.log( | |
| { | |
| "epoch": epoch, | |
| "train_loss": running_total / examples, | |
| "train_reconstruction_bce": running_reconstruction / examples, | |
| "train_kl": running_kl / examples, | |
| "validation_reconstruction_mse": validation["reconstruction_mse"], | |
| "validation_kl": validation["mean_kl"], | |
| "learning_rate": scheduler.get_last_lr()[0], | |
| } | |
| ) | |
| trackio.finish() | |
| assert best_state is not None | |
| model.load_state_dict(best_state) | |
| reconstruction = evaluate_reconstruction(model, test_loader) | |
| generation, generated, labels = generation_metrics(model, judge) | |
| results = { | |
| "model": "GlyphForge Conditional VAE", | |
| "parameters": parameter_count(model), | |
| "latent_dimensions": model.latent_dimensions, | |
| "best_epoch": best_epoch, | |
| "test": reconstruction, | |
| "generation": generation, | |
| "judge": "Tiny Vision labels-only student, 98.52% real-image test accuracy", | |
| } | |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) | |
| save_file(model.state_dict(), ARTIFACT_DIR / "model.safetensors") | |
| save_grid(generated, labels, ARTIFACT_DIR / "samples.png") | |
| (ARTIFACT_DIR / "evaluation.json").write_text( | |
| json.dumps(results, indent=2), | |
| encoding="utf-8", | |
| ) | |
| print(json.dumps(results, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |