Download validation/end_to_end_binary.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 4.72 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/end_to_end_binary.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/validation/end_to_end_binary.py
-
curl -L -o end_to_end_binary.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/end_to_end_binary.py
4.72 kB
| #!/usr/bin/env python3 | |
| """Run a deterministic full-split binary fine-tuning parity workload.""" | |
| from __future__ import annotations | |
| import argparse | |
| from pathlib import Path | |
| import dgl | |
| import numpy as np | |
| import torch | |
| from validation.artifacts import load_artifact | |
| from validation.forward import _graph | |
| def loss(logits, labels, weights): | |
| values = torch.nn.functional.binary_cross_entropy_with_logits( | |
| logits.reshape(-1), labels, reduction="none" | |
| ) | |
| result = logits.new_zeros(()) | |
| for label in torch.unique(labels): | |
| mask = labels == label | |
| result = result + (weights[mask] * values[mask]).sum() / weights[mask].sum() | |
| return result / len(torch.unique(labels)) | |
| def make_model(graph, globals_, checkpoint, trainable): | |
| from gnn4colliders.models.root_gnn import ( | |
| EdgeNetwork, | |
| FineTunedEdgeNetwork, | |
| load_legacy_edge_network_state_dict, | |
| ) | |
| backbone = EdgeNetwork(graph, globals_[:1], 64, 12, 4, 4, dropout=0.0) | |
| payload = torch.load(checkpoint, map_location="cpu", weights_only=False) | |
| load_legacy_edge_network_state_dict(backbone, payload) | |
| return FineTunedEdgeNetwork(backbone, 1, freeze_backbone=not trainable) | |
| def run_batches(model, artifact, indices, batch_size, optimizer=None): | |
| values = [] | |
| for start in range(0, len(indices), batch_size): | |
| batch_indices = indices[start : start + batch_size] | |
| graphs = [_graph(artifact, int(index)) for index in batch_indices] | |
| graph = dgl.batch(graphs) | |
| labels = torch.from_numpy(artifact.labels[batch_indices]).to(torch.float32) | |
| weights = torch.from_numpy(artifact.weights[batch_indices]).to(torch.float32) | |
| if optimizer is not None: | |
| optimizer.zero_grad() | |
| logits = model(graph, None).reshape(-1) | |
| batch_loss = loss(logits, labels, weights) | |
| if optimizer is not None: | |
| batch_loss.backward() | |
| optimizer.step() | |
| values.append( | |
| (batch_indices, logits.detach().numpy(), float(batch_loss.detach())) | |
| ) | |
| return values | |
| def main() -> int: | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--artifact", type=Path, required=True) | |
| parser.add_argument("--checkpoint", type=Path, required=True) | |
| parser.add_argument("--output", type=Path, required=True) | |
| parser.add_argument("--batch-size", type=int, default=1024) | |
| parser.add_argument("--epochs", type=int, default=1) | |
| parser.add_argument("--trainable-backbone", action="store_true") | |
| args = parser.parse_args() | |
| torch.manual_seed(20260818) | |
| artifact = load_artifact(args.artifact) | |
| all_indices = np.arange(artifact.event_count) | |
| split = artifact.folds % 5 | |
| train = all_indices[split < 3] | |
| validation = all_indices[split == 3] | |
| test = all_indices[split == 4] | |
| generator = np.random.default_rng(20260818) | |
| train = train[generator.permutation(len(train))] | |
| first_graph = _graph(artifact, int(train[0])) | |
| model = make_model( | |
| first_graph, | |
| torch.empty((1, 0), dtype=torch.float32), | |
| args.checkpoint, | |
| args.trainable_backbone, | |
| ) | |
| # Match the fixed-head initialization used by the preceding parity gates. | |
| classifier = model.classifier | |
| torch.manual_seed(20260818) | |
| classifier.load_state_dict(torch.nn.Linear(64, 1).state_dict()) | |
| optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) | |
| history = [] | |
| for epoch in range(args.epochs): | |
| model.train() | |
| train_values = run_batches(model, artifact, train, args.batch_size, optimizer) | |
| history.append(float(np.mean([value[2] for value in train_values]))) | |
| print(f"epoch={epoch + 1} train_loss={history[-1]:.12g}", flush=True) | |
| model.eval() | |
| with torch.inference_mode(): | |
| test_values = run_batches(model, artifact, test, args.batch_size) | |
| test_indices = np.concatenate([value[0] for value in test_values]) | |
| test_logits = np.concatenate([value[1] for value in test_values]) | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| torch.save( | |
| { | |
| "model_state_dict": model.state_dict(), | |
| "optimizer_state_dict": optimizer.state_dict(), | |
| "epoch": args.epochs, | |
| }, | |
| args.output.with_suffix(".pt"), | |
| ) | |
| np.savez_compressed( | |
| args.output.with_suffix(".npz"), | |
| indices=test_indices, | |
| logits=test_logits, | |
| labels=artifact.labels[test_indices], | |
| weights=artifact.weights[test_indices], | |
| history=np.asarray(history), | |
| ) | |
| print( | |
| f"events train={len(train)} validation={len(validation)} test={len(test)}", | |
| flush=True, | |
| ) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |