Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion monai/data/wsi_datasets.py
Original file line number Diff line number Diff line change
Expand Up @@ -306,7 +306,7 @@ def _evaluate_patch_locations(self, sample):
)
)
# convert locations to mask_location
mask_locations = np.round((patch_locations + patch_size_0 // 2) / float(mask_ratio))
mask_locations = np.round((patch_locations + patch_size_0 // 2) / float(mask_ratio)).astype(int)

# fill out samples with location and metadata
sample[WSIPatchKeys.SIZE.value] = patch_size
Expand Down
4 changes: 3 additions & 1 deletion monai/handlers/probability_maps.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,7 +108,9 @@ def __call__(self, engine: Engine) -> None:
locs = engine.state.batch[CommonKeys.IMAGE].meta[ProbMapKeys.LOCATION]
probs = engine.state.output[self.prob_key]
for name, loc, prob in zip(names, locs, probs):
self.prob_map[name][tuple(loc)] = prob
# numpy 2.x rejects non-integer array indices, so coerce here for datasets that
# still supply float locations. The bundled WSI datasets already emit integers.
self.prob_map[name][tuple(int(i) for i in loc)] = prob
with self.lock:
self.counter[name] -= 1
if self.counter[name] == 0:
Expand Down
5 changes: 4 additions & 1 deletion tests/data/test_sliding_patch_wsi_dataset.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@
from parameterized import parameterized

from monai.data import SlidingPatchWSIDataset
from monai.utils import WSIPatchKeys, optional_import, set_determinism
from monai.utils import ProbMapKeys, WSIPatchKeys, optional_import, set_determinism
from tests.test_utils import download_url_or_skip_test, testing_data_config

set_determinism(0)
Expand Down Expand Up @@ -250,6 +250,9 @@ def test_read_patches_large(self, input_parameters, expected):
steps = [round(expected[i]["ratio"] * s) for s in expected[i]["patch_size"]]
expected_location = tuple(expected[i]["step_loc"][j] * steps[j] for j in range(len(steps)))
assert_array_equal(sample["image"].meta[WSIPatchKeys.LOCATION], expected_location)
# `ProbMapProducer` uses these as probability-map indices, which numpy 2.x
# only accepts as integers.
self.assertTrue(np.issubdtype(sample["image"].meta[ProbMapKeys.LOCATION].dtype, np.integer))


@skipUnless(has_cucim, "Requires cucim")
Expand Down
28 changes: 24 additions & 4 deletions tests/handlers/test_handler_prob_map_producer.py
Original file line number Diff line number Diff line change
Expand Up @@ -65,6 +65,21 @@ def __getitem__(self, index):
return {"image": MetaTensor(x=image, meta=metadata), "pred": index + 1}


class FloatLocationDataset(TestDataset):
"""A test fixture whose probability-map locations are floats.

The bundled WSI datasets emit integer locations, so this stands in for a third-party
dataset that does not, which the handler still has to index the probability map with.
"""

__test__ = False # indicate to pytest that this class is not intended for collection

def __init__(self, name, size):
Comment thread
coderabbitai[bot] marked this conversation as resolved.
super().__init__(name, size)
for sample in self.data:
sample[ProbMapKeys.LOCATION.value] = sample[ProbMapKeys.LOCATION.value].astype(float)


class TestEvaluator(Evaluator):
__test__ = False # indicate to pytest that this class is not intended for collection

Expand All @@ -73,10 +88,7 @@ def _iteration(self, engine, batchdata):


class TestHandlerProbMapGenerator(unittest.TestCase):
@parameterized.expand([TEST_CASE_0, TEST_CASE_1, TEST_CASE_2])
def test_prob_map_generator(self, name, size):
# set up dataset
dataset = TestDataset(name, size)
def run_and_check(self, dataset, name, size):
batch_size = 2
data_loader = DataLoader(dataset, batch_size=batch_size)

Expand Down Expand Up @@ -106,6 +118,14 @@ def inference(engine, batch):
self.assertListEqual(np.vstack(prob_map.nonzero()).T.tolist(), [[i, i + 1] for i in range(size)])
self.assertListEqual(prob_map[prob_map.nonzero()].tolist(), [i + 1 for i in range(size)])

@parameterized.expand([TEST_CASE_0, TEST_CASE_1, TEST_CASE_2])
def test_prob_map_generator(self, name, size):
self.run_and_check(TestDataset(name, size), name, size)

@parameterized.expand([TEST_CASE_0, TEST_CASE_1, TEST_CASE_2])
def test_prob_map_generator_float_locations(self, name, size):
self.run_and_check(FloatLocationDataset(name, size), name, size)


if __name__ == "__main__":
unittest.main()