Download scripts/train.py from OneScience-Group/NNCAM: direct link, hf CLI and curl.
- Browser
- Download file 4.02 kB
-
https://huggingface.co/OneScience-Group/NNCAM/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/NNCAM/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/NNCAM/resolve/main/scripts/train.py
4.02 kB
| #!/usr/bin/env python3 | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from torch.utils.data import DataLoader, TensorDataset | |
| from model.nncam import build_model, fit_normalizer, normalize_input, scale_output | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def main(): | |
| parser = argparse.ArgumentParser(description="Train NNCAM on the prepared NPZ dataset.") | |
| parser.add_argument("--data", type=Path, default=ROOT / "data/nncam_fake.npz") | |
| parser.add_argument("--checkpoint", type=Path, default=ROOT / "result/checkpoints/nncam.pt") | |
| parser.add_argument("--metrics", type=Path, default=ROOT / "result/training/metrics.json") | |
| parser.add_argument("--epochs", type=int) | |
| parser.add_argument("--batch-size", type=int) | |
| parser.add_argument("--width", type=int) | |
| parser.add_argument("--depth", type=int) | |
| parser.add_argument("--paper-model", action="store_true", help="Explicitly use depth=9, width=256, epochs=18, batch_size=1024 (567361 parameters).") | |
| parser.add_argument("--lr", type=float, default=1e-3) | |
| parser.add_argument("--seed", type=int, default=42) | |
| args = parser.parse_args() | |
| if not args.data.is_file(): | |
| raise FileNotFoundError(f"missing dataset {args.data}; run python scripts/fake_data.py first") | |
| defaults = {"depth": 9, "width": 256, "epochs": 18, "batch_size": 1024} if args.paper_model else {"depth": 4, "width": 32, "epochs": 3, "batch_size": 64} | |
| depth, width = args.depth or defaults["depth"], args.width or defaults["width"] | |
| epochs, batch_size = args.epochs or defaults["epochs"], args.batch_size or defaults["batch_size"] | |
| torch.manual_seed(args.seed) | |
| with np.load(args.data) as data: | |
| x, y = data["x"].astype(np.float32), data["y"].astype(np.float32) | |
| if x.ndim != 2 or x.shape[1] != 94 or y.shape != (x.shape[0], 65): | |
| raise ValueError(f"expected x=[N,94], y=[N,65], got {x.shape}, {y.shape}") | |
| input_mean, input_scale = fit_normalizer(x) | |
| scaled_y = scale_output(y) | |
| target_mean, target_scale = fit_normalizer(scaled_y) | |
| dataset = TensorDataset(torch.from_numpy(normalize_input(x, input_mean, input_scale)), torch.from_numpy(normalize_input(scaled_y, target_mean, target_scale))) | |
| loader = DataLoader(dataset, batch_size=batch_size, shuffle=True) | |
| model = build_model(width=width, depth=depth) | |
| optimizer = torch.optim.Adam(model.parameters(), lr=args.lr) | |
| scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=3, gamma=0.2) | |
| history = [] | |
| for epoch in range(epochs): | |
| total = 0.0 | |
| for xb, yb in loader: | |
| optimizer.zero_grad(set_to_none=True) | |
| loss = torch.nn.functional.mse_loss(model(xb), yb) | |
| loss.backward() | |
| optimizer.step() | |
| total += loss.item() * len(xb) | |
| history.append(total / len(dataset)) | |
| scheduler.step() | |
| print(f"epoch={epoch + 1:02d} loss={history[-1]:.6f}") | |
| parameter_count = sum(parameter.numel() for parameter in model.parameters()) | |
| checkpoint = { | |
| "format_version": 1, | |
| "model": model.state_dict(), | |
| "model_config": model.model_config, | |
| "normalization": {"input_mean": torch.from_numpy(input_mean), "input_scale": torch.from_numpy(input_scale), "target_mean": torch.from_numpy(target_mean), "target_scale": torch.from_numpy(target_scale)}, | |
| "training": {"epochs": epochs, "batch_size": batch_size, "learning_rate": args.lr, "paper_model": args.paper_model, "parameters": parameter_count}, | |
| } | |
| args.checkpoint.parent.mkdir(parents=True, exist_ok=True) | |
| args.metrics.parent.mkdir(parents=True, exist_ok=True) | |
| torch.save(checkpoint, args.checkpoint) | |
| args.metrics.write_text(json.dumps({"loss": history, "final_loss": history[-1], "parameters": parameter_count, "model_config": model.model_config}, indent=2), encoding="utf-8") | |
| print(f"saved {args.checkpoint}; parameters={parameter_count}; paper_model={args.paper_model}") | |
| if __name__ == "__main__": | |
| main() | |