#!/usr/bin/env python3 """Run fixed-weight ROOT-GNN inference on a normalized validation artifact.""" from __future__ import annotations import argparse from dataclasses import replace from pathlib import Path import dgl import numpy as np import torch from validation.artifacts import load_artifact, save_artifact def _graph(artifact, index: int): nodes = torch.from_numpy(artifact.event_nodes(index)).to(torch.float32) src, dst, edges = artifact.event_edges(index) graph = dgl.graph( (torch.from_numpy(src), torch.from_numpy(dst)), num_nodes=nodes.shape[0] ) graph.ndata["features"] = nodes graph.edata["features"] = torch.from_numpy(edges).to(torch.float32) return 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) args = parser.parse_args() artifact = load_artifact(args.artifact) event_count = ( artifact.event_count if args.events is None else min(args.events, artifact.event_count) ) graphs = [_graph(artifact, index) for index in range(event_count)] batch_graph = dgl.batch(graphs) first = graphs[0] empty_globals = torch.empty((1, 0), dtype=torch.float32) from gnn4colliders.models.root_gnn import ( EdgeNetwork, load_legacy_edge_network_state_dict, ) model = EdgeNetwork(first, 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(model, payload) model.eval() with torch.inference_mode(): logits = model(batch_graph, None).cpu().numpy() args.output.mkdir(parents=True, exist_ok=True) reload_path = args.output / "checkpoint_reload.pt" torch.save({"model_state_dict": model.state_dict()}, reload_path) fresh = type(model)(first, empty_globals, 64, 12, 4, 4, dropout=0.0) fresh.load_state_dict( torch.load(reload_path, map_location="cpu")["model_state_dict"] ) fresh.eval() with torch.inference_mode(): reloaded_logits = fresh(batch_graph, None).cpu().numpy() scores = torch.softmax(torch.from_numpy(logits), dim=1).numpy() predictions = scores.argmax(axis=1) updated = replace( artifact, logits=logits, scores=scores, predictions=predictions, manifest={ **artifact.manifest, "implementation": "rewrite", "checkpoint": str(args.checkpoint), "reload_max_abs": float(np.abs(logits - reloaded_logits).max()), }, ) save_artifact(updated, args.output) return 0 if __name__ == "__main__": raise SystemExit(main())