Source code for boundlab.diff.zono3.heaviside

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