Generated from the full canonical file for this source snapshot. Line numbers match the library source.
Source SHA256: 5837deeeace08b277b2c0859254a41e413e6901dfbb9b8e8803076526b24ef68
1"""CUDA transmission with explicit storage, numerical and differentiation contracts.23Importing this module does not initialise CUDA or import Warp. Checked calls4validate device contents; ``validate=False`` is the asynchronous integration5boundary for callers that guarantee the current inputs satisfy these contracts.6No output arrays are allocated by an evaluation. See ``TransmissionWorkspace``7for stream ownership and checked completion of custom tape gradients.8"""910from __future__ import annotations1112# Scientific API names mirror the chapter symbols; private scratch is module-owned.13# ruff: noqa: N803, N81514# pyright: reportPrivateUsage=false15import math16from dataclasses import dataclass, field17from typing import Any, Literal1819from dpt._reductions import ReductionTree, prepare_reduction20from dpt._runtime import DeviceContext, ensure_tape, load_kernels, prepare_context, require_no_tape21from dpt.contracts import ContractError, NumericalError, binary32_scalar2223BeamMode = Literal["none", "scalar", "device-scalar", "per-pixel"]24_BEAM_MODES: dict[str, int] = {"none": 0, "scalar": 1, "device-scalar": 2, "per-pixel": 3}25_LIMIT = 2**31 - 126_TILE = 256272829class TransmissionError(ContractError):30"""Input, layout, aliasing or ownership violates the operator contract."""313233class GradientRangeError(NumericalError):34"""A requested gradient cannot be represented in its binary32 destination."""353637@dataclass(frozen=True, slots=True)38class TransmissionSpec:39"""Compile-independent input policy; the beam scalar's value remains dynamic."""4041beam: BeamMode = "none"42active_L: bool = True43active_beam: bool = False44block_dim: int = 2564546def __post_init__(self) -> None:47if self.beam not in _BEAM_MODES:48raise TransmissionError(f"unsupported beam representation: {self.beam!r}")49if self.active_beam and self.beam not in ("device-scalar", "per-pixel"):50raise TransmissionError("an active beam must be a CUDA array, not a Python scalar")51if self.block_dim not in (128, 256):52raise TransmissionError("block_dim must be 128 or 256")535455@dataclass(slots=True)56class TransmissionWorkspace:57"""Prepared scratch on one CUDA stream; never concurrently share a workspace.5859Inputs must stay unchanged until backward completes. Custom tape callbacks60enqueue range checks asynchronously; call ``check_status()`` after backward61before accepting the gradients. Use ``tape.zero()`` between independent62reverse passes. Second derivatives through the custom callbacks are not63supported. Workspace memory is O(P / 256) only for an active scalar beam.64"""6566spec: TransmissionSpec67max_pixels: int68context: DeviceContext69_kernels: Any = field(repr=False)70_empty: Any = field(repr=False)71_status: Any = field(repr=False)72_reduction: ReductionTree | None = field(repr=False)7374@property75def device(self) -> Any:76return self.context.device7778@property79def stream(self) -> Any:80return self.context.stream8182@property83def _wp(self) -> Any:84return self.context.wp8586@property87def scratch_bytes(self) -> int:88"""Device scratch excluding caller-owned inputs, outputs and their adjoints."""89return 4 + (self._reduction.scratch_bytes if self._reduction is not None else 0)9091def clear_status(self) -> None:92"""Clear accumulated diagnostics explicitly, on the owning stream."""93with self._wp.ScopedStream(self.stream, sync_enter=False):94self._status.zero_()9596def check_status(self) -> None:97"""Synchronise the owning stream and reject any accumulated numerical error."""98self._wp.synchronize_stream(self.stream)99code = int(self._status.numpy()[0])100if code & 2:101raise GradientRangeError("gradient overflow: discard this reverse pass")102if code:103raise TransmissionError("non-finite or out-of-domain device input")104105106def prepare_transmission(107spec: TransmissionSpec | None = None,108*,109device: str = "cuda:0",110max_pixels: int,111stream: Any = None,112) -> TransmissionWorkspace:113"""Allocate persistent diagnostics and reduction scratch, never image outputs."""114if isinstance(max_pixels, bool) or type(max_pixels) is not int:115raise TransmissionError("max_pixels must be an integer")116if not 0 <= max_pixels <= _LIMIT:117raise TransmissionError(f"max_pixels must lie in [0, {_LIMIT}]")118context = prepare_context(device=device, stream=stream, contract_error=TransmissionError)119wp, kernels = context.wp, load_kernels("dpt.kernels.transmission")120spec = spec or TransmissionSpec()121reduction = None122with context.scope():123empty = wp.empty(0, dtype=wp.float32, device=context.device)124status = wp.zeros(1, dtype=wp.int32, device=context.device)125if spec.active_beam and spec.beam == "device-scalar":126reduction = prepare_reduction(context, max_pixels, tile=_TILE)127return TransmissionWorkspace(spec, max_pixels, context, kernels, empty, status, reduction)128129130def _array(131value: Any, name: str, workspace: TransmissionWorkspace, length: int | None = None132) -> int:133size = workspace.context.array(134value,135name,136dtype=workspace._wp.float32,137shape=(length,) if length is not None else None,138)139if size > workspace.max_pixels and length != 1:140raise TransmissionError(f"{name} exceeds workspace capacity")141return size142143144def _beam(n0: Any, count: int, workspace: TransmissionWorkspace) -> tuple[Any, float]:145mode = workspace.spec.beam146if mode == "none":147if n0 is not None:148raise TransmissionError("this workspace declares no open-beam input")149return workspace._empty, 0.0150if mode == "scalar":151try:152return workspace._empty, binary32_scalar(n0, "n0", minimum=0.0)153except ContractError as error:154raise TransmissionError(str(error)) from error155_array(n0, "n0", workspace, 1 if mode == "device-scalar" else count)156return n0, 0.0157158159def _values(workspace: TransmissionWorkspace, arrays: list[tuple[str, Any, str]]) -> None:160wp, kernels = workspace._wp, workspace._kernels161workspace.clear_status()162for _, array, kind in arrays:163if array.size:164wp.launch(165getattr(kernels, kind),166dim=array.size,167inputs=[array],168outputs=[workspace._status],169stream=workspace.stream,170record_tape=False,171)172try:173workspace.check_status()174except TransmissionError as error:175for name, array, kind in arrays:176for index, item in enumerate(array.numpy()):177value = float(item)178valid = math.isfinite(value)179if kind != "finite_seed":180valid = valid and value >= 0181if kind == "decrement_domain":182valid = valid and value < 1183if not valid:184raise TransmissionError(f"{name}[{index}]={value!r} violates {kind}") from error185raise186187188# region book:transmission-contract189def transmit(190L: Any,191n0: Any = None,192*,193out_T: Any = None,194out_counts: Any = None,195out_log_T: Any = None,196out_removed: Any = None,197workspace: TransmissionWorkspace,198stream: Any = None,199tape: Any = None,200validate: bool = True,201) -> None:202"""Evaluate selected deterministic outputs in caller-owned binary32 buffers.203204L is finite, non-negative optical depth. n0 is the declared open-beam205expectation at the detector. No sampling, clipping or implicit transfer is206performed. The default checks device contents before writing outputs.207With validate=False the caller guarantees valid *current* device inputs;208launches remain asynchronous. Pass tape explicitly to record the custom209first-order adjoint, and retain original inputs until backward completes.210"""211# endregion book:transmission-contract212ensure_tape(tape, contract_error=TransmissionError)213wp, kernels = workspace._wp, workspace._kernels214selected = workspace.context.assert_stream(stream)215count = _array(L, "L", workspace)216beam, scalar = _beam(n0, count, workspace)217outputs = (out_T, out_counts, out_log_T, out_removed)218mask = sum(1 << i for i, value in enumerate(outputs) if value is not None)219if mask == 0:220raise TransmissionError("request at least one output")221if out_counts is not None and workspace.spec.beam == "none":222raise TransmissionError("counts require an explicit open-beam input")223writes: list[tuple[str, Any]] = []224for name, value in zip(225("out_T", "out_counts", "out_log_T", "out_removed"), outputs, strict=True226):227if value is not None:228_array(value, name, workspace, count)229writes.append((name, value))230workspace.context.disjoint([("L", L), ("n0", beam)], writes)231arrays = _tape_arrays(L, beam, outputs, workspace) if tape is not None else []232# Stream ownership was checked above; callers supply cross-stream ordering.233with wp.ScopedStream(selected, sync_enter=False):234if validate:235_values(workspace, [("L", L, "finite_nonnegative"), ("n0", beam, "finite_nonnegative")])236if count:237wp.launch(238kernels.get_forward_kernel(mask, _BEAM_MODES[workspace.spec.beam]),239dim=count,240inputs=[L, beam, scalar],241outputs=[value if value is not None else workspace._empty for value in outputs],242stream=selected,243block_dim=workspace.spec.block_dim,244record_tape=False,245)246if tape is not None:247_record_transmission(tape, L, n0, beam, outputs, workspace, arrays)248249250def _tape_arrays(251L: Any, beam: Any, outputs: tuple[Any, ...], workspace: TransmissionWorkspace252) -> list[Any]:253arrays = [value for value in outputs if value is not None]254if workspace.spec.active_L:255arrays.append(L)256if workspace.spec.active_beam:257arrays.append(beam)258if not workspace.spec.active_L and not workspace.spec.active_beam:259raise TransmissionError("recording requires at least one active input")260for array in arrays:261if array.grad is None:262raise TransmissionError("allocate participating tape arrays with requires_grad=True")263workspace.context.disjoint(264[("forward input", L), ("beam", beam)]265+ [("forward output", a) for a in outputs if a is not None],266[("gradient", a.grad) for a in arrays],267)268return arrays269270271def _dependencies(272tape: Any, reads: tuple[Any, Any], outputs: tuple[Any, ...], workspace: TransmissionWorkspace273) -> None:274if not workspace._wp.config.verify_autograd_array_access:275return276# Public access markers diagnose recorded writes before backward. A zero-dim277# record gives Tape.reset()/backward() ownership of fixed inputs and view278# parents too, without allocating gradients or launching a marker thread.279for value in outputs:280if value is not None:281value.mark_write()282for value in reads:283value.mark_read()284padded = [value if value is not None else workspace._empty for value in outputs]285padded += [workspace._empty] * (4 - len(padded))286tape.record_launch(287workspace._kernels.dependency_marker,288dim=0,289max_blocks=0,290inputs=list(reads),291outputs=padded,292device=workspace.device,293block_dim=workspace.spec.block_dim,294)295296297def _record_transmission(298tape: Any,299L: Any,300n0: Any,301beam: Any,302outputs: tuple[Any, ...],303workspace: TransmissionWorkspace,304arrays: list[Any],305) -> None:306wp = workspace._wp307_dependencies(tape, (L, beam), outputs, workspace)308309def backward() -> None:310workspace.context.assert_stream()311transmission_vjp(312L,313n0,314seed_T=outputs[0].grad if outputs[0] is not None else None,315seed_counts=outputs[1].grad if outputs[1] is not None else None,316seed_log_T=outputs[2].grad if outputs[2] is not None else None,317seed_removed=outputs[3].grad if outputs[3] is not None else None,318out_grad_L=L.grad if workspace.spec.active_L else None,319out_grad_n0=beam.grad if workspace.spec.active_beam else None,320workspace=workspace,321validate=False,322_accumulate=True,323)324with wp.ScopedStream(workspace.stream, sync_enter=False):325for value in outputs:326if value is not None and not value.retain_grad:327value.grad.zero_()328329tape.record_func(backward, arrays)330331332def transmission_vjp(333L: Any,334n0: Any = None,335*,336seed_T: Any = None,337seed_counts: Any = None,338seed_log_T: Any = None,339seed_removed: Any = None,340out_grad_L: Any = None,341out_grad_n0: Any = None,342workspace: TransmissionWorkspace,343stream: Any = None,344validate: bool = True,345_accumulate: bool = False,346) -> None:347"""Overwrite requested first-order VJPs; caller cotangents are preserved.348349Missing seeds are zero. Active scalar illumination is reduced in a fixed350FP64 tree. validate=False defers range-error reporting to check_status().351_accumulate is reserved for the tape adapter, not the standalone API.352"""353require_no_tape(contract_error=TransmissionError)354wp, kernels = workspace._wp, workspace._kernels355selected = workspace.context.assert_stream(stream)356count = _array(L, "L", workspace)357beam, scalar = _beam(n0, count, workspace)358seeds = (seed_T, seed_counts, seed_log_T, seed_removed)359mask = sum(1 << i for i, value in enumerate(seeds) if value is not None)360if seed_counts is not None and workspace.spec.beam == "none":361raise TransmissionError("a count cotangent requires an open-beam input")362reads = [("L", L), ("n0", beam)]363values = [("L", L, "finite_nonnegative"), ("n0", beam, "finite_nonnegative")]364for name, seed in zip(365("seed_T", "seed_counts", "seed_log_T", "seed_removed"), seeds, strict=True366):367if seed is not None:368_array(seed, name, workspace, count)369reads.append((name, seed))370values.append((name, seed, "finite_seed"))371writes: list[tuple[str, Any]] = []372if out_grad_L is not None:373if not workspace.spec.active_L:374raise TransmissionError("L is declared fixed")375_array(out_grad_L, "out_grad_L", workspace, count)376writes.append(("out_grad_L", out_grad_L))377if out_grad_n0 is not None:378if not workspace.spec.active_beam:379raise TransmissionError("n0 is declared fixed")380_array(381out_grad_n0,382"out_grad_n0",383workspace,3841 if workspace.spec.beam == "device-scalar" else count,385)386writes.append(("out_grad_n0", out_grad_n0))387if not writes:388raise TransmissionError("request at least one active input gradient")389workspace.context.disjoint(reads, writes)390# Stream ownership was checked above; callers supply cross-stream ordering.391with wp.ScopedStream(selected, sync_enter=False):392if validate:393_values(workspace, values)394per_pixel_beam = out_grad_n0 is not None and workspace.spec.beam == "per-pixel"395if count and (out_grad_L is not None or per_pixel_beam):396wp.launch(397kernels.get_vjp_kernel(398mask,399_BEAM_MODES[workspace.spec.beam],400out_grad_L is not None,401per_pixel_beam,402_accumulate,403),404dim=count,405inputs=[L, beam, scalar]406+ [v if v is not None else workspace._empty for v in seeds],407outputs=[408out_grad_L if out_grad_L is not None else workspace._empty,409out_grad_n0 if per_pixel_beam else workspace._empty,410workspace._status,411],412stream=selected,413block_dim=workspace.spec.block_dim,414record_tape=False,415)416if out_grad_n0 is not None and workspace.spec.beam == "device-scalar":417if count and seed_counts is not None:418reduction = workspace._reduction419assert reduction is not None420blocks = (count + _TILE - 1) // _TILE421wp.launch_tiled(422kernels.get_beam_partial_kernel(),423dim=blocks,424inputs=[L, seed_counts],425outputs=[reduction.partials[0]],426block_dim=_TILE,427stream=selected,428record_tape=False,429)430reduction.finish(431blocks,432out_grad_n0,433workspace._status,434accumulate=_accumulate,435)436elif not _accumulate:437out_grad_n0.zero_()438if validate:439workspace.check_status()440441442def optical_depth_from_removed(443delta: Any,444*,445out_L: Any,446workspace: TransmissionWorkspace,447stream: Any = None,448tape: Any = None,449validate: bool = True,450) -> None:451"""Evaluate -log1p(-delta), for finite 0 <= delta < 1, without cancellation."""452ensure_tape(tape, contract_error=TransmissionError)453wp, kernels = workspace._wp, workspace._kernels454selected = workspace.context.assert_stream(stream)455count = _array(delta, "delta", workspace)456_array(out_L, "out_L", workspace, count)457workspace.context.disjoint([("delta", delta)], [("out_L", out_L)])458if tape is not None:459if delta.grad is None or out_L.grad is None:460raise TransmissionError("inverse tape arrays require preallocated gradients")461workspace.context.disjoint(462[("delta", delta), ("out_L", out_L)],463[("delta.grad", delta.grad), ("out_L.grad", out_L.grad)],464)465# Stream ownership was checked above; callers supply cross-stream ordering.466with wp.ScopedStream(selected, sync_enter=False):467if validate:468_values(workspace, [("delta", delta, "decrement_domain")])469if count:470wp.launch(471kernels.get_inverse_kernel(),472dim=count,473inputs=[delta],474outputs=[out_L],475stream=selected,476block_dim=workspace.spec.block_dim,477record_tape=False,478)479if tape is not None:480_dependencies(tape, (delta, workspace._empty), (out_L,), workspace)481482def backward() -> None:483workspace.context.assert_stream()484optical_depth_from_removed_vjp(485delta,486out_L.grad,487out_grad_delta=delta.grad,488workspace=workspace,489validate=False,490_accumulate=True,491)492if not out_L.retain_grad:493out_L.grad.zero_()494495tape.record_func(backward, [delta, out_L])496497498def optical_depth_from_removed_vjp(499delta: Any,500seed_L: Any,501*,502out_grad_delta: Any,503workspace: TransmissionWorkspace,504stream: Any = None,505validate: bool = True,506_accumulate: bool = False,507) -> None:508"""First-order inverse-decrement VJP; preserve seeds and overwrite the result."""509require_no_tape(contract_error=TransmissionError)510wp, kernels = workspace._wp, workspace._kernels511selected = workspace.context.assert_stream(stream)512count = _array(delta, "delta", workspace)513_array(seed_L, "seed_L", workspace, count)514_array(out_grad_delta, "out_grad_delta", workspace, count)515workspace.context.disjoint(516[("delta", delta), ("seed_L", seed_L)], [("out_grad_delta", out_grad_delta)]517)518# Stream ownership was checked above; callers supply cross-stream ordering.519with wp.ScopedStream(selected, sync_enter=False):520if validate:521_values(522workspace, [("delta", delta, "decrement_domain"), ("seed_L", seed_L, "finite_seed")]523)524if count:525wp.launch(526kernels.get_inverse_vjp_kernel(_accumulate),527dim=count,528inputs=[delta, seed_L],529outputs=[out_grad_delta, workspace._status],530stream=selected,531block_dim=workspace.spec.block_dim,532record_tape=False,533)534if validate:535workspace.check_status()536