GNN4Colliders / tests /unit /data /test_batching.py
ho22joshua's picture
rewriting codebase (#7)
916755e
Raw History Blame Contribute Delete
1.38 kB
import pytest
import torch
from gnn4colliders.data import EventMetadata, GraphSample, batch_graph_samples
from gnn4colliders.graphs import build_dgl_graph
dgl = pytest.importorskip("dgl")
def _sample(index: int, node_count: int) -> GraphSample:
nodes = torch.arange(node_count * 3, dtype=torch.float32).reshape(node_count, 3)
return GraphSample(
graph=build_dgl_graph(nodes),
label=torch.tensor(index),
global_features=torch.tensor([float(index), -float(index)]),
metadata=EventMetadata(index, index + 0.5, f"sample_{index}"),
)
@pytest.mark.parametrize("counts", [(1,), (1, 2), (1, 3, 2)])
def test_graph_batch_preserves_heterogeneous_order_and_metadata(counts):
samples = [_sample(index, count) for index, count in enumerate(counts)]
batch = batch_graph_samples(samples)
assert batch.labels.tolist() == list(range(len(counts)))
assert batch.metadata.sample_id == tuple(f"sample_{i}" for i in range(len(counts)))
assert batch.metadata.weight.tolist() == pytest.approx(
[i + 0.5 for i in range(len(counts))]
)
assert batch.global_features[:, 0].tolist() == list(map(float, range(len(counts))))
assert batch.graph.batch_num_nodes().tolist() == list(counts)
assert batch.graph.batch_num_edges().tolist() == [
count * (count - 1) if count > 1 else 1 for count in counts
]