GNN4Colliders / validation /transfer.py
ho22joshua's picture
Remove historical implementation from active tree
de46a3c
Raw History Blame Contribute Delete
1.88 kB
#!/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())