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 (20ObjectiveSpec,21ObjectiveWorkspace,22prepare_objective,23reduce_objective_components,24)25from dpt.registration import Vector2627from .derivatives import TransportParameter, derivative_histories28from .estimators import (29EstimatorWorkspace,30history_mean,31independent_squared_gradient,32independent_squared_loss,33prepare_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`.4546Chart coordinates are absolute log density for each selected material and47log source amplitude for a selected `log-source-amplitude` parameter.48Its chain factor is formed in the per-history score before reduction. Unselected49densities use `base_density`; an unselected amplitude uses `fixed_amplitude`.50This positive chart excludes exact zero amplitude from optimisation; the51lower-level derivative operator still supports its one-sided zero boundary.5253Capacity must cover the controller's maximum batch. Instances are mutable54single-stream workspaces and must not be shared concurrently. Arrays derived55from a batch are overwritten by the next call. Observations and pixel weights56remain immutable for this oracle's lifetime. `histories_traced` counts actual57forward/replay work, unlike the controller's unique-random-history budget.58"""5960workspace: TransportWorkspace61source: ParallelBeam62parameters: tuple[TransportParameter, ...]63observation: Any64pixel_weights: Any65fixed_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)79histories_traced: int = 080parameter_upload_bytes: int = 081scalar_download_bytes: int = 08283@property84def deterministic_sampling(self) -> bool:85"""Preparation established zero scattering and a fixed ray source."""86return self._deterministic_sampling8788@property89def allocated_bytes(self) -> int:90"""Owned device bytes, excluding shared transport/static input buffers."""91wp = self.workspace.context.wp92return (93sum(94int(value.size) * wp.types.type_size_in_bytes(value.dtype)95for 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)101102def _chart(self, values: Vector) -> float:103if len(values) != len(self.parameters):104raise TransportError("parameter vector does not match the prepared transport chart")105values = tuple(finite_scalar(value, "chart coordinate") for value in values)106amplitude = self.fixed_amplitude107for index, parameter in enumerate(self.parameters):108try:109physical = math.exp(values[index])110except OverflowError as error:111raise TrialDomainError(112"log parameter overflows its physical representation"113) from error114if not math.isfinite(physical) or physical <= 0:115raise TrialDomainError(116"log parameter has no positive finite binary64 representation"117)118if parameter.kind == "log-source-amplitude":119amplitude = physical120if values != self._last_parameters:121for index, value in enumerate(values):122self._parameter_view[index] = value123context = self.workspace.context124context.wp.copy(self._arrays["chart"], self._parameter_host, stream=context.stream)125self.workspace._launch(126self.workspace._kernels.update_density_chart,127self.workspace.spec.grid.materials,128[129self._arrays["chart"],130self._arrays["material_parameter"],131self._arrays["base_density"],132self._arrays["density"],133self.workspace._status,134],135)136self.workspace.check_status()137self.parameter_upload_bytes += 8 * len(values)138self._last_parameters = values139return amplitude140141def _source_arrays(self, batch: HistoryBatch) -> list[Any]:142if batch.count > self.workspace.max_histories:143raise TransportError("oracle batch exceeds prepared capacity")144arrays = [self._arrays[name][: batch.count] for name in ("position", "direction", "weight")]145sample_parallel_beam(146self.source,147batch=batch,148workspace=self.workspace,149out_position=arrays[0],150out_direction=arrays[1],151out_weight=arrays[2],152stream=self.workspace.context.stream,153)154return [*arrays, self._arrays["density"]]155156def _mean(self, batch: HistoryBatch, amplitude: float, destination: Any) -> None:157inputs = self._source_arrays(batch)158outputs = {159f"out_{name}": self._arrays[name][: batch.count]160for name in ("pixel", "score", "energy", "events", "status")161}162trace_histories(163*inputs,164batch=batch,165workspace=self.workspace,166source_amplitude=amplitude,167stream=self.workspace.context.stream,168validate=False,169**outputs,170)171self.histories_traced += batch.count172history_mean(173outputs["out_pixel"],174outputs["out_score"],175outputs["out_status"],176batch=batch,177workspace=self._moments,178out_mean=destination,179stream=self.workspace.context.stream,180validate=False,181)182183def _scalar(self) -> float:184context = self.workspace.context185reduce_objective_components(186self._arrays["components"],187out_loss=self._arrays["scalar"],188workspace=self._objective,189stream=context.stream,190validate=False,191)192context.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.195self.workspace.check_status()196self._objective.check_status()197self.scalar_download_bytes += 8198return float(self._scalar_view[0])199200def gradient_replicate(201self,202parameters: Vector,203mean_batch: HistoryBatch,204derivative_batch: HistoryBatch,205) -> Vector:206try:207return self._gradient_replicate(parameters, mean_batch, derivative_batch)208finally:209# An exception after asynchronous H2D upload must not release pinned210# staging for the next call while CUDA still reads the previous values.211self.workspace.context.wp.synchronize_stream(self.workspace.context.stream)212213def change_replicate(214self,215before: Vector,216after: Vector,217first: HistoryBatch,218second: HistoryBatch,219) -> float:220try:221return self._change_replicate(before, after, first, second)222finally:223self.workspace.context.wp.synchronize_stream(self.workspace.context.stream)224225# region book:transport-inverse-oracle226def _gradient_replicate(227self,228parameters: Vector,229mean_batch: HistoryBatch,230derivative_batch: HistoryBatch,231capture_model: bool = False,232) -> Vector:233"""Use disjoint source/transport samples for the two nonlinear factors."""234require_no_tape()235require_independent(mean_batch, derivative_batch)236amplitude = self._chart(parameters)237self._mean(mean_batch, amplitude, self._arrays["mean_a"])238inputs = self._source_arrays(derivative_batch)239pixel = self._arrays["pixel"][: derivative_batch.count]240derivative = self._arrays["score"][: derivative_batch.count]241status = self._arrays["status"][: derivative_batch.count]242result: list[float] = []243for column, parameter in enumerate(self.parameters):244derivative_histories(245*inputs,246parameter=parameter,247batch=derivative_batch,248workspace=self.workspace,249out_pixel=pixel,250out_derivative=derivative,251out_status=status,252source_amplitude=amplitude,253stream=self.workspace.context.stream,254validate=False,255)256self.histories_traced += derivative_batch.count257history_mean(258pixel,259derivative,260status,261batch=derivative_batch,262workspace=self._moments,263out_mean=self._arrays["mean_b"],264stream=self.workspace.context.stream,265validate=False,266)267if capture_model:268self.workspace._launch(269self._model_kernels.store_column,270self.workspace.spec.detector.pixels,271[272self._arrays["mean_b"],273column,274self.workspace.spec.detector.pixels,275self._arrays["jacobian"],276],277)278independent_squared_gradient(279self._arrays["mean_a"],280self._arrays["mean_b"],281self.observation,282self.pixel_weights,283batches=(mean_batch, derivative_batch),284workspace=self.workspace,285out_components=self._arrays["components"],286stream=self.workspace.context.stream,287validate=False,288)289value = self._scalar()290if not math.isfinite(value):291raise NumericalError("inverse chart derivative overflow")292result.append(value)293return tuple(result)294295def model_replicate(296self,297parameters: Vector,298mean_batch: HistoryBatch,299derivative_batch: HistoryBatch,300) -> tuple[Vector, tuple[Vector, ...]]:301"""Return an unbiased gradient and a PSD *proposal* metric, not an unbiased Hessian.302303Jacobian columns and all pixel contractions stay on CUDA. Only the small304dense metric crosses to prepared pinned staging. Sampling variance biases305its diagonal upwards; fresh independent acceptance decides whether to move.306"""307if not self._model_partials:308raise TransportError("prepare the inverse oracle with local_model=True")309context = self.workspace.context310wp = context.wp311try:312gradient = self._gradient_replicate(parameters, mean_batch, derivative_batch, True)313dimension = len(self.parameters)314pixels = self.workspace.spec.detector.pixels315count = (pixels + 255) // 256316wp.launch_tiled(317self._model_kernels.gram_tiles,318dim=dimension * dimension * count,319block_dim=256,320inputs=[321self._arrays["jacobian"],322self.pixel_weights,323pixels,324dimension,325count,326self._model_partials[0],327self.workspace._status,328],329device=context.device,330stream=context.stream,331record_tape=False,332)333previous = self._model_partials[0]334for destination in self._model_partials[1:]:335next_count = (count + 255) // 256336wp.launch_tiled(337self._model_kernels.sum_gram_tiles,338dim=dimension * dimension * next_count,339block_dim=256,340inputs=[previous, count, next_count, destination],341device=context.device,342stream=context.stream,343record_tape=False,344)345previous, count = destination, next_count346wp.copy(self._model_host, previous, stream=context.stream)347self.workspace.check_status()348self.scalar_download_bytes += 8 * dimension * dimension349curvature = tuple(350tuple(float(self._model_view[i * dimension + j]) for j in range(dimension))351for i in range(dimension)352)353if not all(math.isfinite(x) for row in curvature for x in row):354raise NumericalError("proposal curvature exceeds the finite chart range")355return gradient, curvature356finally:357wp.synchronize_stream(context.stream)358359def _change_replicate(360self,361before: Vector,362after: Vector,363first: HistoryBatch,364second: HistoryBatch,365) -> float:366"""Independent product factors; common random numbers across parameter points."""367require_no_tape()368require_independent(first, second)369losses: list[float] = []370for parameters in (before, after):371amplitude = self._chart(parameters)372self._mean(first, amplitude, self._arrays["mean_a"])373self._mean(second, amplitude, self._arrays["mean_b"])374independent_squared_loss(375self._arrays["mean_a"],376self._arrays["mean_b"],377self.observation,378self.pixel_weights,379batches=(first, second),380workspace=self.workspace,381out_components=self._arrays["components"],382stream=self.workspace.context.stream,383validate=False,384)385losses.append(self._scalar())386difference = losses[1] - losses[0]387if not math.isfinite(difference):388raise NumericalError("objective-change estimate overflow")389return difference390391# endregion book:transport-inverse-oracle392393394def prepare_transport_inverse(395workspace: TransportWorkspace,396*,397source: ParallelBeam,398parameters: tuple[TransportParameter, ...],399observation: Any,400pixel_weights: Any,401base_density: tuple[float, ...],402fixed_amplitude: float = 1.0,403local_model: bool = False,404local_model_max_bytes: int = 256 * 1024 * 1024,405) -> TransportSquaredOracle:406"""Allocate a complete reusable inverse oracle; no physics execution is implied.407408The preparation call binds immutable device observations/weights and uploads409only fixed density values and the material-to-parameter map. Dynamic small410chart uploads and scalar downloads are counted explicitly on the oracle.411"""412require_no_tape()413if type(local_model) is not bool:414raise TransportError("local_model must be a boolean preparation choice")415integer(local_model_max_bytes, "local_model_max_bytes", minimum=1)416parameters = tuple(parameters)417base_density = tuple(base_density)418if any(parameter.kind == "source-amplitude" for parameter in parameters):419raise TransportError(420"the inverse chart requires log-source-amplitude, not direct amplitude"421)422if not parameters or len(set(parameters)) != len(parameters):423raise TransportError("select at least one distinct supported transport parameter")424if len(base_density) != workspace.spec.grid.materials:425raise TransportError("base_density must contain one scale per material")426if any(finite_scalar(value, "base density", minimum=0.0) == 0 for value in base_density):427raise TransportError("base density must be positive")428finite_scalar(fixed_amplitude, "fixed source amplitude", minimum=0.0)429if source.lower_mm[2] >= workspace.spec.detector.z_mm:430raise TransportError("source plane must be below the detector plane")431context = workspace.context432wp = context.wp433pixels = workspace.spec.detector.pixels434dimension = len(parameters)435partial_sizes: list[int] = []436if local_model:437if dimension > 16:438raise TransportError("dense local models support at most 16 active parameters")439count = (pixels + 255) // 256440while True:441partial_sizes.append(dimension * dimension * count)442if count == 1:443break444count = (count + 255) // 256445if dimension * pixels >= 2**31 or any(size >= 2**31 for size in partial_sizes):446raise TransportError("local model exceeds signed 32-bit indexing")447required_bytes = 8 * (dimension * pixels + sum(partial_sizes) + dimension * dimension)448if required_bytes > local_model_max_bytes:449raise TransportError(450f"local model needs {required_bytes} bytes, exceeding preparation budget"451)452for name, value in (("observation", observation), ("pixel_weights", pixel_weights)):453context.array(value, name, dtype=wp.float64, shape=(pixels,))454workspace._launch(455workspace._kernels.validate_measurement,456pixels,457[observation, pixel_weights, workspace._status],458)459workspace.check_status()460mapping = [-1] * workspace.spec.grid.materials461for index, parameter in enumerate(parameters):462if parameter.material is not None:463if parameter.material >= len(mapping):464raise TransportError("active material is outside this model")465mapping[parameter.material] = index466arrays: dict[str, Any] = {}467with context.scope():468for 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):478arrays[name] = wp.empty(workspace.max_histories, dtype=dtype, device=context.device)479for name in ("mean_a", "mean_b", "components"):480arrays[name] = wp.empty(pixels, dtype=wp.float64, device=context.device)481arrays["scalar"] = wp.empty(1, dtype=wp.float64, device=context.device)482arrays["density"] = wp.empty(len(base_density), dtype=wp.float64, device=context.device)483arrays["base_density"] = wp.array(484list(base_density), dtype=wp.float64, device=context.device485)486arrays["material_parameter"] = wp.array(mapping, dtype=wp.int32, device=context.device)487arrays["chart"] = wp.empty(len(parameters), dtype=wp.float64, device=context.device)488parameter_host = wp.empty(len(parameters), dtype=wp.float64, device="cpu", pinned=True)489scalar_host = wp.empty(1, dtype=wp.float64, device="cpu", pinned=True)490model_kernels = None491model_partials: tuple[Any, ...] = ()492model_host = None493model_view = None494deterministic = False495if local_model or workspace.spec.estimator == "continuous-absorption":496model_kernels = load_kernels("dpt.transport.recovery_kernels")497with context.scope():498if local_model:499arrays["jacobian"] = wp.empty(500dimension * pixels, dtype=wp.float64, device=context.device501)502model_partials = tuple(503wp.empty(n, dtype=wp.float64, device=context.device) for n in partial_sizes504)505model_host = wp.empty(506dimension * dimension, dtype=wp.float64, device="cpu", pinned=True507)508model_view = model_host.numpy()509if workspace.spec.estimator == "continuous-absorption" and source.extent_xy_mm == (5100.0,5110.0,512):513assert model_kernels is not None514flag = wp.zeros(1, dtype=wp.int32, device=context.device)515workspace._launch(516model_kernels.flag_scattering,517int(workspace.scattering.size),518[workspace.scattering, flag],519)520wp.synchronize_stream(context.stream)521deterministic = int(flag.numpy()[0]) == 0522moments = prepare_estimators(workspace)523objective = prepare_objective(524ObjectiveSpec(), max_pixels=pixels, device=context.device, stream=context.stream525)526return TransportSquaredOracle(527workspace,528source,529parameters,530observation,531pixel_weights,532fixed_amplitude,533moments,534objective,535arrays,536parameter_host,537parameter_host.numpy(),538scalar_host,539scalar_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