Download validation/binary_step.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 3.94 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/binary_step.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/validation/binary_step.py
-
curl -L -o binary_step.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/validation/binary_step.py
3.94 kB
| #!/usr/bin/env python3 | |
| """Run deterministic binary fine-tuning steps on a balanced event batch.""" | |
| 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: torch.Tensor, labels: torch.Tensor, weights: torch.Tensor | |
| ) -> torch.Tensor: | |
| logits = logits.reshape(-1) | |
| elementwise = torch.nn.functional.binary_cross_entropy_with_logits( | |
| logits, labels.to(dtype=logits.dtype), reduction="none" | |
| ) | |
| result = logits.new_zeros(()) | |
| for label in torch.unique(labels): | |
| mask = labels == label | |
| result = ( | |
| result + (weights[mask] * elementwise[mask]).sum() / weights[mask].sum() | |
| ) | |
| return result / len(torch.unique(labels)) | |
| def _flat(values): | |
| # Legacy and rewrite register the corresponding modules in the same | |
| # architectural order; their parameter names intentionally differ. | |
| return torch.cat([value.detach().cpu().reshape(-1) for _, value in values]).numpy() | |
| 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("--per-class-events", type=int, default=512) | |
| parser.add_argument("--epochs", type=int, default=5) | |
| parser.add_argument("--trainable-backbone", action="store_true") | |
| args = parser.parse_args() | |
| artifact = load_artifact(args.artifact) | |
| count = args.per_class_events | |
| indices = np.concatenate([np.arange(count), np.arange(100000, 100000 + count)]) | |
| graphs = [_graph(artifact, int(index)) for index in indices] | |
| graph = dgl.batch(graphs) | |
| empty_globals = torch.empty((len(graphs), 0), dtype=torch.float32) | |
| from gnn4colliders.models.root_gnn import ( | |
| EdgeNetwork, | |
| FineTunedEdgeNetwork, | |
| load_legacy_edge_network_state_dict, | |
| ) | |
| backbone = EdgeNetwork(graphs[0], empty_globals[:1], 64, 12, 4, 4, dropout=0.0) | |
| payload = torch.load(args.checkpoint, map_location="cpu", weights_only=False) | |
| load_legacy_edge_network_state_dict(backbone, payload) | |
| model = FineTunedEdgeNetwork( | |
| backbone, 1, freeze_backbone=not args.trainable_backbone | |
| ) | |
| classifier = model.classifier | |
| torch.manual_seed(20260818) | |
| classifier.load_state_dict(torch.nn.Linear(64, 1).state_dict()) | |
| model.train() | |
| labels = torch.from_numpy(artifact.labels[indices]).to(torch.float32) | |
| weights = torch.from_numpy(artifact.weights[indices]).to(torch.float32) | |
| optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) | |
| losses = [] | |
| logits = None | |
| gradients = None | |
| for _ in range(args.epochs): | |
| optimizer.zero_grad() | |
| logits = model(graph, None).reshape(-1) | |
| loss = _loss(logits, labels, weights) | |
| loss.backward() | |
| gradients = _flat( | |
| [ | |
| (name, parameter.grad) | |
| for name, parameter in model.named_parameters() | |
| if parameter.grad is not None | |
| ] | |
| ) | |
| optimizer.step() | |
| losses.append(float(loss.detach())) | |
| assert logits is not None and gradients is not None | |
| parameters = _flat(list(model.named_parameters())) | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed( | |
| args.output, | |
| logits=logits.detach().numpy(), | |
| labels=labels.numpy(), | |
| weights=weights.numpy(), | |
| loss=np.asarray(losses[-1]), | |
| history=np.asarray(losses, dtype=np.float64), | |
| gradients=gradients, | |
| parameters=parameters, | |
| ) | |
| print( | |
| f"rewrite: {len(indices)} events, " | |
| f"epochs={args.epochs}, final_loss={losses[-1]:.12g}" | |
| ) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |