(
self,
transform: InvertibleTransform,
batch_size: int,
num_workers: int = 0,
inferrer_fn: Callable = _identity,
device: str | torch.device = "cpu",
image_key=CommonKeys.IMAGE,
orig_key=CommonKeys.LABEL,
nearest_interp: bool = True,
orig_meta_keys: str | None = None,
meta_key_postfix=DEFAULT_POST_FIX,
to_tensor: bool = True,
output_device: str | torch.device = "cpu",
post_func: Callable = _identity,
return_full_data: bool = False,
progress: bool = True,
apply_inverse_to_pred: bool = True,
)
| 113 | __test__ = False # indicate to pytest that this class is not intended for collection |
| 114 | |
| 115 | def __init__( |
| 116 | self, |
| 117 | transform: InvertibleTransform, |
| 118 | batch_size: int, |
| 119 | num_workers: int = 0, |
| 120 | inferrer_fn: Callable = _identity, |
| 121 | device: str | torch.device = "cpu", |
| 122 | image_key=CommonKeys.IMAGE, |
| 123 | orig_key=CommonKeys.LABEL, |
| 124 | nearest_interp: bool = True, |
| 125 | orig_meta_keys: str | None = None, |
| 126 | meta_key_postfix=DEFAULT_POST_FIX, |
| 127 | to_tensor: bool = True, |
| 128 | output_device: str | torch.device = "cpu", |
| 129 | post_func: Callable = _identity, |
| 130 | return_full_data: bool = False, |
| 131 | progress: bool = True, |
| 132 | apply_inverse_to_pred: bool = True, |
| 133 | ) -> None: |
| 134 | self.transform = transform |
| 135 | self.batch_size = batch_size |
| 136 | self.num_workers = num_workers |
| 137 | self.inferrer_fn = inferrer_fn |
| 138 | self.device = device |
| 139 | self.image_key = image_key |
| 140 | self.return_full_data = return_full_data |
| 141 | self.progress = progress |
| 142 | self.apply_inverse_to_pred = apply_inverse_to_pred |
| 143 | self._pred_key = CommonKeys.PRED |
| 144 | self.inverter = Invertd( |
| 145 | keys=self._pred_key, |
| 146 | transform=transform, |
| 147 | orig_keys=orig_key, |
| 148 | orig_meta_keys=orig_meta_keys, |
| 149 | meta_key_postfix=meta_key_postfix, |
| 150 | nearest_interp=nearest_interp, |
| 151 | to_tensor=to_tensor, |
| 152 | device=output_device, |
| 153 | post_func=post_func, |
| 154 | ) |
| 155 | |
| 156 | # check that the transform has at least one random component, and that all random transforms are invertible |
| 157 | self._check_transforms() |
| 158 | |
| 159 | def _check_transforms(self): |
| 160 | """Should be at least 1 random transform, and all random transforms should be invertible.""" |
nothing calls this directly
no test coverage detected