Download tests/unit/data/test_batching.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 1.38 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/tests/unit/data/test_batching.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/tests/unit/data/test_batching.py
-
curl -L -o test_batching.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/tests/unit/data/test_batching.py
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}"), | |
| ) | |
| 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 | |
| ] | |