import unittest import numpy as np from dataset import _ExpoBaseDataset from targets import RESERVED_IDX class ReservedEventFilteringTests(unittest.TestCase): @staticmethod def _base(extra_info_types): dataset = _ExpoBaseDataset.__new__(_ExpoBaseDataset) dataset.extra_info_types = list(extra_info_types) dataset.event_data = np.asarray( [ [101, 10, RESERVED_IDX], [101, 20, 2], [101, 30, 3], ], dtype=np.float64, ) return dataset def _assert_reserved_event_removed(self, extra_info_types): dataset = self._base(extra_info_types) rows = list(dataset._iter_patient_events(impute_no_event_gaps=False)) self.assertEqual(len(rows), 1) eid, times, labels = rows[0] self.assertEqual(eid, 101) np.testing.assert_array_equal(times, np.asarray([20, 30], dtype=np.float32)) np.testing.assert_array_equal(labels, np.asarray([3, 4], dtype=np.int64)) self.assertNotIn(RESERVED_IDX, labels.tolist()) def test_empty_extra_info_removes_legacy_reserved_event(self): self._assert_reserved_event_removed([]) def test_selected_extra_info_removes_legacy_reserved_event(self): self._assert_reserved_event_removed([11]) if __name__ == "__main__": unittest.main()