#!/usr/bin/env python3 """Run canonical ROOT-GNN transfer-learning forward with a fixed binary head.""" 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 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("--events", type=int, default=96) parser.add_argument("--trainable-backbone", action="store_true") args = parser.parse_args() artifact = load_artifact(args.artifact) count = min(args.events, artifact.event_count) graphs = [_graph(artifact, index) for index in range(count)] batch_graph = dgl.batch(graphs) empty_globals = torch.empty((1, 0), dtype=torch.float32) from gnn4colliders.models.root_gnn import ( EdgeNetwork, FineTunedEdgeNetwork, load_legacy_edge_network_state_dict, ) backbone = EdgeNetwork(graphs[0], empty_globals, 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.eval() with torch.inference_mode(): logits = model(batch_graph, None).cpu().numpy() args.output.parent.mkdir(parents=True, exist_ok=True) np.savez_compressed(args.output, logits=logits) return 0 if __name__ == "__main__": raise SystemExit(main())