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 @torch.inference_mode() 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, } @torch.inference_mode() 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()