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