File size: 4,723 Bytes
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
116
117
118
119
120
121
122
123
124
125
126
127
128
129
#!/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())