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

Method test_shape

tests/data/test_gdsdataset.py:134–206  ·  view source on GitHub ↗
(self, transform, expected_shape)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 4

GDSDatasetClass · 0.90
astypeMethod · 0.80
saveMethod · 0.80
set_dataMethod · 0.45

Tested by

no test coverage detected