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