"""Outcome-guided preprocessor shipped with the model repository. Implements the outcome preprocessor contract expected by pyRadPlan's ``OutcomeCNN`` objective, on top of :class:`pyRadPlan.ai_models.BasePreprocessor`: - ``configure(ct, cst, cst_masks=None, device=None)`` resamples CT and masks onto the fixed model input grid once. - ``set_dose_grid(grid)`` precomputes dose<->model sampling coordinates for an arbitrary optimization dose grid. - ``preprocess(dose, requires_grad=False)`` resamples the dose onto the model grid under ``torch.no_grad()``; with ``requires_grad`` the model-grid dose becomes the autograd leaf, so backward only spans the model. Gradient smoothness is expected from the model architecture (e.g. BlurPool3d), not from any explicit smoothing. - ``postprocess(outputs)`` maps the model logit to a probability (sigmoid). - ``gradient_to_dose_grid()`` applies the preprocessing chain rule to ``dose_leaf.grad`` and interpolates it linearly onto the dose grid. All parameters (grid size/spacing, HU window, dose normalization, mask collapsing, input ordering) come from the ``model_preprocessing`` section of ``model_config.json``. ``type_order`` must match the model's forward signature, as the objective calls ``model(*preprocess(dose))``. """ from typing import Any, Optional, Union import logging import numpy as np import SimpleITK as sitk from pyRadPlan.ai_models import BasePreprocessor try: import torch import torch.nn.functional as F # noqa: N812 except ImportError: torch = None # type: ignore F = None # type: ignore logger = logging.getLogger(__name__) class OutcomeCnnPreprocessor(BasePreprocessor): """Differentiable preprocessor mapping planning data onto a fixed model grid. Configuration is read from the ``model_preprocessing`` section of the model's ``model_config.json``: - ``input_dimensions`` : model grid size (X, Y, Z) - ``input_spacing`` : model grid voxel spacing in mm - ``center`` : ``"target"`` centers the model grid on the target center of mass (currently the only supported mode) - ``type_order`` : channel/argument order for :meth:`assemble` (default ``["dose", "ct", "mask"]``) - ``modality.dose`` : ``normalization_value`` (Gy), ``extract`` (mask the dose channel with the structure masks) - ``modality.ct`` : ``window`` HU window (default ``[-1024, 3071]``), mapped linearly to [0, 1] - ``modality.mask`` : ``collapse`` (merge all masks into one channel) Parameters ---------- config : dict, optional The ``model_preprocessing`` dictionary shipped with the model. """ def __init__(self, config: Optional[dict] = None) -> None: if torch is None: raise ImportError( "PyTorch is required for outcome-guided preprocessing. " "Install it e.g. via: pip install torch" ) super().__init__(config) self.input_dimensions: tuple[int, ...] = tuple(self.config["input_dimensions"]) self.input_spacing: tuple[float, ...] = tuple(self.config["input_spacing"]) self.center_mode: str = self.config.get("center", "target") self.type_order: list[str] = list(self.config.get("type_order", ["dose", "ct", "mask"])) modality = self.config.get("modality", {}) self.dose_config: dict = modality.get("dose", {}) self.ct_config: dict = modality.get("ct", {}) self.mask_config: dict = modality.get("mask", {}) self.device: "torch.device" = torch.device("cpu") # Static tensors, set by configure() self._ct_tensor: Optional[torch.Tensor] = None # (1, 1, mZ, mY, mX) self._mask_tensor: Optional[torch.Tensor] = None # (1, C, mZ, mY, mX) # Grid geometry self._model_grid: Optional[dict] = None self._dose_grid: Optional[dict] = None self._coords_dose_to_model: Optional[torch.Tensor] = None self._coords_model_to_dose: Optional[torch.Tensor] = None # Leaf of the last preprocess(requires_grad=True) call self._dose_leaf: Optional[torch.Tensor] = None # ------------------------------------------------------------------ # Configuration # ------------------------------------------------------------------ def configure( self, ct, cst, cst_masks: Optional[list[str]] = None, device: Optional[Union[str, "torch.device"]] = None, ) -> None: """One-time setup of model-grid geometry, CT and mask tensors. Parameters ---------- ct : CT Planning CT (``ct.cube_hu`` is a SimpleITK image). cst : StructureSet Structure set providing ``target_center_of_mass()`` and the VOIs. cst_masks : list[str], optional VOI names used as mask channels (order matters unless the config collapses them). When *None*, all VOIs are used. device : str or torch.device, optional Compute device; should match the model's device. Defaults to CUDA if available, else CPU. """ if device is not None: self.device = torch.device(device) else: self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") if self.center_mode == "target": center = np.asarray(cst.target_center_of_mass(), dtype=np.float64) else: raise ValueError(f"Unknown center mode: {self.center_mode}") ct_grid = self._grid_to_dict(ct.grid) self._model_grid = self._centered_model_grid(center, ct_grid["direction"]) coords_ct_to_model = self._compute_sample_coords(ct_grid, self._model_grid).to(self.device) self._prepare_ct(ct, cst, coords_ct_to_model) self._prepare_masks(cst, cst_masks, coords_ct_to_model) logger.info( "Outcome preprocessor configured: model grid %s @ %s mm, device=%s", self._model_grid["size"], self._model_grid["spacing"], self.device, ) def set_dose_grid(self, dose_grid) -> None: """Precompute dose<->model sampling coordinates for an arbitrary dose grid. Parameters ---------- dose_grid : Grid The grid the optimization dose vector lives on. """ if self._model_grid is None: raise RuntimeError("Preprocessor not configured - call configure() first.") self._dose_grid = self._grid_to_dict(dose_grid) self._coords_dose_to_model = self._compute_sample_coords( self._dose_grid, self._model_grid ).to(self.device) self._coords_model_to_dose = self._compute_sample_coords( self._model_grid, self._dose_grid ).to(self.device) logger.debug( "Dose grid set: %s @ %s mm", self._dose_grid["size"], self._dose_grid["spacing"] ) # ------------------------------------------------------------------ # Per-iteration API # ------------------------------------------------------------------ def preprocess(self, inputs: Any, requires_grad: bool = False) -> Any: """Build the model inputs from a dose array. Parameters ---------- inputs : Array 1-D Fortran-order flat dose array or 3-D (X, Y, Z) array on the dose grid, in any array namespace (numpy, cupy, torch). requires_grad : bool When *True* the model-grid dose becomes an autograd leaf, afterwards accessible as :attr:`dose_leaf`. Returns ------- Any The result of :meth:`assemble` (by default a tuple ordered by ``type_order``, to be passed as ``model(*inputs)``). """ if self._coords_dose_to_model is None: raise RuntimeError("Dose grid not set - call set_dose_grid() first.") with torch.no_grad(): dose_src = self._dose_to_tensor(inputs) dose_5d = dose_src.unsqueeze(0).unsqueeze(0) # (1, 1, dZ, dY, dX) dose_model = F.grid_sample( dose_5d, self._coords_dose_to_model, mode="bilinear", padding_mode="zeros", align_corners=True, ) # (1, 1, mZ, mY, mX) norm_value = self.dose_config.get("normalization_value") if norm_value is not None: dose_model = dose_model / norm_value if self.dose_config.get("extract", False): dose_model = dose_model * self._mask_tensor if requires_grad: dose_model = dose_model.detach().requires_grad_(True) self._dose_leaf = dose_model if requires_grad else None return self.assemble(dose_model, self._ct_tensor, self._mask_tensor) def assemble( self, dose: "torch.Tensor", ct: "torch.Tensor", mask: "torch.Tensor" ) -> tuple["torch.Tensor", ...]: """Arrange the channel tensors into the model's input signature. The default returns a tuple ordered by ``type_order``; the model is then called as ``model(*inputs)``. Override for models expecting e.g. a single channel-stacked tensor. """ channels = {"dose": dose, "ct": ct, "mask": mask} return tuple(channels[key] for key in self.type_order) def postprocess(self, outputs: Any) -> "torch.Tensor": """Map raw model output to a scalar outcome probability. The default assumes the model outputs a logit and applies a sigmoid. Override for models that already output probabilities. """ return torch.sigmoid(outputs).sum() @property def dose_leaf(self) -> "torch.Tensor": """The model-grid dose leaf of the last ``preprocess(requires_grad=True)``.""" if self._dose_leaf is None: raise RuntimeError("No dose leaf - call preprocess(..., requires_grad=True) first.") return self._dose_leaf def gradient_to_dose_grid(self) -> "torch.Tensor": """Map the model-grid gradient back onto the dose grid. Applies the chain rule of the preprocessing (extract mask, normalization) to ``dose_leaf.grad`` and interpolates the result linearly onto the dose grid. Returns ------- torch.Tensor Flat (Fortran-order) float32 gradient of length ``prod(dose grid)``. """ grad_model = self.dose_leaf.grad if grad_model is None: raise RuntimeError("dose_leaf has no gradient - run backward() first.") with torch.no_grad(): if self.dose_config.get("extract", False): grad_model = (grad_model * self._mask_tensor).sum(dim=1, keepdim=True) norm_value = self.dose_config.get("normalization_value") if norm_value is not None: grad_model = grad_model / norm_value grad_dose = F.grid_sample( grad_model, self._coords_model_to_dose, mode="bilinear", padding_mode="zeros", align_corners=True, ) # (1, 1, dZ, dY, dX) # C-order flatten of (Z, Y, X) equals Fortran-order flatten of (X, Y, Z) return grad_dose[0, 0].contiguous().to(dtype=torch.float32).reshape(-1) # ------------------------------------------------------------------ # Introspection (mainly for debugging / tests) # ------------------------------------------------------------------ @property def ct_tensor(self) -> Optional["torch.Tensor"]: """CT tensor on the model grid, (1, 1, mZ, mY, mX).""" return self._ct_tensor @property def mask_tensor(self) -> Optional["torch.Tensor"]: """Mask tensor on the model grid, (1, C, mZ, mY, mX).""" return self._mask_tensor @property def model_grid(self) -> Optional[dict]: """Model grid geometry (size, origin, spacing, direction).""" return self._model_grid # ------------------------------------------------------------------ # Internals # ------------------------------------------------------------------ @staticmethod def _grid_to_dict(grid) -> dict: """Reduce a pyRadPlan Grid to the geometry needed for sampling.""" return { "size": tuple(int(d) for d in grid.dimensions), "origin": tuple(float(v) for v in grid.origin), "spacing": tuple(float(v) for v in grid.resolution_vector), "direction": tuple(float(v) for v in np.asarray(grid.direction).flatten()), } def _centered_model_grid(self, center: np.ndarray, direction: tuple) -> dict: """Model grid with the given center at its geometric center.""" size = self.input_dimensions spacing = self.input_spacing dir_mat = np.asarray(direction, dtype=np.float64).reshape(3, 3) center_idx = np.array([(s - 1) / 2.0 for s in size]) offset = dir_mat @ (np.asarray(spacing) * center_idx) origin = tuple((center - offset).tolist()) return { "size": tuple(size), "origin": origin, "spacing": tuple(spacing), "direction": tuple(direction), } @staticmethod def _compute_sample_coords(src_grid: dict, dst_grid: dict) -> "torch.Tensor": """Compute normalized [-1, 1] coords mapping ``dst_grid`` voxel centers into ``src_grid``. This is the grid ``F.grid_sample`` expects when the sampled tensor lives on ``src_grid``. Returns a tensor of shape (1, dZ, dY, dX, 3). """ dx, dy, dz = dst_grid["size"] ix = torch.arange(dx, dtype=torch.float32) iy = torch.arange(dy, dtype=torch.float32) iz = torch.arange(dz, dtype=torch.float32) gz, gy, gx = torch.meshgrid(iz, iy, ix, indexing="ij") # (dZ, dY, dX) indices = torch.stack( [gx.reshape(-1), gy.reshape(-1), gz.reshape(-1)], dim=1 ) # (N, 3) - x, y, z d_origin = torch.tensor(dst_grid["origin"], dtype=torch.float32) d_spacing = torch.tensor(dst_grid["spacing"], dtype=torch.float32) d_dir = torch.tensor(dst_grid["direction"], dtype=torch.float32).reshape(3, 3) phys = d_origin + (indices * d_spacing) @ d_dir.T # (N, 3) s_origin = torch.tensor(src_grid["origin"], dtype=torch.float32) s_spacing = torch.tensor(src_grid["spacing"], dtype=torch.float32) s_dir = torch.tensor(src_grid["direction"], dtype=torch.float32).reshape(3, 3) s_dir_inv = torch.linalg.inv(s_dir) src_idx = ((phys - s_origin) @ s_dir_inv.T) / s_spacing # (N, 3) s_size = torch.tensor(src_grid["size"], dtype=torch.float32) normalized = 2.0 * src_idx / (s_size - 1) - 1.0 # align_corners=True convention # grid_sample 5-D convention: grid[..., 0]=W(X), grid[..., 1]=H(Y), grid[..., 2]=D(Z) return normalized.reshape(1, int(dz), int(dy), int(dx), 3) def _prepare_ct(self, ct, cst, coords: "torch.Tensor") -> None: """Window/normalize the CT and resample it to the model grid.""" window = self.ct_config.get("window", [-1024, 3071]) lo, hi = float(window[0]), float(window[1]) ct_np = sitk.GetArrayFromImage(ct.cube_hu).astype(np.float32) # (Z, Y, X) ct_np = (np.clip(ct_np, lo, hi) - lo) / (hi - lo) # Zero (= window minimum) outside the body so padding and exterior match body = next((v for v in cst.vois if v.name.upper() == "BODY"), None) if body is not None: body_np = sitk.GetArrayViewFromImage(body.mask).astype(np.float32) ct_np = ct_np * (body_np > 0) ct_t = torch.from_numpy(ct_np).unsqueeze(0).unsqueeze(0).to(self.device) self._ct_tensor = F.grid_sample( ct_t, coords, mode="bilinear", padding_mode="zeros", align_corners=True ) def _prepare_masks(self, cst, cst_masks: Optional[list[str]], coords: "torch.Tensor") -> None: """Resample the requested VOI masks to the model grid (nearest neighbor).""" if cst_masks is not None: vois = [] for name in cst_masks: voi = next((v for v in cst.vois if v.name.lower() == name.lower()), None) if voi is None: available = [v.name for v in cst.vois] raise ValueError(f"VOI '{name}' not found. Available: {available}") vois.append(voi) else: vois = list(cst.vois) channels = [] for voi in vois: mask_np = sitk.GetArrayViewFromImage(voi.mask).astype(np.float32) mask_t = torch.from_numpy(mask_np).unsqueeze(0).unsqueeze(0).to(self.device) channels.append( F.grid_sample( mask_t, coords, mode="nearest", padding_mode="zeros", align_corners=True ) ) mask = torch.cat(channels, dim=1) # (1, C, mZ, mY, mX) if self.mask_config.get("collapse", self.mask_config.get("collaps", False)): mask = mask.amax(dim=1, keepdim=True) self._mask_tensor = mask # `extract` multiplies the single-channel dose by the mask; with more than # one mask channel this would broadcast the dose to C channels and break the # single-dose-channel model. Fail early and clearly instead. if self.dose_config.get("extract", False) and mask.shape[1] > 1: raise ValueError( "dose 'extract' requires a single mask channel; got " f"{mask.shape[1]} channels. Set mask 'collapse': true or pass a single VOI." ) def _dose_to_tensor(self, dose_values) -> "torch.Tensor": """Convert a dose array of any namespace to a (dZ, dY, dX) tensor on device.""" from pyRadPlan.core import xp_utils # noqa: PLC0415 - avoid import cycle at module load # Single device+dtype cast: a cupy/numpy dose is brought onto the model # device and to float32 in one step (dlpack keeps it zero-copy where possible). t = xp_utils.to_namespace(torch, dose_values).to( device=self.device, dtype=torch.float32 ) if t.ndim == 1: dx, dy, dz = self._dose_grid["size"] # Fortran-order reshape to (X, Y, Z) == C-order reshape to (Z, Y, X) t = t.reshape(dz, dy, dx) else: t = t.permute(2, 1, 0) # (X, Y, Z) -> (Z, Y, X) return t.contiguous().detach()