flow-pocket / source /train.py
ARotting's picture
Publish Exactly invertible RealNVP and density controls
d486718 verified
Raw
History Blame Contribute Delete
6.68 kB
from __future__ import annotations
import json
from pathlib import Path
import joblib
import numpy as np
import pandas as pd
import torch
import trackio
from data import generate_pinwheel
from model import RealNVP, parameter_count
from safetensors.torch import save_file
from sklearn.mixture import GaussianMixture
PROJECT_DIR = Path(__file__).resolve().parent
ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "flow-pocket"
DATA_DIR = PROJECT_DIR / "data"
def gaussian_fit(data: np.ndarray) -> dict:
return {"mean": data.mean(0), "covariance": np.cov(data.T)}
def gaussian_nll(data: np.ndarray, fit: dict) -> float:
centered = data - fit["mean"]
covariance = fit["covariance"]
inverse = np.linalg.inv(covariance)
log_determinant = np.linalg.slogdet(covariance)[1]
quadratic = np.einsum("bi,ij,bj->b", centered, inverse, centered)
return float(np.mean(np.log(2 * np.pi) + 0.5 * log_determinant + 0.5 * quadratic))
def gaussian_sample(fit: dict, samples: int, seed: int) -> np.ndarray:
return np.random.default_rng(seed).multivariate_normal(
fit["mean"], fit["covariance"], size=samples
)
def rbf_mmd(first: np.ndarray, second: np.ndarray) -> float:
rng = np.random.default_rng(2043)
first = first[rng.choice(len(first), 1000, replace=False)]
second = second[rng.choice(len(second), 1000, replace=False)]
combined = np.concatenate([first, second])
pairs = rng.choice(len(combined), size=(4000, 2), replace=True)
distances = np.sum(
(combined[pairs[:, 0]] - combined[pairs[:, 1]]) ** 2, axis=1
)
bandwidth = max(float(np.median(distances[distances > 0])), 1e-4)
def kernel_mean(left: np.ndarray, right: np.ndarray) -> float:
distances = ((left[:, None, :] - right[None, :, :]) ** 2).sum(2)
return float(np.exp(-distances / (2 * bandwidth)).mean())
return kernel_mean(first, first) + kernel_mean(second, second) - 2 * kernel_mean(
first, second
)
def main() -> None:
torch.manual_seed(2043)
torch.set_num_threads(1)
train_data = generate_pinwheel(40_000, 2043)
validation_data = generate_pinwheel(5_000, 3043)
test_data = generate_pinwheel(10_000, 4043)
model = RealNVP()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=1e-6)
tensor = torch.from_numpy(train_data)
validation = torch.from_numpy(validation_data)
rng = np.random.default_rng(2043)
trackio.init(
project="flow-pocket",
name="realnvp-pinwheel-v1",
config={
"parameters": parameter_count(model),
"coupling_layers": len(model.layers),
"training_examples": len(train_data),
"training_steps": 4_000,
},
)
best_state = None
best_validation = float("inf")
history = []
for step in range(1, 4_001):
batch = tensor[rng.choice(len(tensor), 512, replace=False)]
loss = -model.log_probability(batch).mean()
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
optimizer.step()
if step % 100 == 0:
model.eval()
with torch.inference_mode():
validation_nll = float(
-model.log_probability(validation).mean()
)
record = {
"training_step": step,
"training_nll": float(loss.detach()),
"validation_nll": validation_nll,
}
history.append(record)
trackio.log(record)
if validation_nll < best_validation:
best_validation = validation_nll
best_state = {
name: parameter.detach().clone()
for name, parameter in model.state_dict().items()
}
model.train()
if best_state is not None:
model.load_state_dict(best_state)
model.eval()
gaussian = gaussian_fit(train_data)
mixture = GaussianMixture(
n_components=5,
covariance_type="full",
random_state=2043,
max_iter=500,
n_init=3,
).fit(train_data)
with torch.inference_mode():
flow_nll = float(
-model.log_probability(torch.from_numpy(test_data)).mean()
)
generated_flow = model.sample(5_000, seed=5043).numpy()
latent, _ = model(torch.from_numpy(test_data[:2_000]))
reconstructed = model.inverse(latent)
cycle_error = float(
torch.max(torch.abs(reconstructed - torch.from_numpy(test_data[:2_000])))
)
generated_gaussian = gaussian_sample(gaussian, 5_000, 6043)
generated_mixture, _ = mixture.sample(5_000)
results = {
"realnvp": {
"parameters": parameter_count(model),
"test_nll": flow_nll,
"sample_mmd": rbf_mmd(generated_flow, test_data),
"maximum_cycle_error": cycle_error,
},
"full_covariance_gaussian": {
"test_nll": gaussian_nll(test_data, gaussian),
"sample_mmd": rbf_mmd(generated_gaussian, test_data),
},
"five_component_gmm": {
"test_nll": float(-mixture.score(test_data)),
"sample_mmd": rbf_mmd(generated_mixture, test_data),
},
}
report = {
"benchmark": "Five-arm pinwheel density estimation",
"training_examples": len(train_data),
"heldout_examples": len(test_data),
"results": results,
"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 / "realnvp.safetensors")
joblib.dump(
{"gaussian": gaussian, "mixture": mixture},
ARTIFACT_DIR / "classical_controls.joblib",
)
(ARTIFACT_DIR / "evaluation.json").write_text(
json.dumps(report, indent=2), encoding="utf-8"
)
np.savez_compressed(
ARTIFACT_DIR / "generated_samples.npz",
realnvp=generated_flow,
gaussian=generated_gaussian,
gmm=generated_mixture,
)
pd.DataFrame(test_data, columns=["x", "y"]).to_parquet(
DATA_DIR / "pinwheel_test.parquet", index=False
)
trackio.log(
{
"flow_test_nll": results["realnvp"]["test_nll"],
"flow_sample_mmd": results["realnvp"]["sample_mmd"],
"gmm_test_nll": results["five_component_gmm"]["test_nll"],
"gmm_sample_mmd": results["five_component_gmm"]["sample_mmd"],
}
)
trackio.finish()
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()