"""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",
]