"""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)