| 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 ( |
| ConditionalEnergyNetwork, |
| TinyVisionJudge, |
| langevin_sample, |
| 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" |
| JUDGE_WEIGHTS = ( |
| VISION_DIR / "artifacts" / "tiny-student-scratch" / "model.safetensors" |
| ) |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "energy-pocket" |
| DATA_DIR = PROJECT_DIR / "data" |
| SEED = 2069 |
|
|
|
|
| def load_data() -> tuple[torch.Tensor, torch.Tensor]: |
| frame = pd.read_parquet(VISION_DIR / "data" / "train.parquet") |
| pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16 |
| labels = frame["label"].to_numpy(dtype=np.int64, copy=True) |
| return torch.from_numpy(pixels), torch.from_numpy(labels) |
|
|
|
|
| def nearest_metrics( |
| generated: torch.Tensor, |
| labels: torch.Tensor, |
| real: torch.Tensor, |
| real_labels: torch.Tensor, |
| ) -> dict: |
| nearest = [] |
| for label in range(10): |
| fake = generated[labels == label] |
| reference = real[real_labels == label] |
| distances = (fake[:, None] - reference[None]).square().mean(2) |
| nearest.extend(distances.min(1).values.tolist()) |
| values = np.asarray(nearest) |
| return { |
| "mean_nearest_training_mse": float(values.mean()), |
| "exact_training_copy_fraction": float((values < 1e-8).mean()), |
| } |
|
|
|
|
| def evaluate( |
| model: ConditionalEnergyNetwork, |
| judge: TinyVisionJudge, |
| real: torch.Tensor, |
| real_labels: torch.Tensor, |
| ) -> tuple[dict, torch.Tensor, torch.Tensor, torch.Tensor]: |
| labels = torch.arange(10).repeat_interleave(100) |
| generator = torch.Generator().manual_seed(SEED + 10_000) |
| initial = torch.rand(len(labels), 64, generator=generator) |
| generated = langevin_sample( |
| model, |
| initial, |
| labels, |
| steps=40, |
| step_size=0.08, |
| noise_scale=0.008, |
| generator=generator, |
| ) |
| with torch.inference_mode(): |
| predictions = judge(generated.reshape(-1, 1, 8, 8)).argmax(1) |
| positive_energy = float(model(real, real_labels).mean()) |
| negative_energy = float(model(generated, labels).mean()) |
| 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(0).mean()) |
| for label in range(10) |
| } |
| rounded = torch.round(generated * 16).to(torch.uint8).numpy() |
| uniqueness = { |
| str(label): len( |
| {row.tobytes() for row in rounded[labels.numpy() == label]} |
| ) |
| / 100 |
| for label in range(10) |
| } |
| report = { |
| "judge_accuracy": float((predictions == labels).float().mean()), |
| "judge_accuracy_by_class": per_class, |
| "mean_pixel_variance_by_class": diversity, |
| "mean_quantized_unique_fraction": float(np.mean(list(uniqueness.values()))), |
| "saturated_pixel_fraction": float( |
| ((generated < 0.02) | (generated > 0.98)).float().mean() |
| ), |
| "positive_training_energy": positive_energy, |
| "generated_energy": negative_energy, |
| "energy_gap_generated_minus_real": negative_energy - positive_energy, |
| "memorization": nearest_metrics(generated, labels, real, real_labels), |
| "samples": len(labels), |
| "langevin_steps": 40, |
| } |
| 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: |
| random.seed(SEED) |
| np.random.seed(SEED) |
| torch.manual_seed(SEED) |
| torch.set_num_threads(1) |
| real, real_labels = load_data() |
| model = ConditionalEnergyNetwork() |
| judge = TinyVisionJudge() |
| judge.load_state_dict(load_file(JUDGE_WEIGHTS)) |
| judge.eval() |
| optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4, weight_decay=1e-5) |
| replay = torch.rand(2_000, 64) |
| replay_labels = torch.arange(10).repeat_interleave(200) |
| rng = np.random.default_rng(SEED) |
| history = [] |
| trackio.init( |
| project="energy-pocket", |
| name="persistent-contrastive-divergence-v1", |
| config={ |
| "parameters": parameter_count(model), |
| "steps": 2_500, |
| "langevin_steps_per_update": 20, |
| "replay_examples": len(replay), |
| }, |
| ) |
| for step in range(1, 2_501): |
| indexes = torch.from_numpy(rng.choice(len(real), 128, replace=False)) |
| positive = real[indexes] |
| labels = real_labels[indexes] |
| replay_indexes = torch.from_numpy(rng.integers(0, len(replay), size=128)) |
| negative = replay[replay_indexes].clone() |
| refresh = torch.from_numpy(rng.random(128) < 0.05) |
| negative[refresh] = torch.rand(int(refresh.sum()), 64) |
| negative = langevin_sample( |
| model, |
| negative, |
| labels, |
| steps=20, |
| step_size=0.08, |
| noise_scale=0.01, |
| ) |
| replay[replay_indexes] = negative |
| replay_labels[replay_indexes] = labels |
| positive_energy = model(positive, labels) |
| negative_energy = model(negative, labels) |
| classification = F.cross_entropy(-model.all_energies(positive), labels) |
| energy_regularizer = positive_energy.square().mean() |
| energy_regularizer = energy_regularizer + negative_energy.square().mean() |
| loss = ( |
| positive_energy.mean() |
| - negative_energy.mean() |
| + classification |
| + 0.001 * energy_regularizer |
| ) |
| optimizer.zero_grad(set_to_none=True) |
| loss.backward() |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 10) |
| optimizer.step() |
| if step == 1 or step % 250 == 0: |
| record = { |
| "training_step": step, |
| "loss": float(loss.detach()), |
| "positive_energy": float(positive_energy.mean().detach()), |
| "negative_energy": float(negative_energy.mean().detach()), |
| "classification_loss": float(classification.detach()), |
| } |
| history.append(record) |
| trackio.log(record) |
| generation, generated, labels, predictions = evaluate( |
| model, judge, real, real_labels |
| ) |
| report = { |
| "model": "Energy Pocket", |
| "method": "Class-conditional energy network with persistent contrastive divergence", |
| "parameters": parameter_count(model), |
| "training_steps": 2_500, |
| "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(model.state_dict(), ARTIFACT_DIR / "model.safetensors") |
| save_grid(generated, labels, ARTIFACT_DIR / "samples.png") |
| (ARTIFACT_DIR / "evaluation.json").write_text( |
| json.dumps(report, indent=2), encoding="utf-8" |
| ) |
| pd.DataFrame( |
| { |
| "label": labels.numpy(), |
| "judge_prediction": predictions.numpy(), |
| "pixels": list(generated.numpy()), |
| } |
| ).to_parquet(DATA_DIR / "langevin_samples.parquet", index=False) |
| trackio.log( |
| { |
| "judge_accuracy": generation["judge_accuracy"], |
| "energy_gap": generation["energy_gap_generated_minus_real"], |
| "unique_fraction": generation["mean_quantized_unique_fraction"], |
| } |
| ) |
| trackio.finish() |
| print(json.dumps(report, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|