#!/usr/bin/env python3 """Run a deterministic full-split binary fine-tuning parity workload.""" 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, labels, weights): values = torch.nn.functional.binary_cross_entropy_with_logits( logits.reshape(-1), labels, reduction="none" ) result = logits.new_zeros(()) for label in torch.unique(labels): mask = labels == label result = result + (weights[mask] * values[mask]).sum() / weights[mask].sum() return result / len(torch.unique(labels)) def make_model(graph, globals_, checkpoint, trainable): from gnn4colliders.models.root_gnn import ( EdgeNetwork, FineTunedEdgeNetwork, load_legacy_edge_network_state_dict, ) backbone = EdgeNetwork(graph, globals_[:1], 64, 12, 4, 4, dropout=0.0) payload = torch.load(checkpoint, map_location="cpu", weights_only=False) load_legacy_edge_network_state_dict(backbone, payload) return FineTunedEdgeNetwork(backbone, 1, freeze_backbone=not trainable) def run_batches(model, artifact, indices, batch_size, optimizer=None): values = [] for start in range(0, len(indices), batch_size): batch_indices = indices[start : start + batch_size] graphs = [_graph(artifact, int(index)) for index in batch_indices] graph = dgl.batch(graphs) labels = torch.from_numpy(artifact.labels[batch_indices]).to(torch.float32) weights = torch.from_numpy(artifact.weights[batch_indices]).to(torch.float32) if optimizer is not None: optimizer.zero_grad() logits = model(graph, None).reshape(-1) batch_loss = loss(logits, labels, weights) if optimizer is not None: batch_loss.backward() optimizer.step() values.append( (batch_indices, logits.detach().numpy(), float(batch_loss.detach())) ) return values 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("--batch-size", type=int, default=1024) parser.add_argument("--epochs", type=int, default=1) parser.add_argument("--trainable-backbone", action="store_true") args = parser.parse_args() torch.manual_seed(20260818) artifact = load_artifact(args.artifact) all_indices = np.arange(artifact.event_count) split = artifact.folds % 5 train = all_indices[split < 3] validation = all_indices[split == 3] test = all_indices[split == 4] generator = np.random.default_rng(20260818) train = train[generator.permutation(len(train))] first_graph = _graph(artifact, int(train[0])) model = make_model( first_graph, torch.empty((1, 0), dtype=torch.float32), args.checkpoint, args.trainable_backbone, ) # Match the fixed-head initialization used by the preceding parity gates. classifier = model.classifier torch.manual_seed(20260818) classifier.load_state_dict(torch.nn.Linear(64, 1).state_dict()) optimizer = torch.optim.Adam(model.parameters(), lr=1e-4) history = [] for epoch in range(args.epochs): model.train() train_values = run_batches(model, artifact, train, args.batch_size, optimizer) history.append(float(np.mean([value[2] for value in train_values]))) print(f"epoch={epoch + 1} train_loss={history[-1]:.12g}", flush=True) model.eval() with torch.inference_mode(): test_values = run_batches(model, artifact, test, args.batch_size) test_indices = np.concatenate([value[0] for value in test_values]) test_logits = np.concatenate([value[1] for value in test_values]) args.output.parent.mkdir(parents=True, exist_ok=True) torch.save( { "model_state_dict": model.state_dict(), "optimizer_state_dict": optimizer.state_dict(), "epoch": args.epochs, }, args.output.with_suffix(".pt"), ) np.savez_compressed( args.output.with_suffix(".npz"), indices=test_indices, logits=test_logits, labels=artifact.labels[test_indices], weights=artifact.weights[test_indices], history=np.asarray(history), ) print( f"events train={len(train)} validation={len(validation)} test={len(test)}", flush=True, ) return 0 if __name__ == "__main__": raise SystemExit(main())