| 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 PocketDenoiser, 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" / "pocket-diffusion" |
| STEPS = 50 |
|
|
|
|
| def seed_everything(seed: int) -> None: |
| random.seed(seed) |
| np.random.seed(seed) |
| torch.manual_seed(seed) |
|
|
|
|
| def load_training_data() -> DataLoader: |
| frame = pd.read_parquet(DATA_DIR / "train.parquet") |
| pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 8.0 - 1.0 |
| labels = frame["label"].to_numpy(dtype=np.int64, copy=True) |
| return DataLoader( |
| TensorDataset(torch.from_numpy(pixels), torch.from_numpy(labels)), |
| batch_size=128, |
| shuffle=True, |
| generator=torch.Generator().manual_seed(2032), |
| ) |
|
|
|
|
| def schedule() -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: |
| betas = torch.linspace(1e-4, 0.025, STEPS) |
| alphas = 1.0 - betas |
| cumulative = torch.cumprod(alphas, dim=0) |
| return betas, alphas, cumulative |
|
|
|
|
| @torch.inference_mode() |
| def sample( |
| model: PocketDenoiser, |
| labels: torch.Tensor, |
| guidance: float, |
| seed: int, |
| ) -> torch.Tensor: |
| generator = torch.Generator().manual_seed(seed) |
| betas, alphas, cumulative = schedule() |
| pixels = torch.randn(len(labels), 64, generator=generator) |
| null_labels = torch.full_like(labels, 10) |
| model.eval() |
| for step in reversed(range(STEPS)): |
| timesteps = torch.full((len(labels),), step, dtype=torch.long) |
| conditional = model(pixels, timesteps, labels) |
| unconditional = model(pixels, timesteps, null_labels) |
| predicted_noise = unconditional + guidance * (conditional - unconditional) |
| alpha = alphas[step] |
| cumulative_alpha = cumulative[step] |
| mean = ( |
| pixels - (1 - alpha) / torch.sqrt(1 - cumulative_alpha) * predicted_noise |
| ) / torch.sqrt(alpha) |
| if step: |
| noise = torch.randn(pixels.shape, generator=generator) |
| pixels = mean + torch.sqrt(betas[step]) * noise |
| else: |
| pixels = mean |
| return torch.clamp((pixels + 1) / 2, 0, 1) |
|
|
|
|
| @torch.inference_mode() |
| def generation_metrics( |
| model: PocketDenoiser, |
| judge: TinyVisionJudge, |
| guidance: float, |
| ) -> tuple[dict, torch.Tensor, torch.Tensor]: |
| labels = torch.arange(10).repeat_interleave(100) |
| generated = sample(model, labels, guidance=guidance, seed=2032) |
| predictions = judge(generated.reshape(-1, 1, 8, 8)).argmax(dim=1) |
| accuracy_by_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": accuracy_by_class, |
| "mean_pixel_variance_by_class": diversity, |
| "samples": len(labels), |
| "guidance": guidance, |
| }, |
| generated, |
| labels, |
| ) |
|
|
|
|
| def save_grid(generated: torch.Tensor, labels: torch.Tensor, path: Path) -> None: |
| selected = [generated[labels == label][:10] for label in range(10)] |
| images = torch.cat(selected).reshape(10, 10, 8, 8).cpu().numpy() |
| 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, |
| ] = 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(2032) |
| loader = load_training_data() |
| model = PocketDenoiser() |
| judge = TinyVisionJudge() |
| judge.load_state_dict(load_file(JUDGE_WEIGHTS)) |
| judge.eval() |
| betas, _, cumulative = schedule() |
| optimizer = torch.optim.AdamW(model.parameters(), lr=0.0015, weight_decay=0.001) |
| epochs = 300 |
| scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs) |
| trackio.init( |
| project="pocket-diffusion", |
| name="conditional-ddpm-cfg-v1", |
| config={ |
| "parameters": parameter_count(model), |
| "diffusion_steps": STEPS, |
| "epochs": epochs, |
| "label_dropout": 0.12, |
| }, |
| ) |
| for epoch in range(1, epochs + 1): |
| model.train() |
| running_loss = 0.0 |
| examples = 0 |
| for pixels, labels in loader: |
| timesteps = torch.randint(0, STEPS, (len(labels),)) |
| noise = torch.randn_like(pixels) |
| cumulative_alpha = cumulative[timesteps].unsqueeze(1) |
| noisy = ( |
| torch.sqrt(cumulative_alpha) * pixels |
| + torch.sqrt(1 - cumulative_alpha) * noise |
| ) |
| conditioned_labels = labels.clone() |
| drop = torch.rand(len(labels)) < 0.12 |
| conditioned_labels[drop] = 10 |
| prediction = model(noisy, timesteps, conditioned_labels) |
| loss = F.mse_loss(prediction, noise) |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| optimizer.step() |
| running_loss += loss.item() * len(labels) |
| examples += len(labels) |
| scheduler.step() |
| if epoch == 1 or epoch % 10 == 0: |
| trackio.log( |
| { |
| "epoch": epoch, |
| "noise_prediction_mse": running_loss / examples, |
| "learning_rate": scheduler.get_last_lr()[0], |
| } |
| ) |
| trackio.finish() |
|
|
| guidance_candidates = {} |
| for guidance in [1.0, 1.5, 2.0, 2.5, 3.0]: |
| metrics, _, _ = generation_metrics(model, judge, guidance) |
| guidance_candidates[str(guidance)] = metrics["judge_accuracy"] |
| best_guidance = float(max(guidance_candidates, key=guidance_candidates.get)) |
| generation, generated, labels = generation_metrics(model, judge, best_guidance) |
| results = { |
| "model": "PocketDiffusion", |
| "parameters": parameter_count(model), |
| "diffusion_steps": STEPS, |
| "epochs": epochs, |
| "guidance_search": guidance_candidates, |
| "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() |
|
|