| 128 | """ |
| 129 | |
| 130 | def __init__( |
| 131 | self, |
| 132 | splitter: Splitter | None = None, |
| 133 | merger_cls: type[Merger] | str = AvgMerger, |
| 134 | batch_size: int = 1, |
| 135 | preprocessing: Callable | None = None, |
| 136 | postprocessing: Callable | None = None, |
| 137 | output_keys: Sequence | None = None, |
| 138 | match_spatial_shape: bool = True, |
| 139 | buffer_size: int = 0, |
| 140 | **merger_kwargs: Any, |
| 141 | ) -> None: |
| 142 | Inferer.__init__(self) |
| 143 | # splitter |
| 144 | if not isinstance(splitter, (Splitter, type(None))): |
| 145 | if not isinstance(splitter, Splitter): |
| 146 | raise TypeError( |
| 147 | f"'splitter' should be a `Splitter` object that returns: " |
| 148 | "an iterable of pairs of (patch, location) or a MetaTensor that has `PatchKeys.LOCATION` metadata)." |
| 149 | f"{type(splitter)} is given." |
| 150 | ) |
| 151 | self.splitter = splitter |
| 152 | |
| 153 | # merger |
| 154 | if isinstance(merger_cls, str): |
| 155 | valid_merger_cls: type[Merger] |
| 156 | # search amongst implemented mergers in MONAI |
| 157 | valid_merger_cls, merger_found = optional_import("monai.inferers.merger", name=merger_cls) |
| 158 | if not merger_found: |
| 159 | # try to locate the requested merger class (with dotted path) |
| 160 | valid_merger_cls = locate(merger_cls) # type: ignore |
| 161 | if valid_merger_cls is None: |
| 162 | raise ValueError(f"The requested `merger_cls` ['{merger_cls}'] does not exist.") |
| 163 | merger_cls = valid_merger_cls |
| 164 | if not issubclass(merger_cls, Merger): |
| 165 | raise TypeError(f"'merger' should be a subclass of `Merger`, {merger_cls} is given.") |
| 166 | self.merger_cls = merger_cls |
| 167 | self.merger_kwargs = merger_kwargs |
| 168 | |
| 169 | # pre-processor (process patch before the network) |
| 170 | if preprocessing is not None and not callable(preprocessing): |
| 171 | raise TypeError(f"'preprocessing' should be a callable object, {type(preprocessing)} is given.") |
| 172 | self.preprocessing = preprocessing |
| 173 | |
| 174 | # post-processor (process the output of the network) |
| 175 | if postprocessing is not None and not callable(postprocessing): |
| 176 | raise TypeError(f"'postprocessing' should be a callable object, {type(postprocessing)} is given.") |
| 177 | self.postprocessing = postprocessing |
| 178 | |
| 179 | # batch size for patches |
| 180 | if batch_size < 1: |
| 181 | raise ValueError(f"`batch_size` must be a positive number, {batch_size} is given.") |
| 182 | self.batch_size = batch_size |
| 183 | |
| 184 | # model output keys |
| 185 | self.output_keys = output_keys |
| 186 | |
| 187 | # whether to crop the output to match the input shape |