python/dpt/transport/inverse.py

Generated from the full canonical file for this source snapshot. Line numbers match the library source.

Source SHA256: 808100ab4db1e9a17a5a8343126ce00b702dcb81569ba5ce9ca4364be62b3d1e

1"""Concrete CUDA oracle for squared-expected-signal recovery in a fixed chart.23The controller sees only small parameter/gradient vectors and scalar objective4checkpoints. Source histories, transport, detector sums, independent estimator5products and objective reductions remain on one CUDA stream. No Python history6loop, full image download or per-replicate device allocation is used.7"""89# Internal workspaces are composed by this package, never exposed as a backend API.10# pyright: reportPrivateUsage=false11from __future__ import annotations1213import math14from dataclasses import dataclass, field15from typing import Any1617from dpt._runtime import load_kernels, require_no_tape18from dpt.contracts import NumericalError, TrialDomainError, finite_scalar, integer19from dpt.objectives import (20    ObjectiveSpec,21    ObjectiveWorkspace,22    prepare_objective,23    reduce_objective_components,24)25from dpt.registration import Vector2627from .derivatives import TransportParameter, derivative_histories28from .estimators import (29    EstimatorWorkspace,30    history_mean,31    independent_squared_gradient,32    independent_squared_loss,33    prepare_estimators,34)35from .forward import TransportWorkspace, trace_histories36from .model import TransportError37from .rng import HistoryBatch, require_independent38from .source import ParallelBeam as ParallelBeam39from .source import sample_parallel_beam404142@dataclass(slots=True)43class TransportSquaredOracle:44    """Prepared implementation of `IndependentSquaredOracle`.4546    Chart coordinates are absolute log density for each selected material and47    log source amplitude for a selected `log-source-amplitude` parameter.48    Its chain factor is formed in the per-history score before reduction. Unselected49    densities use `base_density`; an unselected amplitude uses `fixed_amplitude`.50    This positive chart excludes exact zero amplitude from optimisation; the51    lower-level derivative operator still supports its one-sided zero boundary.5253    Capacity must cover the controller's maximum batch. Instances are mutable54    single-stream workspaces and must not be shared concurrently. Arrays derived55    from a batch are overwritten by the next call. Observations and pixel weights56    remain immutable for this oracle's lifetime. `histories_traced` counts actual57    forward/replay work, unlike the controller's unique-random-history budget.58    """5960    workspace: TransportWorkspace61    source: ParallelBeam62    parameters: tuple[TransportParameter, ...]63    observation: Any64    pixel_weights: Any65    fixed_amplitude: float66    _moments: EstimatorWorkspace = field(repr=False)67    _objective: ObjectiveWorkspace = field(repr=False)68    _arrays: dict[str, Any] = field(repr=False)69    _parameter_host: Any = field(repr=False)70    _parameter_view: Any = field(repr=False)71    _scalar_host: Any = field(repr=False)72    _scalar_view: Any = field(repr=False)73    _model_kernels: Any = field(default=None, repr=False)74    _model_partials: tuple[Any, ...] = field(default=(), repr=False)75    _model_host: Any = field(default=None, repr=False)76    _model_view: Any = field(default=None, repr=False)77    _deterministic_sampling: bool = field(default=False, repr=False)78    _last_parameters: Vector | None = field(default=None, repr=False)79    histories_traced: int = 080    parameter_upload_bytes: int = 081    scalar_download_bytes: int = 08283    @property84    def deterministic_sampling(self) -> bool:85        """Preparation established zero scattering and a fixed ray source."""86        return self._deterministic_sampling8788    @property89    def allocated_bytes(self) -> int:90        """Owned device bytes, excluding shared transport/static input buffers."""91        wp = self.workspace.context.wp92        return (93            sum(94                int(value.size) * wp.types.type_size_in_bytes(value.dtype)95                for value in self._arrays.values()96            )97            + sum(int(value.size) * 8 for value in self._model_partials)98            + self._moments.scratch_bytes99            + self._objective.scratch_bytes100        )101102    def _chart(self, values: Vector) -> float:103        if len(values) != len(self.parameters):104            raise TransportError("parameter vector does not match the prepared transport chart")105        values = tuple(finite_scalar(value, "chart coordinate") for value in values)106        amplitude = self.fixed_amplitude107        for index, parameter in enumerate(self.parameters):108            try:109                physical = math.exp(values[index])110            except OverflowError as error:111                raise TrialDomainError(112                    "log parameter overflows its physical representation"113                ) from error114            if not math.isfinite(physical) or physical <= 0:115                raise TrialDomainError(116                    "log parameter has no positive finite binary64 representation"117                )118            if parameter.kind == "log-source-amplitude":119                amplitude = physical120        if values != self._last_parameters:121            for index, value in enumerate(values):122                self._parameter_view[index] = value123            context = self.workspace.context124            context.wp.copy(self._arrays["chart"], self._parameter_host, stream=context.stream)125            self.workspace._launch(126                self.workspace._kernels.update_density_chart,127                self.workspace.spec.grid.materials,128                [129                    self._arrays["chart"],130                    self._arrays["material_parameter"],131                    self._arrays["base_density"],132                    self._arrays["density"],133                    self.workspace._status,134                ],135            )136            self.workspace.check_status()137            self.parameter_upload_bytes += 8 * len(values)138            self._last_parameters = values139        return amplitude140141    def _source_arrays(self, batch: HistoryBatch) -> list[Any]:142        if batch.count > self.workspace.max_histories:143            raise TransportError("oracle batch exceeds prepared capacity")144        arrays = [self._arrays[name][: batch.count] for name in ("position", "direction", "weight")]145        sample_parallel_beam(146            self.source,147            batch=batch,148            workspace=self.workspace,149            out_position=arrays[0],150            out_direction=arrays[1],151            out_weight=arrays[2],152            stream=self.workspace.context.stream,153        )154        return [*arrays, self._arrays["density"]]155156    def _mean(self, batch: HistoryBatch, amplitude: float, destination: Any) -> None:157        inputs = self._source_arrays(batch)158        outputs = {159            f"out_{name}": self._arrays[name][: batch.count]160            for name in ("pixel", "score", "energy", "events", "status")161        }162        trace_histories(163            *inputs,164            batch=batch,165            workspace=self.workspace,166            source_amplitude=amplitude,167            stream=self.workspace.context.stream,168            validate=False,169            **outputs,170        )171        self.histories_traced += batch.count172        history_mean(173            outputs["out_pixel"],174            outputs["out_score"],175            outputs["out_status"],176            batch=batch,177            workspace=self._moments,178            out_mean=destination,179            stream=self.workspace.context.stream,180            validate=False,181        )182183    def _scalar(self) -> float:184        context = self.workspace.context185        reduce_objective_components(186            self._arrays["components"],187            out_loss=self._arrays["scalar"],188            workspace=self._objective,189            stream=context.stream,190            validate=False,191        )192        context.wp.copy(self._scalar_host, self._arrays["scalar"], stream=context.stream)193        # Check both producers before accepting the scalar; a failed transport194        # kernel may still leave finite zeros that a loss kernel cannot diagnose.195        self.workspace.check_status()196        self._objective.check_status()197        self.scalar_download_bytes += 8198        return float(self._scalar_view[0])199200    def gradient_replicate(201        self,202        parameters: Vector,203        mean_batch: HistoryBatch,204        derivative_batch: HistoryBatch,205    ) -> Vector:206        try:207            return self._gradient_replicate(parameters, mean_batch, derivative_batch)208        finally:209            # An exception after asynchronous H2D upload must not release pinned210            # staging for the next call while CUDA still reads the previous values.211            self.workspace.context.wp.synchronize_stream(self.workspace.context.stream)212213    def change_replicate(214        self,215        before: Vector,216        after: Vector,217        first: HistoryBatch,218        second: HistoryBatch,219    ) -> float:220        try:221            return self._change_replicate(before, after, first, second)222        finally:223            self.workspace.context.wp.synchronize_stream(self.workspace.context.stream)224225    # region book:transport-inverse-oracle226    def _gradient_replicate(227        self,228        parameters: Vector,229        mean_batch: HistoryBatch,230        derivative_batch: HistoryBatch,231        capture_model: bool = False,232    ) -> Vector:233        """Use disjoint source/transport samples for the two nonlinear factors."""234        require_no_tape()235        require_independent(mean_batch, derivative_batch)236        amplitude = self._chart(parameters)237        self._mean(mean_batch, amplitude, self._arrays["mean_a"])238        inputs = self._source_arrays(derivative_batch)239        pixel = self._arrays["pixel"][: derivative_batch.count]240        derivative = self._arrays["score"][: derivative_batch.count]241        status = self._arrays["status"][: derivative_batch.count]242        result: list[float] = []243        for column, parameter in enumerate(self.parameters):244            derivative_histories(245                *inputs,246                parameter=parameter,247                batch=derivative_batch,248                workspace=self.workspace,249                out_pixel=pixel,250                out_derivative=derivative,251                out_status=status,252                source_amplitude=amplitude,253                stream=self.workspace.context.stream,254                validate=False,255            )256            self.histories_traced += derivative_batch.count257            history_mean(258                pixel,259                derivative,260                status,261                batch=derivative_batch,262                workspace=self._moments,263                out_mean=self._arrays["mean_b"],264                stream=self.workspace.context.stream,265                validate=False,266            )267            if capture_model:268                self.workspace._launch(269                    self._model_kernels.store_column,270                    self.workspace.spec.detector.pixels,271                    [272                        self._arrays["mean_b"],273                        column,274                        self.workspace.spec.detector.pixels,275                        self._arrays["jacobian"],276                    ],277                )278            independent_squared_gradient(279                self._arrays["mean_a"],280                self._arrays["mean_b"],281                self.observation,282                self.pixel_weights,283                batches=(mean_batch, derivative_batch),284                workspace=self.workspace,285                out_components=self._arrays["components"],286                stream=self.workspace.context.stream,287                validate=False,288            )289            value = self._scalar()290            if not math.isfinite(value):291                raise NumericalError("inverse chart derivative overflow")292            result.append(value)293        return tuple(result)294295    def model_replicate(296        self,297        parameters: Vector,298        mean_batch: HistoryBatch,299        derivative_batch: HistoryBatch,300    ) -> tuple[Vector, tuple[Vector, ...]]:301        """Return an unbiased gradient and a PSD *proposal* metric, not an unbiased Hessian.302303        Jacobian columns and all pixel contractions stay on CUDA. Only the small304        dense metric crosses to prepared pinned staging. Sampling variance biases305        its diagonal upwards; fresh independent acceptance decides whether to move.306        """307        if not self._model_partials:308            raise TransportError("prepare the inverse oracle with local_model=True")309        context = self.workspace.context310        wp = context.wp311        try:312            gradient = self._gradient_replicate(parameters, mean_batch, derivative_batch, True)313            dimension = len(self.parameters)314            pixels = self.workspace.spec.detector.pixels315            count = (pixels + 255) // 256316            wp.launch_tiled(317                self._model_kernels.gram_tiles,318                dim=dimension * dimension * count,319                block_dim=256,320                inputs=[321                    self._arrays["jacobian"],322                    self.pixel_weights,323                    pixels,324                    dimension,325                    count,326                    self._model_partials[0],327                    self.workspace._status,328                ],329                device=context.device,330                stream=context.stream,331                record_tape=False,332            )333            previous = self._model_partials[0]334            for destination in self._model_partials[1:]:335                next_count = (count + 255) // 256336                wp.launch_tiled(337                    self._model_kernels.sum_gram_tiles,338                    dim=dimension * dimension * next_count,339                    block_dim=256,340                    inputs=[previous, count, next_count, destination],341                    device=context.device,342                    stream=context.stream,343                    record_tape=False,344                )345                previous, count = destination, next_count346            wp.copy(self._model_host, previous, stream=context.stream)347            self.workspace.check_status()348            self.scalar_download_bytes += 8 * dimension * dimension349            curvature = tuple(350                tuple(float(self._model_view[i * dimension + j]) for j in range(dimension))351                for i in range(dimension)352            )353            if not all(math.isfinite(x) for row in curvature for x in row):354                raise NumericalError("proposal curvature exceeds the finite chart range")355            return gradient, curvature356        finally:357            wp.synchronize_stream(context.stream)358359    def _change_replicate(360        self,361        before: Vector,362        after: Vector,363        first: HistoryBatch,364        second: HistoryBatch,365    ) -> float:366        """Independent product factors; common random numbers across parameter points."""367        require_no_tape()368        require_independent(first, second)369        losses: list[float] = []370        for parameters in (before, after):371            amplitude = self._chart(parameters)372            self._mean(first, amplitude, self._arrays["mean_a"])373            self._mean(second, amplitude, self._arrays["mean_b"])374            independent_squared_loss(375                self._arrays["mean_a"],376                self._arrays["mean_b"],377                self.observation,378                self.pixel_weights,379                batches=(first, second),380                workspace=self.workspace,381                out_components=self._arrays["components"],382                stream=self.workspace.context.stream,383                validate=False,384            )385            losses.append(self._scalar())386        difference = losses[1] - losses[0]387        if not math.isfinite(difference):388            raise NumericalError("objective-change estimate overflow")389        return difference390391    # endregion book:transport-inverse-oracle392393394def prepare_transport_inverse(395    workspace: TransportWorkspace,396    *,397    source: ParallelBeam,398    parameters: tuple[TransportParameter, ...],399    observation: Any,400    pixel_weights: Any,401    base_density: tuple[float, ...],402    fixed_amplitude: float = 1.0,403    local_model: bool = False,404    local_model_max_bytes: int = 256 * 1024 * 1024,405) -> TransportSquaredOracle:406    """Allocate a complete reusable inverse oracle; no physics execution is implied.407408    The preparation call binds immutable device observations/weights and uploads409    only fixed density values and the material-to-parameter map. Dynamic small410    chart uploads and scalar downloads are counted explicitly on the oracle.411    """412    require_no_tape()413    if type(local_model) is not bool:414        raise TransportError("local_model must be a boolean preparation choice")415    integer(local_model_max_bytes, "local_model_max_bytes", minimum=1)416    parameters = tuple(parameters)417    base_density = tuple(base_density)418    if any(parameter.kind == "source-amplitude" for parameter in parameters):419        raise TransportError(420            "the inverse chart requires log-source-amplitude, not direct amplitude"421        )422    if not parameters or len(set(parameters)) != len(parameters):423        raise TransportError("select at least one distinct supported transport parameter")424    if len(base_density) != workspace.spec.grid.materials:425        raise TransportError("base_density must contain one scale per material")426    if any(finite_scalar(value, "base density", minimum=0.0) == 0 for value in base_density):427        raise TransportError("base density must be positive")428    finite_scalar(fixed_amplitude, "fixed source amplitude", minimum=0.0)429    if source.lower_mm[2] >= workspace.spec.detector.z_mm:430        raise TransportError("source plane must be below the detector plane")431    context = workspace.context432    wp = context.wp433    pixels = workspace.spec.detector.pixels434    dimension = len(parameters)435    partial_sizes: list[int] = []436    if local_model:437        if dimension > 16:438            raise TransportError("dense local models support at most 16 active parameters")439        count = (pixels + 255) // 256440        while True:441            partial_sizes.append(dimension * dimension * count)442            if count == 1:443                break444            count = (count + 255) // 256445        if dimension * pixels >= 2**31 or any(size >= 2**31 for size in partial_sizes):446            raise TransportError("local model exceeds signed 32-bit indexing")447        required_bytes = 8 * (dimension * pixels + sum(partial_sizes) + dimension * dimension)448        if required_bytes > local_model_max_bytes:449            raise TransportError(450                f"local model needs {required_bytes} bytes, exceeding preparation budget"451            )452    for name, value in (("observation", observation), ("pixel_weights", pixel_weights)):453        context.array(value, name, dtype=wp.float64, shape=(pixels,))454    workspace._launch(455        workspace._kernels.validate_measurement,456        pixels,457        [observation, pixel_weights, workspace._status],458    )459    workspace.check_status()460    mapping = [-1] * workspace.spec.grid.materials461    for index, parameter in enumerate(parameters):462        if parameter.material is not None:463            if parameter.material >= len(mapping):464                raise TransportError("active material is outside this model")465            mapping[parameter.material] = index466    arrays: dict[str, Any] = {}467    with context.scope():468        for name, dtype in (469            ("position", wp.vec3d),470            ("direction", wp.vec3d),471            ("weight", wp.float64),472            ("pixel", wp.int32),473            ("score", wp.float64),474            ("energy", wp.float64),475            ("events", wp.int32),476            ("status", wp.int32),477        ):478            arrays[name] = wp.empty(workspace.max_histories, dtype=dtype, device=context.device)479        for name in ("mean_a", "mean_b", "components"):480            arrays[name] = wp.empty(pixels, dtype=wp.float64, device=context.device)481        arrays["scalar"] = wp.empty(1, dtype=wp.float64, device=context.device)482        arrays["density"] = wp.empty(len(base_density), dtype=wp.float64, device=context.device)483        arrays["base_density"] = wp.array(484            list(base_density), dtype=wp.float64, device=context.device485        )486        arrays["material_parameter"] = wp.array(mapping, dtype=wp.int32, device=context.device)487        arrays["chart"] = wp.empty(len(parameters), dtype=wp.float64, device=context.device)488        parameter_host = wp.empty(len(parameters), dtype=wp.float64, device="cpu", pinned=True)489        scalar_host = wp.empty(1, dtype=wp.float64, device="cpu", pinned=True)490    model_kernels = None491    model_partials: tuple[Any, ...] = ()492    model_host = None493    model_view = None494    deterministic = False495    if local_model or workspace.spec.estimator == "continuous-absorption":496        model_kernels = load_kernels("dpt.transport.recovery_kernels")497    with context.scope():498        if local_model:499            arrays["jacobian"] = wp.empty(500                dimension * pixels, dtype=wp.float64, device=context.device501            )502            model_partials = tuple(503                wp.empty(n, dtype=wp.float64, device=context.device) for n in partial_sizes504            )505            model_host = wp.empty(506                dimension * dimension, dtype=wp.float64, device="cpu", pinned=True507            )508            model_view = model_host.numpy()509        if workspace.spec.estimator == "continuous-absorption" and source.extent_xy_mm == (510            0.0,511            0.0,512        ):513            assert model_kernels is not None514            flag = wp.zeros(1, dtype=wp.int32, device=context.device)515            workspace._launch(516                model_kernels.flag_scattering,517                int(workspace.scattering.size),518                [workspace.scattering, flag],519            )520            wp.synchronize_stream(context.stream)521            deterministic = int(flag.numpy()[0]) == 0522    moments = prepare_estimators(workspace)523    objective = prepare_objective(524        ObjectiveSpec(), max_pixels=pixels, device=context.device, stream=context.stream525    )526    return TransportSquaredOracle(527        workspace,528        source,529        parameters,530        observation,531        pixel_weights,532        fixed_amplitude,533        moments,534        objective,535        arrays,536        parameter_host,537        parameter_host.numpy(),538        scalar_host,539        scalar_host.numpy(),540        _model_kernels=model_kernels,541        _model_partials=model_partials,542        _model_host=model_host,543        _model_view=model_view,544        _deterministic_sampling=deterministic,545    )546