pocket-wgan / source /train.py
ARotting's picture
Publish Projection-conditioned WGAN-GP with collapse diagnostics
f3a1fe5 verified
Raw
History Blame Contribute Delete
12.3 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 (
ConditionalGenerator,
ProjectionCritic,
TinyVisionJudge,
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"
VISION_DATA = VISION_DIR / "data" / "train.parquet"
JUDGE_WEIGHTS = (
VISION_DIR / "artifacts" / "tiny-student-scratch" / "model.safetensors"
)
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "pocket-wgan"
DATA_DIR = PROJECT_DIR / "data"
SEED = 2047
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
def load_training_data() -> tuple[torch.Tensor, torch.Tensor]:
frame = pd.read_parquet(VISION_DATA)
pixels = np.stack(frame["image"].to_numpy()).astype(np.float32) / 16.0
labels = frame["label"].to_numpy(dtype=np.int64, copy=True)
return torch.from_numpy(pixels), torch.from_numpy(labels)
def gradient_penalty(
critic: ProjectionCritic,
real: torch.Tensor,
fake: torch.Tensor,
labels: torch.Tensor,
) -> torch.Tensor:
alpha = torch.rand(len(real), 1)
interpolated = alpha * real + (1 - alpha) * fake
interpolated.requires_grad_(True)
score, _ = critic(interpolated, labels)
gradient = torch.autograd.grad(
outputs=score.sum(),
inputs=interpolated,
create_graph=True,
)[0]
norm = gradient.flatten(1).norm(2, dim=1)
return (norm - 1).square().mean()
@torch.inference_mode()
def quick_selection_score(
generator: ConditionalGenerator,
judge: TinyVisionJudge,
real: torch.Tensor,
real_labels: torch.Tensor,
seed: int,
) -> dict:
generator.eval()
labels = torch.arange(10).repeat_interleave(40)
generated = generator.generate(labels, seed=seed)
predictions = judge(generated.reshape(-1, 1, 8, 8)).argmax(dim=1)
fidelity = float((predictions == labels).float().mean())
ratios = []
for label in range(10):
generated_variance = generated[labels == label].var(dim=0).mean()
real_variance = real[real_labels == label].var(dim=0).mean()
ratios.append(float(generated_variance / real_variance.clamp_min(1e-8)))
diversity_ratio = float(np.mean(ratios))
return {
"judge_fidelity": fidelity,
"mean_diversity_ratio": diversity_ratio,
"selection_score": fidelity + 0.15 * min(diversity_ratio, 1.0),
}
def nearest_neighbor_metrics(
generated: torch.Tensor,
labels: torch.Tensor,
real: torch.Tensor,
real_labels: torch.Tensor,
) -> dict:
minimum_mse = []
unique_fractions = {}
for label in range(10):
fake_class = generated[labels == label]
real_class = real[real_labels == label]
distances = (fake_class[:, None, :] - real_class[None, :, :]).square().mean(2)
minimum_mse.extend(distances.min(dim=1).values.tolist())
quantized = torch.round(fake_class * 16).to(torch.uint8).numpy()
unique_fractions[str(label)] = float(
len({sample.tobytes() for sample in quantized}) / len(quantized)
)
minimum = np.asarray(minimum_mse)
return {
"mean_nearest_training_mse": float(minimum.mean()),
"median_nearest_training_mse": float(np.median(minimum)),
"exact_training_copy_fraction": float((minimum < 1e-8).mean()),
"quantized_unique_fraction_by_class": unique_fractions,
"mean_quantized_unique_fraction": float(np.mean(list(unique_fractions.values()))),
}
@torch.inference_mode()
def evaluate(
generator: ConditionalGenerator,
critic: ProjectionCritic,
judge: TinyVisionJudge,
real: torch.Tensor,
real_labels: torch.Tensor,
) -> tuple[dict, torch.Tensor, torch.Tensor, torch.Tensor]:
generator.eval()
critic.eval()
judge.eval()
labels = torch.arange(10).repeat_interleave(100)
generated = generator.generate(labels, seed=SEED + 10_000)
predictions = judge(generated.reshape(-1, 1, 8, 8)).argmax(dim=1)
per_class_fidelity = {}
per_class_variance = {}
per_class_real_variance = {}
per_class_diversity_ratio = {}
for label in range(10):
mask = labels == label
real_mask = real_labels == label
generated_variance = float(generated[mask].var(dim=0).mean())
real_variance = float(real[real_mask].var(dim=0).mean())
per_class_fidelity[str(label)] = float(
(predictions[mask] == labels[mask]).float().mean()
)
per_class_variance[str(label)] = generated_variance
per_class_real_variance[str(label)] = real_variance
per_class_diversity_ratio[str(label)] = generated_variance / real_variance
memorization = nearest_neighbor_metrics(
generated,
labels,
real,
real_labels,
)
report = {
"judge_accuracy": float((predictions == labels).float().mean()),
"judge_accuracy_by_class": per_class_fidelity,
"mean_pixel_variance_by_class": per_class_variance,
"real_mean_pixel_variance_by_class": per_class_real_variance,
"diversity_ratio_by_class": per_class_diversity_ratio,
"mean_diversity_ratio": float(np.mean(list(per_class_diversity_ratio.values()))),
"samples": len(labels),
"collapse_and_memorization_checks": memorization,
}
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:
seed_everything(SEED)
torch.set_num_threads(1)
if not JUDGE_WEIGHTS.exists():
raise FileNotFoundError("Train Tiny Vision Foundry before Pocket WGAN.")
real, real_labels = load_training_data()
judge = TinyVisionJudge()
judge.load_state_dict(load_file(JUDGE_WEIGHTS))
judge.eval()
for parameter in judge.parameters():
parameter.requires_grad_(False)
generator = ConditionalGenerator()
critic = ProjectionCritic()
generator_optimizer = torch.optim.Adam(
generator.parameters(), lr=1.5e-4, betas=(0.0, 0.9)
)
critic_optimizer = torch.optim.Adam(
critic.parameters(), lr=1.5e-4, betas=(0.0, 0.9)
)
batch_size = 128
generator_steps = 2_400
critic_steps = 3
rng = np.random.default_rng(SEED)
history = []
best_score = -float("inf")
best_step = 0
best_state = None
trackio.init(
project="pocket-wgan",
name="projection-wgan-gp-v1",
config={
"generator_parameters": parameter_count(generator),
"critic_parameters": parameter_count(critic),
"generator_steps": generator_steps,
"critic_steps_per_generator": critic_steps,
"gradient_penalty": 10.0,
"training_examples": len(real),
},
)
for step in range(1, generator_steps + 1):
generator.train()
critic.train()
critic_loss_value = 0.0
gradient_penalty_value = 0.0
for _ in range(critic_steps):
indexes = torch.from_numpy(
rng.choice(len(real), batch_size, replace=False)
)
real_batch = real[indexes]
label_batch = real_labels[indexes]
noise = torch.randn(batch_size, generator.noise_dimensions)
fake_batch = generator(noise, label_batch).detach()
real_score, real_logits = critic(real_batch, label_batch)
fake_score, fake_logits = critic(fake_batch, label_batch)
penalty = gradient_penalty(critic, real_batch, fake_batch, label_batch)
auxiliary = F.cross_entropy(real_logits, label_batch)
auxiliary = auxiliary + 0.25 * F.cross_entropy(fake_logits, label_batch)
critic_loss = (
fake_score.mean() - real_score.mean() + 10.0 * penalty + 0.35 * auxiliary
)
critic_optimizer.zero_grad(set_to_none=True)
critic_loss.backward()
critic_optimizer.step()
critic_loss_value = float(critic_loss.detach())
gradient_penalty_value = float(penalty.detach())
labels = torch.from_numpy(rng.integers(0, 10, size=batch_size)).long()
noise = torch.randn(batch_size, generator.noise_dimensions)
generated = generator(noise, labels)
score, logits = critic(generated, labels)
generator_loss = -score.mean() + 0.75 * F.cross_entropy(logits, labels)
generator_optimizer.zero_grad(set_to_none=True)
generator_loss.backward()
generator_optimizer.step()
if step == 1 or step % 200 == 0:
selection = quick_selection_score(
generator,
judge,
real,
real_labels,
seed=SEED + step,
)
record = {
"generator_step": step,
"generator_loss": float(generator_loss.detach()),
"critic_loss": critic_loss_value,
"gradient_penalty": gradient_penalty_value,
**selection,
}
history.append(record)
trackio.log(record)
if selection["selection_score"] > best_score:
best_score = selection["selection_score"]
best_step = step
best_state = {
name: value.detach().cpu().clone()
for name, value in generator.state_dict().items()
}
assert best_state is not None
generator.load_state_dict(best_state)
generation, generated, labels, predictions = evaluate(
generator,
critic,
judge,
real,
real_labels,
)
results = {
"model": "Pocket WGAN-GP",
"method": "Projection-conditioned WGAN-GP with auxiliary class supervision",
"generator_parameters": parameter_count(generator),
"critic_parameters": parameter_count(critic),
"training_examples": len(real),
"generator_steps": generator_steps,
"critic_updates": generator_steps * critic_steps,
"best_generator_step": best_step,
"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(generator.state_dict(), ARTIFACT_DIR / "generator.safetensors")
save_file(critic.state_dict(), ARTIFACT_DIR / "critic.safetensors")
save_grid(generated, labels, ARTIFACT_DIR / "samples.png")
np.savez_compressed(
ARTIFACT_DIR / "generated_samples.npz",
pixels=generated.numpy(),
labels=labels.numpy(),
judge_predictions=predictions.numpy(),
)
(ARTIFACT_DIR / "evaluation.json").write_text(
json.dumps(results, indent=2),
encoding="utf-8",
)
pd.DataFrame(
{
"label": labels.numpy(),
"judge_prediction": predictions.numpy(),
"pixels": list(generated.numpy()),
}
).to_parquet(DATA_DIR / "evaluation_samples.parquet", index=False)
trackio.log(
{
"final_judge_accuracy": generation["judge_accuracy"],
"final_diversity_ratio": generation["mean_diversity_ratio"],
"final_unique_fraction": generation[
"collapse_and_memorization_checks"
]["mean_quantized_unique_fraction"],
}
)
trackio.finish()
print(json.dumps(results, indent=2))
if __name__ == "__main__":
main()