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