pocket-diffusion / train.py
ARotting's picture
Publish Interactive classifier-free guided digit diffusion
0213535 verified
Raw
History Blame Contribute Delete
7.33 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 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()