(self, transform, expected_shape)
| 192 | |
| 193 | @parameterized.expand([TEST_CASE_1, TEST_CASE_2]) |
| 194 | def test_hash_as_key(self, transform, expected_shape): |
| 195 | test_image = nib.Nifti1Image(np.random.randint(0, 2, size=[128, 128, 128]).astype(float), np.eye(4)) |
| 196 | with tempfile.TemporaryDirectory() as tempdir: |
| 197 | test_data = [] |
| 198 | for i in ["1", "2", "2", "3", "3"]: |
| 199 | for k in ["image", "label", "extra"]: |
| 200 | nib.save(test_image, os.path.join(tempdir, f"{k}{i}.nii.gz")) |
| 201 | test_data.append({k: os.path.join(tempdir, f"{k}{i}.nii.gz") for k in ["image", "label", "extra"]}) |
| 202 | |
| 203 | dataset = CacheDataset(data=test_data, transform=transform, cache_num=4, num_workers=2, hash_as_key=True) |
| 204 | self.assertEqual(len(dataset), 5) |
| 205 | # ensure no duplicated cache content |
| 206 | self.assertEqual(len(dataset._cache), 3) |
| 207 | self.assertEqual(len(dataset._hash_keys), 3) |
| 208 | self.assertEqual(dataset.cache_num, 3) |
| 209 | data1 = dataset[0] |
| 210 | data2 = dataset[1] |
| 211 | data3 = dataset[-1] |
| 212 | # test slice indices |
| 213 | data4 = dataset[0:-1] |
| 214 | self.assertEqual(len(data4), 4) |
| 215 | |
| 216 | if transform is None: |
| 217 | self.assertEqual(data1["image"], os.path.join(tempdir, "image1.nii.gz")) |
| 218 | self.assertEqual(data2["label"], os.path.join(tempdir, "label2.nii.gz")) |
| 219 | self.assertEqual(data3["image"], os.path.join(tempdir, "image3.nii.gz")) |
| 220 | else: |
| 221 | self.assertTupleEqual(data1["image"].shape, expected_shape) |
| 222 | self.assertTupleEqual(data2["label"].shape, expected_shape) |
| 223 | self.assertTupleEqual(data3["image"].shape, expected_shape) |
| 224 | for d in data4: |
| 225 | self.assertTupleEqual(d["image"].shape, expected_shape) |
| 226 | |
| 227 | test_data2 = test_data[:3] |
| 228 | dataset.set_data(data=test_data2) |
| 229 | self.assertEqual(len(dataset), 3) |
| 230 | # ensure no duplicated cache content |
| 231 | self.assertEqual(len(dataset._cache), 2) |
| 232 | self.assertEqual(dataset.cache_num, 2) |
| 233 | |
| 234 | |
| 235 | if __name__ == "__main__": |
nothing calls this directly
no test coverage detected