Download tests/unit/sequences/test_sequence_sample.py from HWresearch/GNN4Colliders: direct link, hf CLI and curl.
- Browser
- Download file 1.76 kB
-
https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/tests/unit/sequences/test_sequence_sample.py
- Command line
-
hf download hf://HWresearch/GNN4Colliders/tests/unit/sequences/test_sequence_sample.py
-
curl -L -o test_sequence_sample.py https://huggingface.co/HWresearch/GNN4Colliders/resolve/main/tests/unit/sequences/test_sequence_sample.py
1.76 kB
| import torch | |
| from gnn4colliders.data import EventMetadata, EventSample | |
| from gnn4colliders.sequences import ( | |
| SequenceDataLoader, | |
| SequenceSample, | |
| batch_sequence_samples, | |
| sequence_sample_from_event, | |
| ) | |
| def _metadata(index: int) -> EventMetadata: | |
| return EventMetadata(fold=index, weight=1.0, sample_id=f"event:{index}") | |
| def test_sequence_feature_builder_reuses_shared_event_features(): | |
| event = EventSample( | |
| objects={ | |
| "pt": [10.0, 20.0], | |
| "eta": [0.1, -0.2], | |
| "phi": [0.3, -0.4], | |
| }, | |
| label=1, | |
| global_features=torch.empty(0), | |
| event_index=0, | |
| metadata=_metadata(0), | |
| ) | |
| sample = sequence_sample_from_event( | |
| event, | |
| [["pt"], ["eta"], ["phi"], "CALC_E", [1.0], [0.0], "NODE_TYPE"], | |
| ["vector"], | |
| [1.0] * 7, | |
| ) | |
| assert sample.tokens.shape == (2, 7) | |
| assert sample.label.item() == 1 | |
| assert sample.metadata.sample_id == "event:0" | |
| def test_sequence_batch_pads_and_masks_in_source_order(): | |
| samples = [ | |
| SequenceSample( | |
| tokens=torch.ones(2, 3), | |
| label=torch.tensor(1), | |
| global_features=None, | |
| metadata=_metadata(1), | |
| ), | |
| SequenceSample( | |
| tokens=torch.ones(1, 3), | |
| label=torch.tensor(0), | |
| global_features=None, | |
| metadata=_metadata(2), | |
| ), | |
| ] | |
| batch = batch_sequence_samples(samples) | |
| assert batch.tokens.shape == (2, 2, 3) | |
| assert batch.token_mask.tolist() == [[True, True], [True, False]] | |
| assert batch.metadata.sample_id == ("event:1", "event:2") | |
| loader = SequenceDataLoader(samples, batch_size=1, shuffle=True, seed=9) | |
| loader.set_epoch(2) | |
| assert len(list(loader)) == 2 | |