ARotting's picture
Publish 9.5K parameter conditional variational autoencoder
36072d3 verified
Raw
History Blame Contribute Delete
8.03 kB
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()