diff --git a/cebra/data/base.py b/cebra/data/base.py index f5491e51..2b08c65a 100644 --- a/cebra/data/base.py +++ b/cebra/data/base.py @@ -154,8 +154,6 @@ def expand_index_in_trial(self, index, trial_ids, trial_borders): trial_ids is in size of a length of self.index and indicate the trial id of the index belong to. trial_borders is in size of a length of self.idnex and indicate the border of each trial. - Todo: - - rewrite """ # TODO(stes) potential room for speed improvements by pre-allocating these tensors/ @@ -163,16 +161,15 @@ def expand_index_in_trial(self, index, trial_ids, trial_borders): offset = torch.arange(-self.offset.left, self.offset.right, device=index.device) - index = torch.tensor( - [ - torch.clamp( - i, - trial_borders[trial_ids[i]] + self.offset.left, - trial_borders[trial_ids[i] + 1] - self.offset.right, - ) for i in index - ], - device=self.device, - ) + + trial_ids = torch.as_tensor(trial_ids, device=index.device) + trial_borders = torch.as_tensor(trial_borders, device=index.device) + batch_trial_ids = trial_ids[index] + min_borders = trial_borders[batch_trial_ids] + self.offset.left + max_borders = trial_borders[batch_trial_ids + 1] - self.offset.right + + index = torch.clamp(index, min=min_borders, max=max_borders) + return index[:, None] + offset[None, :] @abc.abstractmethod diff --git a/tests/test_loader.py b/tests/test_loader.py index cb6be9a7..792258d6 100644 --- a/tests/test_loader.py +++ b/tests/test_loader.py @@ -103,6 +103,38 @@ def test_offset(): offset = cebra.data.Offset(4, -2) +def _expand_index_in_trial_old(dataset, index, trial_ids, trial_borders): + """Reference implementation of Dataset.expand_index_in_trial.""" + offset = torch.arange(-dataset.offset.left, + dataset.offset.right, + device=index.device) + index = torch.tensor( + [ + torch.clamp( + i, + trial_borders[trial_ids[i]] + dataset.offset.left, + trial_borders[trial_ids[i] + 1] - dataset.offset.right, + ) for i in index + ], + device=dataset.device, + ) + return index[:, None] + offset[None, :] + + +def test_expand_index_in_trial_matches_previous_implementation(): + dataset = RandomDataset(N=12) + dataset.offset = cebra.data.Offset(2, 2) + index = torch.tensor([0, 2, 5, 7, 11]) + trial_ids = np.repeat(np.arange(3), 4) + trial_borders = [0, 4, 8, 12] + + expected = _expand_index_in_trial_old(dataset, index, trial_ids, + trial_borders) + actual = dataset.expand_index_in_trial(index, trial_ids, trial_borders) + + torch.testing.assert_close(actual, expected) + + def _assert_dataset_on_correct_device(loader, device): assert hasattr(loader, "dataset") assert hasattr(loader, "device")