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 ]