#!/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())