python/dpt/transmission.py

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."""4041    beam: BeamMode = "none"42    active_L: bool = True43    active_beam: bool = False44    block_dim: int = 2564546    def __post_init__(self) -> None:47        if self.beam not in _BEAM_MODES:48            raise TransmissionError(f"unsupported beam representation: {self.beam!r}")49        if self.active_beam and self.beam not in ("device-scalar", "per-pixel"):50            raise TransmissionError("an active beam must be a CUDA array, not a Python scalar")51        if self.block_dim not in (128, 256):52            raise 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.5859    Inputs must stay unchanged until backward completes. Custom tape callbacks60    enqueue range checks asynchronously; call ``check_status()`` after backward61    before accepting the gradients. Use ``tape.zero()`` between independent62    reverse passes. Second derivatives through the custom callbacks are not63    supported. Workspace memory is O(P / 256) only for an active scalar beam.64    """6566    spec: TransmissionSpec67    max_pixels: int68    context: 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    @property75    def device(self) -> Any:76        return self.context.device7778    @property79    def stream(self) -> Any:80        return self.context.stream8182    @property83    def _wp(self) -> Any:84        return self.context.wp8586    @property87    def scratch_bytes(self) -> int:88        """Device scratch excluding caller-owned inputs, outputs and their adjoints."""89        return 4 + (self._reduction.scratch_bytes if self._reduction is not None else 0)9091    def clear_status(self) -> None:92        """Clear accumulated diagnostics explicitly, on the owning stream."""93        with self._wp.ScopedStream(self.stream, sync_enter=False):94            self._status.zero_()9596    def check_status(self) -> None:97        """Synchronise the owning stream and reject any accumulated numerical error."""98        self._wp.synchronize_stream(self.stream)99        code = int(self._status.numpy()[0])100        if code & 2:101            raise GradientRangeError("gradient overflow: discard this reverse pass")102        if code:103            raise TransmissionError("non-finite or out-of-domain device input")104105106def prepare_transmission(107    spec: TransmissionSpec | None = None,108    *,109    device: str = "cuda:0",110    max_pixels: int,111    stream: Any = None,112) -> TransmissionWorkspace:113    """Allocate persistent diagnostics and reduction scratch, never image outputs."""114    if isinstance(max_pixels, bool) or type(max_pixels) is not int:115        raise TransmissionError("max_pixels must be an integer")116    if not 0 <= max_pixels <= _LIMIT:117        raise TransmissionError(f"max_pixels must lie in [0, {_LIMIT}]")118    context = prepare_context(device=device, stream=stream, contract_error=TransmissionError)119    wp, kernels = context.wp, load_kernels("dpt.kernels.transmission")120    spec = spec or TransmissionSpec()121    reduction = None122    with context.scope():123        empty = wp.empty(0, dtype=wp.float32, device=context.device)124        status = wp.zeros(1, dtype=wp.int32, device=context.device)125        if spec.active_beam and spec.beam == "device-scalar":126            reduction = prepare_reduction(context, max_pixels, tile=_TILE)127    return TransmissionWorkspace(spec, max_pixels, context, kernels, empty, status, reduction)128129130def _array(131    value: Any, name: str, workspace: TransmissionWorkspace, length: int | None = None132) -> int:133    size = workspace.context.array(134        value,135        name,136        dtype=workspace._wp.float32,137        shape=(length,) if length is not None else None,138    )139    if size > workspace.max_pixels and length != 1:140        raise TransmissionError(f"{name} exceeds workspace capacity")141    return size142143144def _beam(n0: Any, count: int, workspace: TransmissionWorkspace) -> tuple[Any, float]:145    mode = workspace.spec.beam146    if mode == "none":147        if n0 is not None:148            raise TransmissionError("this workspace declares no open-beam input")149        return workspace._empty, 0.0150    if mode == "scalar":151        try:152            return workspace._empty, binary32_scalar(n0, "n0", minimum=0.0)153        except ContractError as error:154            raise TransmissionError(str(error)) from error155    _array(n0, "n0", workspace, 1 if mode == "device-scalar" else count)156    return n0, 0.0157158159def _values(workspace: TransmissionWorkspace, arrays: list[tuple[str, Any, str]]) -> None:160    wp, kernels = workspace._wp, workspace._kernels161    workspace.clear_status()162    for _, array, kind in arrays:163        if array.size:164            wp.launch(165                getattr(kernels, kind),166                dim=array.size,167                inputs=[array],168                outputs=[workspace._status],169                stream=workspace.stream,170                record_tape=False,171            )172    try:173        workspace.check_status()174    except TransmissionError as error:175        for name, array, kind in arrays:176            for index, item in enumerate(array.numpy()):177                value = float(item)178                valid = math.isfinite(value)179                if kind != "finite_seed":180                    valid = valid and value >= 0181                if kind == "decrement_domain":182                    valid = valid and value < 1183                if not valid:184                    raise TransmissionError(f"{name}[{index}]={value!r} violates {kind}") from error185        raise186187188# region book:transmission-contract189def transmit(190    L: Any,191    n0: Any = None,192    *,193    out_T: Any = None,194    out_counts: Any = None,195    out_log_T: Any = None,196    out_removed: Any = None,197    workspace: TransmissionWorkspace,198    stream: Any = None,199    tape: Any = None,200    validate: bool = True,201) -> None:202    """Evaluate selected deterministic outputs in caller-owned binary32 buffers.203204    L is finite, non-negative optical depth. n0 is the declared open-beam205    expectation at the detector. No sampling, clipping or implicit transfer is206    performed. The default checks device contents before writing outputs.207    With validate=False the caller guarantees valid *current* device inputs;208    launches remain asynchronous. Pass tape explicitly to record the custom209    first-order adjoint, and retain original inputs until backward completes.210    """211    # endregion book:transmission-contract212    ensure_tape(tape, contract_error=TransmissionError)213    wp, kernels = workspace._wp, workspace._kernels214    selected = workspace.context.assert_stream(stream)215    count = _array(L, "L", workspace)216    beam, scalar = _beam(n0, count, workspace)217    outputs = (out_T, out_counts, out_log_T, out_removed)218    mask = sum(1 << i for i, value in enumerate(outputs) if value is not None)219    if mask == 0:220        raise TransmissionError("request at least one output")221    if out_counts is not None and workspace.spec.beam == "none":222        raise TransmissionError("counts require an explicit open-beam input")223    writes: list[tuple[str, Any]] = []224    for name, value in zip(225        ("out_T", "out_counts", "out_log_T", "out_removed"), outputs, strict=True226    ):227        if value is not None:228            _array(value, name, workspace, count)229            writes.append((name, value))230    workspace.context.disjoint([("L", L), ("n0", beam)], writes)231    arrays = _tape_arrays(L, beam, outputs, workspace) if tape is not None else []232    # Stream ownership was checked above; callers supply cross-stream ordering.233    with wp.ScopedStream(selected, sync_enter=False):234        if validate:235            _values(workspace, [("L", L, "finite_nonnegative"), ("n0", beam, "finite_nonnegative")])236        if count:237            wp.launch(238                kernels.get_forward_kernel(mask, _BEAM_MODES[workspace.spec.beam]),239                dim=count,240                inputs=[L, beam, scalar],241                outputs=[value if value is not None else workspace._empty for value in outputs],242                stream=selected,243                block_dim=workspace.spec.block_dim,244                record_tape=False,245            )246        if tape is not None:247            _record_transmission(tape, L, n0, beam, outputs, workspace, arrays)248249250def _tape_arrays(251    L: Any, beam: Any, outputs: tuple[Any, ...], workspace: TransmissionWorkspace252) -> list[Any]:253    arrays = [value for value in outputs if value is not None]254    if workspace.spec.active_L:255        arrays.append(L)256    if workspace.spec.active_beam:257        arrays.append(beam)258    if not workspace.spec.active_L and not workspace.spec.active_beam:259        raise TransmissionError("recording requires at least one active input")260    for array in arrays:261        if array.grad is None:262            raise TransmissionError("allocate participating tape arrays with requires_grad=True")263    workspace.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    )268    return arrays269270271def _dependencies(272    tape: Any, reads: tuple[Any, Any], outputs: tuple[Any, ...], workspace: TransmissionWorkspace273) -> None:274    if not workspace._wp.config.verify_autograd_array_access:275        return276    # 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.279    for value in outputs:280        if value is not None:281            value.mark_write()282    for value in reads:283        value.mark_read()284    padded = [value if value is not None else workspace._empty for value in outputs]285    padded += [workspace._empty] * (4 - len(padded))286    tape.record_launch(287        workspace._kernels.dependency_marker,288        dim=0,289        max_blocks=0,290        inputs=list(reads),291        outputs=padded,292        device=workspace.device,293        block_dim=workspace.spec.block_dim,294    )295296297def _record_transmission(298    tape: Any,299    L: Any,300    n0: Any,301    beam: Any,302    outputs: tuple[Any, ...],303    workspace: TransmissionWorkspace,304    arrays: list[Any],305) -> None:306    wp = workspace._wp307    _dependencies(tape, (L, beam), outputs, workspace)308309    def backward() -> None:310        workspace.context.assert_stream()311        transmission_vjp(312            L,313            n0,314            seed_T=outputs[0].grad if outputs[0] is not None else None,315            seed_counts=outputs[1].grad if outputs[1] is not None else None,316            seed_log_T=outputs[2].grad if outputs[2] is not None else None,317            seed_removed=outputs[3].grad if outputs[3] is not None else None,318            out_grad_L=L.grad if workspace.spec.active_L else None,319            out_grad_n0=beam.grad if workspace.spec.active_beam else None,320            workspace=workspace,321            validate=False,322            _accumulate=True,323        )324        with wp.ScopedStream(workspace.stream, sync_enter=False):325            for value in outputs:326                if value is not None and not value.retain_grad:327                    value.grad.zero_()328329    tape.record_func(backward, arrays)330331332def transmission_vjp(333    L: Any,334    n0: Any = None,335    *,336    seed_T: Any = None,337    seed_counts: Any = None,338    seed_log_T: Any = None,339    seed_removed: Any = None,340    out_grad_L: Any = None,341    out_grad_n0: Any = None,342    workspace: TransmissionWorkspace,343    stream: Any = None,344    validate: bool = True,345    _accumulate: bool = False,346) -> None:347    """Overwrite requested first-order VJPs; caller cotangents are preserved.348349    Missing seeds are zero. Active scalar illumination is reduced in a fixed350    FP64 tree. validate=False defers range-error reporting to check_status().351    _accumulate is reserved for the tape adapter, not the standalone API.352    """353    require_no_tape(contract_error=TransmissionError)354    wp, kernels = workspace._wp, workspace._kernels355    selected = workspace.context.assert_stream(stream)356    count = _array(L, "L", workspace)357    beam, scalar = _beam(n0, count, workspace)358    seeds = (seed_T, seed_counts, seed_log_T, seed_removed)359    mask = sum(1 << i for i, value in enumerate(seeds) if value is not None)360    if seed_counts is not None and workspace.spec.beam == "none":361        raise TransmissionError("a count cotangent requires an open-beam input")362    reads = [("L", L), ("n0", beam)]363    values = [("L", L, "finite_nonnegative"), ("n0", beam, "finite_nonnegative")]364    for name, seed in zip(365        ("seed_T", "seed_counts", "seed_log_T", "seed_removed"), seeds, strict=True366    ):367        if seed is not None:368            _array(seed, name, workspace, count)369            reads.append((name, seed))370            values.append((name, seed, "finite_seed"))371    writes: list[tuple[str, Any]] = []372    if out_grad_L is not None:373        if not workspace.spec.active_L:374            raise TransmissionError("L is declared fixed")375        _array(out_grad_L, "out_grad_L", workspace, count)376        writes.append(("out_grad_L", out_grad_L))377    if out_grad_n0 is not None:378        if not workspace.spec.active_beam:379            raise TransmissionError("n0 is declared fixed")380        _array(381            out_grad_n0,382            "out_grad_n0",383            workspace,384            1 if workspace.spec.beam == "device-scalar" else count,385        )386        writes.append(("out_grad_n0", out_grad_n0))387    if not writes:388        raise TransmissionError("request at least one active input gradient")389    workspace.context.disjoint(reads, writes)390    # Stream ownership was checked above; callers supply cross-stream ordering.391    with wp.ScopedStream(selected, sync_enter=False):392        if validate:393            _values(workspace, values)394        per_pixel_beam = out_grad_n0 is not None and workspace.spec.beam == "per-pixel"395        if count and (out_grad_L is not None or per_pixel_beam):396            wp.launch(397                kernels.get_vjp_kernel(398                    mask,399                    _BEAM_MODES[workspace.spec.beam],400                    out_grad_L is not None,401                    per_pixel_beam,402                    _accumulate,403                ),404                dim=count,405                inputs=[L, beam, scalar]406                + [v if v is not None else workspace._empty for v in seeds],407                outputs=[408                    out_grad_L if out_grad_L is not None else workspace._empty,409                    out_grad_n0 if per_pixel_beam else workspace._empty,410                    workspace._status,411                ],412                stream=selected,413                block_dim=workspace.spec.block_dim,414                record_tape=False,415            )416        if out_grad_n0 is not None and workspace.spec.beam == "device-scalar":417            if count and seed_counts is not None:418                reduction = workspace._reduction419                assert reduction is not None420                blocks = (count + _TILE - 1) // _TILE421                wp.launch_tiled(422                    kernels.get_beam_partial_kernel(),423                    dim=blocks,424                    inputs=[L, seed_counts],425                    outputs=[reduction.partials[0]],426                    block_dim=_TILE,427                    stream=selected,428                    record_tape=False,429                )430                reduction.finish(431                    blocks,432                    out_grad_n0,433                    workspace._status,434                    accumulate=_accumulate,435                )436            elif not _accumulate:437                out_grad_n0.zero_()438        if validate:439            workspace.check_status()440441442def optical_depth_from_removed(443    delta: Any,444    *,445    out_L: Any,446    workspace: TransmissionWorkspace,447    stream: Any = None,448    tape: Any = None,449    validate: bool = True,450) -> None:451    """Evaluate -log1p(-delta), for finite 0 <= delta < 1, without cancellation."""452    ensure_tape(tape, contract_error=TransmissionError)453    wp, kernels = workspace._wp, workspace._kernels454    selected = workspace.context.assert_stream(stream)455    count = _array(delta, "delta", workspace)456    _array(out_L, "out_L", workspace, count)457    workspace.context.disjoint([("delta", delta)], [("out_L", out_L)])458    if tape is not None:459        if delta.grad is None or out_L.grad is None:460            raise TransmissionError("inverse tape arrays require preallocated gradients")461        workspace.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.466    with wp.ScopedStream(selected, sync_enter=False):467        if validate:468            _values(workspace, [("delta", delta, "decrement_domain")])469        if count:470            wp.launch(471                kernels.get_inverse_kernel(),472                dim=count,473                inputs=[delta],474                outputs=[out_L],475                stream=selected,476                block_dim=workspace.spec.block_dim,477                record_tape=False,478            )479        if tape is not None:480            _dependencies(tape, (delta, workspace._empty), (out_L,), workspace)481482            def backward() -> None:483                workspace.context.assert_stream()484                optical_depth_from_removed_vjp(485                    delta,486                    out_L.grad,487                    out_grad_delta=delta.grad,488                    workspace=workspace,489                    validate=False,490                    _accumulate=True,491                )492                if not out_L.retain_grad:493                    out_L.grad.zero_()494495            tape.record_func(backward, [delta, out_L])496497498def optical_depth_from_removed_vjp(499    delta: Any,500    seed_L: Any,501    *,502    out_grad_delta: Any,503    workspace: TransmissionWorkspace,504    stream: Any = None,505    validate: bool = True,506    _accumulate: bool = False,507) -> None:508    """First-order inverse-decrement VJP; preserve seeds and overwrite the result."""509    require_no_tape(contract_error=TransmissionError)510    wp, kernels = workspace._wp, workspace._kernels511    selected = workspace.context.assert_stream(stream)512    count = _array(delta, "delta", workspace)513    _array(seed_L, "seed_L", workspace, count)514    _array(out_grad_delta, "out_grad_delta", workspace, count)515    workspace.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.519    with wp.ScopedStream(selected, sync_enter=False):520        if validate:521            _values(522                workspace, [("delta", delta, "decrement_domain"), ("seed_L", seed_L, "finite_seed")]523            )524        if count:525            wp.launch(526                kernels.get_inverse_vjp_kernel(_accumulate),527                dim=count,528                inputs=[delta, seed_L],529                outputs=[out_grad_delta, workspace._status],530                stream=selected,531                block_dim=workspace.spec.block_dim,532                record_tape=False,533            )534        if validate:535            workspace.check_status()536