Source code for boundlab.diff.ops

"""Custom operators that mark differential structure inside a model.

These are ordinary Python functions that lower to custom-domain ONNX nodes
(``boundlab::DiffPair``, ``boundlab::HeavisidePruning``, …) when the model is
exported with :func:`boundlab.interp.onnx_export`.  The BoundLab interpreter
turns those nodes back into :class:`~boundlab.diff.expr.DiffExpr2` /
:class:`~boundlab.diff.expr.DiffExpr3` values; run eagerly, each op evaluates
the *second* (modified) network so the very same module can be executed
concretely for Monte-Carlo checks.

Importing this module registers the custom ONNX op names with
:mod:`boundlab.interp`, so ``boundlab::`` nodes dispatch to the handler names
``diff_pair``, ``heaviside_pruning``, ``softmax_pruning`` and ``topk_pruning``.
"""

from __future__ import annotations

from typing import Any

import torch
from torch import nn

from boundlab.interp import Interpreter, OpHandler, _ONNX_TO_HANDLER
from boundlab.diff.expr import DiffExpr2


# ---------------------------------------------------------------------------
# ONNX op-name registration
# ---------------------------------------------------------------------------

_ONNX_TO_HANDLER.update(
    {
        "DiffPair": "diff_pair",
        "HeavisidePruning": "heaviside_pruning",
        "SoftmaxPruning": "softmax_pruning",
        "TopKPruning": "topk_pruning",
    }
)


def _tracing() -> bool:
    return torch.compiler.is_exporting() or torch.jit.is_tracing()


# ---------------------------------------------------------------------------
# Operators
# ---------------------------------------------------------------------------


[docs] def diff_pair(x: torch.Tensor, y: torch.Tensor) -> torch.Tensor: """Mark ``x`` and ``y`` as the two branches of one differential value. Exported as ``boundlab::DiffPair``; the interpreter lifts that node into a :class:`~boundlab.diff.expr.DiffExpr2`. Eagerly it is a no-op returning ``x``, so a paired model still runs as network 1. Args: x: Tensor for the first network branch. y: Tensor for the second branch; same shape and dtype as ``x``. Examples -------- >>> import torch >>> from boundlab.diff.ops import diff_pair >>> diff_pair(torch.zeros(4), torch.ones(4)).shape torch.Size([4]) """ assert x.shape == y.shape, "diff_pair operands must have the same shape" if _tracing(): return torch.onnx.ops.symbolic( "boundlab::DiffPair", (x, y), dtype=x.dtype, shape=x.shape, version=1, ) return x
[docs] def heaviside_pruning(scores: torch.Tensor, data: torch.Tensor) -> torch.Tensor: """Mock score-based pruning: network 1 keeps ``data``, network 2 masks it. Network 2 computes ``heaviside(scores) * data`` with ``h(0) = 1``, matching the abstract interpreters' ``scores >= 0`` convention. Eagerly this evaluates the *pruned* network, so the model can be sampled directly. """ assert scores.shape == data.shape[-scores.dim():], ( "scores must match the trailing shape of data" ) if _tracing(): return torch.onnx.ops.symbolic( "boundlab::HeavisidePruning", (scores, data), dtype=data.dtype, shape=data.shape, version=1, ) return torch.heaviside(scores, scores.new_ones(())) * data
[docs] def softmax_pruning( scores: torch.Tensor, data: torch.Tensor, dim: int = -1 ) -> torch.Tensor: """Mock softmax pruning: network 1 is ``softmax(data)``, network 2 masks it. Network 2 computes the mask-renormalised softmax ``h(sⱼ)·exp(dⱼ) / Σₖ h(sₖ)·exp(dₖ)``. Eagerly this evaluates the pruned network. """ assert scores.shape == data.shape, "scores and data must have the same shape" dim = dim if dim >= 0 else data.dim() + dim if _tracing(): return torch.onnx.ops.symbolic( "boundlab::SoftmaxPruning", (scores, data), attrs={"dim": dim}, dtype=data.dtype, shape=data.shape, version=1, ) mask = torch.heaviside(scores, scores.new_ones(())) masked = mask * torch.exp(data) return masked / masked.sum(dim=dim, keepdim=True)
[docs] def topk_pruning( scores: torch.Tensor, data: torch.Tensor, k: int, dim: int = -1 ) -> torch.Tensor: """Mock top-k pruning: network 2 keeps the ``k`` highest-scoring positions. Eagerly this evaluates the pruned network, ``topk_mask(scores) * data``. """ assert scores.shape == data.shape, "scores and data must have the same shape" dim = dim if dim >= 0 else scores.dim() + dim if _tracing(): return torch.onnx.ops.symbolic( "boundlab::TopKPruning", (scores, data), attrs={"k": int(k), "dim": dim}, dtype=data.dtype, shape=data.shape, version=1, ) kept = max(0, min(int(k), scores.shape[dim])) mask = torch.zeros_like(data) if kept > 0: mask.scatter_(dim, torch.topk(scores, kept, dim=dim).indices, 1.0) return mask * data
[docs] class DiffLinear(nn.Module): """Two parallel linear layers paired through :func:`diff_pair`. Concretely equivalent to ``fc1(x)``; under differential interpretation the paired weights become a :class:`~boundlab.diff.expr.DiffExpr2` so both branches propagate at once. Examples -------- >>> import torch >>> from torch import nn >>> from boundlab.diff.ops import DiffLinear >>> DiffLinear(nn.Linear(4, 3), nn.Linear(4, 3))(torch.zeros(4)).shape torch.Size([3]) """
[docs] def __init__(self, fc1: nn.Linear, fc2: nn.Linear): super().__init__() assert fc1.in_features == fc2.in_features, "fc1/fc2 in_features differ" assert fc1.out_features == fc2.out_features, "fc1/fc2 out_features differ" assert (fc1.bias is not None) == (fc2.bias is not None), ( "fc1 and fc2 must both have a bias, or neither" ) self.fc1 = fc1 self.fc2 = fc2
[docs] def forward(self, x: torch.Tensor) -> torch.Tensor: weight = diff_pair(self.fc1.weight, self.fc2.weight) out = x @ weight.t() if self.fc1.bias is not None: assert self.fc2.bias is not None out = out + diff_pair(self.fc1.bias, self.fc2.bias) return out
# --------------------------------------------------------------------------- # Interpreter handler # ---------------------------------------------------------------------------
[docs] class DiffPair(OpHandler): """Lift a ``boundlab::DiffPair`` node into a :class:`DiffExpr2`.""" op = "diff_pair"
[docs] def handle(self, interp: Interpreter, x: Any, y: Any, **kwargs) -> DiffExpr2: del interp, kwargs return DiffExpr2(x, y)
__all__ = [ "DiffLinear", "DiffPair", "diff_pair", "heaviside_pruning", "softmax_pruning", "topk_pruning", ]