Remove legacy event and mixed distribution paths
This commit is contained in:
44
tests/test_dataset_reserved_event.py
Normal file
44
tests/test_dataset_reserved_event.py
Normal file
@@ -0,0 +1,44 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user