energy-pocket / source /train.py
ARotting's picture
Publish Conditional energy model with persistent contrastive divergence
7d90be6 verified
Raw
History Blame Contribute Delete
8.4 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 (
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()