File size: 3,942 Bytes
c93e291 de46a3c c93e291 de46a3c c93e291 de46a3c c93e291 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 | #!/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())
|