(self, transform, expected_shape)
| 132 | @unittest.skipUnless(has_nib, "Requires nibabel package.") |
| 133 | @parameterized.expand([TEST_CASE_1, TEST_CASE_2, TEST_CASE_3]) |
| 134 | def test_shape(self, transform, expected_shape): |
| 135 | test_image = nib.Nifti1Image(np.random.randint(0, 2, size=[128, 128, 128]).astype(float), np.eye(4)) |
| 136 | with tempfile.TemporaryDirectory() as tempdir: |
| 137 | nib.save(test_image, os.path.join(tempdir, "test_image1.nii.gz")) |
| 138 | nib.save(test_image, os.path.join(tempdir, "test_label1.nii.gz")) |
| 139 | nib.save(test_image, os.path.join(tempdir, "test_extra1.nii.gz")) |
| 140 | nib.save(test_image, os.path.join(tempdir, "test_image2.nii.gz")) |
| 141 | nib.save(test_image, os.path.join(tempdir, "test_label2.nii.gz")) |
| 142 | nib.save(test_image, os.path.join(tempdir, "test_extra2.nii.gz")) |
| 143 | test_data = [ |
| 144 | { |
| 145 | "image": os.path.join(tempdir, "test_image1.nii.gz"), |
| 146 | "label": os.path.join(tempdir, "test_label1.nii.gz"), |
| 147 | "extra": os.path.join(tempdir, "test_extra1.nii.gz"), |
| 148 | }, |
| 149 | { |
| 150 | "image": os.path.join(tempdir, "test_image2.nii.gz"), |
| 151 | "label": os.path.join(tempdir, "test_label2.nii.gz"), |
| 152 | "extra": os.path.join(tempdir, "test_extra2.nii.gz"), |
| 153 | }, |
| 154 | ] |
| 155 | |
| 156 | cache_dir = os.path.join(os.path.join(tempdir, "cache"), "data") |
| 157 | dataset_precached = GDSDataset(data=test_data, transform=transform, cache_dir=cache_dir, device=0) |
| 158 | data1_precached = dataset_precached[0] |
| 159 | data2_precached = dataset_precached[1] |
| 160 | |
| 161 | dataset_postcached = GDSDataset(data=test_data, transform=transform, cache_dir=cache_dir, device=0) |
| 162 | data1_postcached = dataset_postcached[0] |
| 163 | data2_postcached = dataset_postcached[1] |
| 164 | data3_postcached = dataset_postcached[0:2] |
| 165 | |
| 166 | if transform is None: |
| 167 | self.assertEqual(data1_precached["image"], os.path.join(tempdir, "test_image1.nii.gz")) |
| 168 | self.assertEqual(data2_precached["label"], os.path.join(tempdir, "test_label2.nii.gz")) |
| 169 | self.assertEqual(data1_postcached["image"], os.path.join(tempdir, "test_image1.nii.gz")) |
| 170 | self.assertEqual(data2_postcached["extra"], os.path.join(tempdir, "test_extra2.nii.gz")) |
| 171 | else: |
| 172 | self.assertTupleEqual(data1_precached["image"].shape, expected_shape) |
| 173 | self.assertTupleEqual(data1_precached["label"].shape, expected_shape) |
| 174 | self.assertTupleEqual(data1_precached["extra"].shape, expected_shape) |
| 175 | self.assertTupleEqual(data2_precached["image"].shape, expected_shape) |
| 176 | self.assertTupleEqual(data2_precached["label"].shape, expected_shape) |
| 177 | self.assertTupleEqual(data2_precached["extra"].shape, expected_shape) |
| 178 | |
| 179 | self.assertTupleEqual(data1_postcached["image"].shape, expected_shape) |
| 180 | self.assertTupleEqual(data1_postcached["label"].shape, expected_shape) |
| 181 | self.assertTupleEqual(data1_postcached["extra"].shape, expected_shape) |
| 182 | self.assertTupleEqual(data2_postcached["image"].shape, expected_shape) |
| 183 | self.assertTupleEqual(data2_postcached["label"].shape, expected_shape) |
| 184 | self.assertTupleEqual(data2_postcached["extra"].shape, expected_shape) |
| 185 | for d in data3_postcached: |
| 186 | self.assertTupleEqual(d["image"].shape, expected_shape) |
| 187 | |
| 188 | # update the data to cache |
| 189 | test_data_new = [ |
| 190 | { |
| 191 | "image": os.path.join(tempdir, "test_image1_new.nii.gz"), |
nothing calls this directly
no test coverage detected