GNN4Colliders / validation /batching.py
ho22joshua's picture
Remove historical implementation from active tree
de46a3c
Raw History Blame Contribute Delete
2.08 kB
#!/usr/bin/env python3
"""Dump deterministic batch membership and aggregate sizes."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import torch
from validation.artifacts import load_artifact
from validation.forward import _graph
def main() -> int:
parser = argparse.ArgumentParser()
parser.add_argument("--artifact", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--batch-size", type=int, default=8)
args = parser.parse_args()
artifact = load_artifact(args.artifact)
from gnn4colliders.data import (
EventMetadata,
GraphDataLoader,
GraphDataset,
GraphSample,
)
samples = []
for index in range(artifact.event_count):
samples.append(
GraphSample(
_graph(artifact, index),
torch.tensor(artifact.labels[index]),
None,
EventMetadata(
int(artifact.folds[index]),
float(artifact.weights[index]),
str(artifact.sample_id[index]),
{"index": index},
),
)
)
loader = GraphDataLoader(GraphDataset(samples), args.batch_size, shuffle=False)
batches = []
for batch in loader:
graph = batch.graph
labels = batch.labels
ids = list(batch.metadata.sample_id)
weights = batch.metadata.weight.tolist()
batches.append(
{
"sample_id": ids,
"num_graphs": len(ids),
"nodes": int(graph.num_nodes()),
"edges": int(graph.num_edges()),
"labels": labels.reshape(-1).tolist(),
"weights": weights,
}
)
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(
json.dumps({"batch_size": args.batch_size, "batches": batches}, indent=2) + "\n"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())