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