"""Differential linearisers for the mock pruning operators.
``boundlab::HeavisidePruning`` and ``boundlab::TopKPruning`` apply a
score-derived 0/1 mask to the **second** network only::
out_x = x
out_y = mask · y
out_d = mask · d + (1 − mask) · x
With *concrete* scores the mask is deterministic, so that rewrite is exact —
the difference carries no relaxation error whatsoever.
Symbolic (input-dependent) scores are rejected rather than approximated: a
single fixed mask cannot soundly enclose an input-varying kept-set. Such
models are handled one level up, by enumerating the reachable kept-sets and
taking the union of the per-case bounds.
"""
from __future__ import annotations
from typing import Any
import torch
from boundlab import Expr
from boundlab.ibp import Bias
from boundlab.interp import Interpreter, OpHandler
from boundlab.diff.expr import DiffExpr2, DiffExpr3
def const_scores(scores: Any) -> torch.Tensor:
"""Extract the concrete score tensor, or refuse loudly."""
if isinstance(scores, (DiffExpr2, DiffExpr3)):
scores = scores.y # only the second network is pruned
if isinstance(scores, Bias):
return scores.arr
if isinstance(scores, torch.Tensor):
return scores
raise NotImplementedError(
"Pruning requires constant scores. Input-dependent scores must be "
"handled by enumerating the reachable kept-sets and unioning the "
"per-case bounds, not by relaxing the mask."
)
def apply_mask(mask: torch.Tensor, data: Any) -> DiffExpr3:
"""Mask the second branch with a concrete 0/1 ``mask``; the first stays whole."""
if isinstance(data, torch.Tensor):
data = Bias(data) # a graph constant: shared between the networks
triple = data.to(DiffExpr3)
mask = torch.broadcast_to(mask, tuple(triple.shape_dtype.shape)).to(
triple.shape_dtype.dtype
)
return DiffExpr3(
triple.x,
triple.y * mask,
triple.diff * mask + triple.x * (1.0 - mask),
)
def heaviside_mask(scores: torch.Tensor) -> torch.Tensor:
"""``h(s) = 1`` for ``s >= 0`` — matching the eager operator's convention."""
dtype = scores.dtype if scores.is_floating_point() else torch.get_default_dtype()
return (scores >= 0).to(dtype)
def topk_mask(scores: torch.Tensor, k: int, dim: int = -1) -> torch.Tensor:
"""1 at the ``k`` highest-scoring positions along ``dim``, 0 elsewhere."""
dtype = scores.dtype if scores.is_floating_point() else torch.get_default_dtype()
kept = max(0, min(int(k), int(scores.shape[dim])))
mask = torch.zeros_like(scores, dtype=dtype)
if kept > 0:
mask.scatter_(dim, torch.topk(scores, kept, dim=dim).indices, 1.0)
return mask
[docs]
class DiffHeavisidePruning(OpHandler):
"""Exact handler for ``boundlab::HeavisidePruning`` with concrete scores."""
op = "heaviside_pruning"
[docs]
def handle(
self, interp: Interpreter, scores: Any, data: Any, **kwargs: Any
) -> Expr:
del interp, kwargs
return apply_mask(heaviside_mask(const_scores(scores)), data)
[docs]
class DiffTopKPruning(OpHandler):
"""Exact handler for ``boundlab::TopKPruning`` with concrete scores."""
op = "topk_pruning"
[docs]
def handle(
self,
interp: Interpreter,
scores: Any,
data: Any,
*,
k: int,
dim: int = -1,
**kwargs: Any,
) -> Expr:
del interp, kwargs
return apply_mask(topk_mask(const_scores(scores), k, dim), data)
__all__ = [
"DiffHeavisidePruning",
"DiffTopKPruning",
"apply_mask",
"const_scores",
"heaviside_mask",
"topk_mask",
]