(self, transform, expected_shape)
| 89 | |
| 90 | @parameterized.expand([TEST_CASE_1, TEST_CASE_2, TEST_CASE_3]) |
| 91 | def test_shape(self, transform, expected_shape): |
| 92 | test_image = nib.Nifti1Image(np.random.randint(0, 2, size=[128, 128, 128]).astype(float), np.eye(4)) |
| 93 | with tempfile.TemporaryDirectory() as tempdir: |
| 94 | nib.save(test_image, os.path.join(tempdir, "test_image1.nii.gz")) |
| 95 | nib.save(test_image, os.path.join(tempdir, "test_label1.nii.gz")) |
| 96 | nib.save(test_image, os.path.join(tempdir, "test_extra1.nii.gz")) |
| 97 | nib.save(test_image, os.path.join(tempdir, "test_image2.nii.gz")) |
| 98 | nib.save(test_image, os.path.join(tempdir, "test_label2.nii.gz")) |
| 99 | nib.save(test_image, os.path.join(tempdir, "test_extra2.nii.gz")) |
| 100 | test_data = [ |
| 101 | { |
| 102 | "image": os.path.join(tempdir, "test_image1.nii.gz"), |
| 103 | "label": os.path.join(tempdir, "test_label1.nii.gz"), |
| 104 | "extra": os.path.join(tempdir, "test_extra1.nii.gz"), |
| 105 | }, |
| 106 | { |
| 107 | "image": os.path.join(tempdir, "test_image2.nii.gz"), |
| 108 | "label": os.path.join(tempdir, "test_label2.nii.gz"), |
| 109 | "extra": os.path.join(tempdir, "test_extra2.nii.gz"), |
| 110 | }, |
| 111 | ] |
| 112 | |
| 113 | cache_dir = os.path.join(os.path.join(tempdir, "cache"), "data") |
| 114 | dataset_precached = PersistentDataset(data=test_data, transform=transform, cache_dir=cache_dir) |
| 115 | data1_precached = dataset_precached[0] |
| 116 | data2_precached = dataset_precached[1] |
| 117 | |
| 118 | dataset_postcached = PersistentDataset(data=test_data, transform=transform, cache_dir=cache_dir) |
| 119 | data1_postcached = dataset_postcached[0] |
| 120 | data2_postcached = dataset_postcached[1] |
| 121 | data3_postcached = dataset_postcached[0:2] |
| 122 | |
| 123 | if transform is None: |
| 124 | self.assertEqual(data1_precached["image"], os.path.join(tempdir, "test_image1.nii.gz")) |
| 125 | self.assertEqual(data2_precached["label"], os.path.join(tempdir, "test_label2.nii.gz")) |
| 126 | self.assertEqual(data1_postcached["image"], os.path.join(tempdir, "test_image1.nii.gz")) |
| 127 | self.assertEqual(data2_postcached["extra"], os.path.join(tempdir, "test_extra2.nii.gz")) |
| 128 | else: |
| 129 | self.assertTupleEqual(data1_precached["image"].shape, expected_shape) |
| 130 | self.assertTupleEqual(data1_precached["label"].shape, expected_shape) |
| 131 | self.assertTupleEqual(data1_precached["extra"].shape, expected_shape) |
| 132 | self.assertTupleEqual(data2_precached["image"].shape, expected_shape) |
| 133 | self.assertTupleEqual(data2_precached["label"].shape, expected_shape) |
| 134 | self.assertTupleEqual(data2_precached["extra"].shape, expected_shape) |
| 135 | |
| 136 | self.assertTupleEqual(data1_postcached["image"].shape, expected_shape) |
| 137 | self.assertTupleEqual(data1_postcached["label"].shape, expected_shape) |
| 138 | self.assertTupleEqual(data1_postcached["extra"].shape, expected_shape) |
| 139 | self.assertTupleEqual(data2_postcached["image"].shape, expected_shape) |
| 140 | self.assertTupleEqual(data2_postcached["label"].shape, expected_shape) |
| 141 | self.assertTupleEqual(data2_postcached["extra"].shape, expected_shape) |
| 142 | for d in data3_postcached: |
| 143 | self.assertTupleEqual(d["image"].shape, expected_shape) |
| 144 | |
| 145 | # update the data to cache |
| 146 | test_data_new = [ |
| 147 | { |
| 148 | "image": os.path.join(tempdir, "test_image1_new.nii.gz"), |
nothing calls this directly
no test coverage detected