#!/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())