samir-schulz's picture
Fix import: pyRadPlan.ml to pyRadPlan.ai_models
5b483d8 verified
Raw
History Blame Contribute Delete
2.03 kB
"""Preprocessor for the ConvBayes_new dose-calc model (single-channel dose output).
Builds a depth-as-channel ``(batch, 1+depth, H, W)`` input tensor from a CT cuboid and
per-bixel energy, and rescales the raw network output back to physical dose [Gy].
"""
import array_api_compat
import torch
from pyRadPlan.ai_models import BasePreprocessor
from pyRadPlan.core.xp_utils import to_namespace
class ConvDoseSinglePreprocessor(BasePreprocessor):
"""CT+energy -> network input; network output -> physical dose."""
def preprocess(self, inputs: dict) -> torch.Tensor:
"""Assemble the depth-as-channel network input from CT cuboid + energy."""
cfg = self.config
ct_cuboid = inputs["ct_cuboid"]
energy = inputs["energy"]
clip_lo, clip_hi = cfg["ct_clip_range"]
xp = array_api_compat.array_namespace(ct_cuboid)
ct_norm = (xp.clip(ct_cuboid, clip_lo, clip_hi) + cfg["ct_offset"]) / cfg["ct_scale"]
energy_norm = energy / cfg["max_energy"]
t = to_namespace("torch", ct_norm).float().detach()
e_t = to_namespace("torch", energy_norm).float().detach()
if e_t.ndim == 0:
e_t = e_t.view(1)
batch_size = e_t.shape[0]
t = t.expand(batch_size, -1, -1, -1)
height, width = t.shape[-2], t.shape[-1]
energy_map = e_t.view(-1, 1, 1, 1).expand(-1, 1, height, width)
return torch.cat((energy_map, t), dim=1)
def postprocess(self, outputs: torch.Tensor) -> dict:
"""Rescale the normalized network output back to physical dose [Gy]."""
return {"physical_dose": outputs * self.config["max_dose"]}
def predict(self, model: torch.nn.Module, model_input: torch.Tensor) -> dict:
"""Deterministic single forward pass. Assumes ``model_input`` is already preprocessed."""
model_input = model_input.to(next(model.parameters()).device)
with torch.inference_mode():
raw_output = model(model_input)
return self.postprocess(raw_output)