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()