MCPcopy Create free account
hub / github.com/Project-MONAI/MONAI / test_shape

Method test_shape

tests/data/test_persistentdataset.py:91–163  ·  view source on GitHub ↗
(self, transform, expected_shape)

Source from the content-addressed store, hash-verified

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"),

Callers

nothing calls this directly

Calls 4

set_dataMethod · 0.95
PersistentDatasetClass · 0.90
astypeMethod · 0.80
saveMethod · 0.80

Tested by

no test coverage detected