File size: 18,619 Bytes
012f65b c246f4d 012f65b c20d956 012f65b c20d956 9eb90c4 c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 012f65b c20d956 c246f4d c20d956 012f65b c246f4d c20d956 012f65b c20d956 012f65b c246f4d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 | """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()
|