Source code for spacr.attribution

"""Attribution — what a trained classifier attends to, and whether that means anything.

spaCR already drew Grad-CAM and saliency maps (``spacr.utils.GradCAMGenerator``,
``spacr.utils.SaliencyMapGenerator``, ``spacr.utils.IntegratedGradients``). This
module is the library those two were the first entries in: the CAM family
(Grad-CAM, Grad-CAM++, Score-CAM, XGrad-CAM, Layer-CAM, Eigen-CAM, HiRes-CAM,
Ablation-CAM), the gradient family (saliency, integrated gradients, guided
backprop, input×gradient, DeepLIFT, and SmoothGrad wrapped around any of them),
the SHAP family (GradientSHAP, DeepSHAP), the perturbation family (occlusion,
feature ablation), and for Vision Transformers attention rollout plus Chefer et
al.'s class-specific relevance propagation.

Not every method applies to every backbone. :func:`method_applicability` says
which do for a model or a ``model_type`` name, and why the others do not, so a
method selector can grey an entry with its reason instead of failing mid-run.

Almost none of it is written here. The CAM variants come from the optional
**torchcam** attribution extra, while the gradient and perturbation methods
come from the core **captum** dependency. What is written here is the part no
library can supply: the
adapters that make every method agree on one output shape, the handling of
spaCR's two classifier head shapes, and — the reason this module exists — the
analyses in the second half.

**An attribution map is not an explanation.** It is a number per pixel produced
by a procedure. It does not show what the model "looked at", it does not
establish that the highlighted pixels caused the prediction, and it will render
a confident, beautiful, plausible picture for a model with random weights. Four
checks stand between a map and wishful thinking, and they are the point of this
module:

* :func:`deletion_curve` / :func:`insertion_curve` — remove (or add) the pixels
  the map ranks highest and watch the score. A map whose deletion curve is flat
  is not describing what the model uses, however good it looks.
* :func:`pointing_game` — does the map's peak land inside the object at all?
  spaCR has the object masks already (``merged/*.npy``), so this costs nothing.
* :func:`randomization_sanity_check` — Adebayo et al. 2018. Randomise the
  model's weights layer by layer and attribute again. Several popular methods
  return a nearly identical map for a randomised model, which means they are
  edge detectors that happen to be plotted over a classifier. This is the single
  most informative check here and the one most often skipped.
* :func:`method_agreement` — rank correlation between methods on the same image.
  Agreement is weak evidence. Disagreement is strong evidence that no single map
  should be trusted.

Every criterion measures a *different* property and they routinely disagree.
None of them is ground truth, because for attribution there is none.

**Two head shapes, one contract.** A spaCR classifier's head emits either one
logit (binary; class 1 when the logit is positive) or ``C`` logits. Code that
assumes one shape is wrong for the other, and silently so: attributing the raw
logit of a single-logit head always explains *class 1*, so for an image the model
called class 0 the map you get is the map for the class it rejected — the
negation of what you asked for. Every method here goes through
:class:`ClassScoreModel`, which presents a single logit ``z`` as the two-column
view ``[-z, +z]``. Both classes then have a real gradient, ``target`` means the
same thing for both head shapes, and no caller has to know which head it has.

:author: spaCR
"""
from __future__ import annotations

import copy
import math
from dataclasses import dataclass, field
from typing import (Any, Callable, Dict, List, Optional, Sequence, Tuple,
                    Union)

import numpy as np

_trapezoid = getattr(np, 'trapezoid', None) or np.trapz
import torch
import torch.nn as nn
import torch.nn.functional as F

__all__ = [
    "Attribution",
    "ATTRIBUTION_METHODS",
    "MethodSpec",
    "AttributionError",
    "NoSpatialLayerError",
    "UnknownMethodError",
    "ClassScoreModel",
    "attribute",
    "smoothgrad",
    "compare_methods",
    "attention_rollout",
    "chefer_relevance",
    "method_applicability",
    "applicable_methods",
    "architecture_kind",
    "cam_type_choices",
    "cam_type_applicability",
    "resolve_cam_type",
    "CAM_TYPE_ALIASES",
    "LEGACY_CAM_TYPES",
    "list_methods",
    "methods_by_family",
    "conv_layer_names",
    "resolve_layer",
    "recommended_layer",
    "class_scores",
    "Curve",
    "deletion_curve",
    "insertion_curve",
    "faithfulness",
    "pointing_game",
    "pointing_game_rate",
    "SanityCheck",
    "randomization_sanity_check",
    "Agreement",
    "method_agreement",
    "AttributionMapGenerator",
    "NOT_AN_EXPLANATION",
    "CRITERION_CAVEATS",
]



#: Attached to every reported attribution result. Never suppressed.
NOT_AN_EXPLANATION = (
    "An attribution map is not an explanation of causation. It is a number per "
    "pixel produced by a procedure, and every method here will render a "
    "confident-looking map for a model with random weights. Read the deletion "
    "and insertion curves, the pointing-game hit rate and the randomisation "
    "sanity check before believing any of it."
)

#: What each search criterion rewards — and what it cannot see.
CRITERION_CAVEATS: Dict[str, str] = {
    "deletion_auc": (
        "removes the highest-ranked pixels first and averages the class score "
        "along the way, so a LOWER value is better: the score collapsed as soon "
        "as the map's top pixels went. It is confounded by the removal baseline "
        "— blanking a region creates an edge the model has never seen, and part "
        "of the score drop is that artefact rather than the information removed."
    ),
    "insertion_auc": (
        "starts from a blanked image and adds the highest-ranked pixels first, "
        "so a HIGHER value is better. It rewards maps that concentrate on a "
        "small sufficient region, and it systematically favours smooth, blobby "
        "maps over sharp per-pixel ones, which is why it can rank the methods in "
        "the opposite order to deletion."
    ),
    "pointing_game": (
        "asks only whether the map's single brightest pixel falls inside the "
        "object mask. It is cheap and spaCR-specific, it says nothing about the "
        "rest of the map, and it is trivially satisfied by any method biased "
        "towards bright or textured regions when the object is the bright thing "
        "in the frame."
    ),
    "sanity_gap": (
        "one minus the rank correlation between the map from the trained model "
        "and the map from the same model with randomised weights, so a HIGHER "
        "value is better. A method scoring near zero produced the same picture "
        "for a random model and is an edge detector, not an explanation. It "
        "measures dependence on the weights, not correctness."
    ),
}


[docs] class AttributionError(RuntimeError): """Base class for the failures this module reports instead of guessing."""
[docs] class UnknownMethodError(AttributionError): """Raised for a method name that is not registered."""
[docs] class NoSpatialLayerError(AttributionError): """Raised when a CAM is asked of a model that has no spatial layer to hook. A CAM is a weighted sum of one convolutional layer's feature maps. A pure transformer has no such layer, and hooking its patch embedding produces a picture that is not a CAM of anything. Rather than return that picture, the CAM adapters raise this and name the model so the caller can switch to :func:`attention_rollout` or to a gradient / perturbation method, which work on any architecture. """
[docs] class ClassScoreModel(nn.Module): """Present any spaCR classifier head as ``C >= 2`` per-class scores. A head emitting one logit ``z`` is presented as ``[-z, +z]``: column 1 is the evidence for class 1, column 0 the evidence for class 0, and ``argmax`` reproduces the ``z > 0`` rule spaCR's binary models use. The obvious alternative, ``[0, z]``, is exactly equivalent under softmax but has zero gradient for class 0, so every attribution for class 0 would be an all-zero map — a silent, plausible-looking wrong answer. A head emitting ``C > 1`` logits is passed through untouched. :param model: the classifier to wrap. :param n_out: optional raw output width. If omitted, the first forward pass infers it from the wrapped model's output. :ivar n_out: the wrapped model's raw output width (1 or C). :ivar n_classes: the number of classes the wrapper exposes (2 or C). :ivar single_logit: True when the wrapped head emits one logit. """ def __init__(self, model: nn.Module, n_out: Optional[int] = None): """Wrap ``model``, recording whether its head is single-logit.""" super().__init__() self.model = model self.n_out = int(n_out) if n_out is not None else None self.single_logit = None if n_out is None else (int(n_out) == 1) def _note_width(self, raw: torch.Tensor) -> torch.Tensor: """Record the head width the first time a real forward pass reveals it.""" if raw.ndim == 1: raw = raw.unsqueeze(-1) if self.n_out is None: self.n_out = int(raw.shape[-1]) self.single_logit = self.n_out == 1 return raw
[docs] def forward(self, x: torch.Tensor) -> torch.Tensor: """Return ``(B, n_classes)`` scores for ``x``. :param x: input batch passed to the wrapped classifier. """ raw = self._note_width(self.model(x)) if raw.shape[-1] == 1: return torch.cat([-raw, raw], dim=-1) return raw
@property
[docs] def n_classes(self) -> int: """How many classes the wrapper exposes.""" if self.n_out is None: raise AttributionError( "the head width is not known yet — run one forward pass " "through the wrapper first") return 2 if self.n_out == 1 else int(self.n_out)
def _to_batch(image: Any) -> torch.Tensor: """Coerce an image into a float ``(1, C, H, W)`` tensor. :param image: tensor or array shaped ``(H, W)``, ``(C, H, W)`` or ``(1, C, H, W)``. :returns: the batched tensor, detached and float. :raises AttributionError: for a batch of more than one image or a rank the adapters cannot interpret. """ if not isinstance(image, torch.Tensor): image = torch.as_tensor(np.asarray(image)) x = image.detach().float() if x.ndim == 2: x = x.unsqueeze(0).unsqueeze(0) elif x.ndim == 3: x = x.unsqueeze(0) elif x.ndim != 4: raise AttributionError( f"image must be (H, W), (C, H, W) or (1, C, H, W); got shape " f"{tuple(image.shape)}.") if x.shape[0] != 1: raise AttributionError( f"attribute() explains one image at a time; got a batch of " f"{x.shape[0]}. Loop over the batch, or use compare_methods() per " f"image — the analyses downstream are per-image too.") return x
[docs] def class_scores(model: nn.Module, x: torch.Tensor, *, probability: bool = True) -> torch.Tensor: """Per-class scores for ``x``, for either head shape. :param model: the classifier (raw or already wrapped). :param x: input batch ``(B, C, H, W)``. :param probability: return softmax probabilities rather than raw scores. Probabilities are what the deletion / insertion curves track, because a bounded quantity makes their areas comparable across images. :returns: ``(B, n_classes)`` tensor. """ wrapped = model if isinstance(model, ClassScoreModel) else ClassScoreModel(model) with torch.no_grad(): scores = wrapped(x) if probability: return torch.softmax(scores, dim=-1) return scores
def _predicted_class(wrapped: ClassScoreModel, x: torch.Tensor) -> int: """The class the model predicts for ``x``, for either head shape.""" with torch.no_grad(): return int(wrapped(x).argmax(dim=-1)[0]) def _resolve_target(wrapped: ClassScoreModel, x: torch.Tensor, target: Optional[int]) -> int: """Validate ``target`` against the head, defaulting to the prediction.""" predicted = _predicted_class(wrapped, x) if target is None: return predicted target = int(target) n = wrapped.n_classes if not 0 <= target < n: raise AttributionError( f"target={target} is not a class of this model. Its head emits " f"{wrapped.n_out} logit(s), which is " f"{'a binary head with classes 0 and 1' if wrapped.n_out == 1 else f'{n} classes 0..{n - 1}'}." ) return target
[docs] def conv_layer_names(model: nn.Module) -> List[str]: """Every ``Conv2d`` layer name in ``model``, in definition order. :param model: the model to scan. :returns: dotted layer names; empty for a model with no convolutions. """ return [name for name, mod in model.named_modules() if isinstance(mod, nn.Conv2d)]
[docs] def resolve_layer(model: nn.Module, name: str) -> nn.Module: """Resolve a dotted layer name against ``model``. :param model: the model to look in. :param name: dotted module path, e.g. ``'features.2'``. :returns: the submodule. :raises AttributionError: naming the closest available layers. A wrong target layer is the most common way a CAM run dies, and a bare ``AttributeError: 'Sequential' object has no attribute 'conv_b'`` does not tell the user what to type instead. """ modules = dict(model.named_modules()) if name in modules: return modules[name] convs = conv_layer_names(model) candidates = convs or [n for n in modules if n] shown = candidates[-25:] if len(candidates) > 25 else candidates more = (f" (and {len(candidates) - len(shown)} earlier ones)" if len(candidates) > len(shown) else "") kind = "convolutional layers" if convs else "layers" raise AttributionError( f"target layer {name!r} does not exist in this model. Available " f"{kind}{more}: {shown}. The last one, {candidates[-1]!r}, is the " f"usual CAM target." if candidates else f"target layer {name!r} does not exist in this model, and the model " f"has no named submodules to target." )
def _spatial_target_layer(model: nn.Module, layer: Optional[str], model_type: Optional[str]) -> Tuple[str, nn.Module]: """Pick and validate the layer a CAM will hook. :param model: the raw (unwrapped) model. :param layer: dotted layer name, or None to use the last convolution. :param model_type: the architecture name, used in the error message. :returns: ``(name, module)``. :raises NoSpatialLayerError: when the model has no convolution to hook. :raises AttributionError: when a named layer does not exist. """ if layer: return str(layer), resolve_layer(model, str(layer)) name = recommended_layer(model) if name is None: raise NoSpatialLayerError( f"model_type={model_type or type(model).__name__!r} has no Conv2d " f"layer, so there is no feature map for a CAM to weight and no " f"honest CAM to compute. Use method='attention_rollout' if this is " f"a transformer with torch MultiheadAttention blocks, or any of the " f"gradient / perturbation methods (saliency, integrated_gradients, " f"occlusion, feature_ablation), which need no spatial layer.") return name, resolve_layer(model, name) def _is_attention_module(module: nn.Module) -> bool: """Whether a module is an attention block, by type or by class name. ``nn.MultiheadAttention`` catches spaCR's and torch's own blocks; the name test catches torchvision's and timm's, which subclass ``nn.Module`` directly (``WindowAttention``, ``RelativePositionalMultiHeadAttention``, ...). """ if isinstance(module, nn.MultiheadAttention): return True return type(module).__name__.endswith("Attention") def _check_spatial_activation(module: nn.Module, wrapped: ClassScoreModel, x: torch.Tensor, layer_name: str, model_type: Optional[str], allow_pre_attention: bool = False) -> None: """Refuse to CAM a layer that cannot carry a CAM's meaning. Two ways that happens, both of which otherwise render a plausible picture: * the layer's output is not a ``(B, C, H, W)`` feature map. A transformer block emits ``(B, tokens, channels)``; reduced over the channel axis and reshaped it makes an image, and that image is not a CAM of anything. * the layer is a **patch embedding** — it runs before every attention block in the model, so no class-discriminative information has reached it yet. This is the pure-ViT trap: ``recommend_target_layers`` happily returns the patch-embed ``Conv2d`` because it *is* a convolution, and Grad-CAM over it is a picture of local image statistics. Hybrids like MaxViT are unaffected — their MBConv layers run after attention blocks, which is why spaCR's default MaxViT target layer keeps working. :param allow_pre_attention: opt out of the second check when the caller really does want the patch embedding. """ order: List[str] = [] captured: List[torch.Tensor] = [] def _target_hook(_m, _inp, out): """Capture the target layer's output and its position in the pass.""" captured.append(out if isinstance(out, torch.Tensor) else out[0]) order.append("target") def _attn_hook(_m, _inp, _out): """Record that an attention block ran.""" order.append("attention") handles = [module.register_forward_hook(_target_hook)] handles += [m.register_forward_hook(_attn_hook) for m in wrapped.model.modules() if m is not module and _is_attention_module(m)] try: with torch.no_grad(): wrapped(x) finally: for handle in handles: handle.remove() if not captured: raise AttributionError( f"target layer {layer_name!r} never ran during the forward pass, so " f"it cannot be the CAM target. Check that the layer is actually on " f"the path this model takes.") act = captured[0] if act.ndim != 4: raise NoSpatialLayerError( f"target layer {layer_name!r} of " f"model_type={model_type or type(wrapped.model).__name__!r} emits a " f"{act.ndim}-D tensor {tuple(act.shape)}, not a (B, C, H, W) feature " f"map. A CAM over it would be a reshaped token vector drawn as an " f"image, which means nothing. Use method='attention_rollout' for a " f"transformer, or a gradient / perturbation method.") if allow_pre_attention or "attention" not in order: return first_target = order.index("target") first_attention = order.index("attention") if first_target < first_attention: raise NoSpatialLayerError( f"target layer {layer_name!r} of " f"model_type={model_type or type(wrapped.model).__name__!r} runs " f"before every one of the model's {order.count('attention')} " f"attention blocks, so it is the patch embedding: no " f"class-discriminative information has reached it and a CAM over it " f"is a picture of local image statistics, not of what the model " f"attends to. Use method='attention_rollout' for this architecture, " f"a gradient or perturbation method, or name a later layer " f"explicitly. Pass allow_pre_attention=True to override.") @dataclass
[docs] class Attribution: """One attribution map plus everything needed to judge it. :param method: Requested attribution-method name; normally a key of :data:`ATTRIBUTION_METHODS`, and retained verbatim on a skipped-failure placeholder. :param map: Finite 2-D float32 map at input spatial resolution; larger values rank pixels higher, while a skipped failure carries an all-zero placeholder. :param target: Class index this map was requested to explain. :param n_classes: Number of classes exposed by the normalized head, including two for a single-logit binary head. :param single_logit: Whether the underlying model head emitted one binary logit. :param predicted: Class predicted by the model for the attributed input. :param raw: Signed per-channel attribution retained by methods that expose it, or ``None`` when no such representation exists. :param layer: Resolved CAM target-layer name, the caller's requested layer on a skipped failure, or ``None`` when no layer applies. :param family: Registered family (``"cam"``, ``"gradient"``, ``"perturbation"``, or ``"attention"``), or an empty string for an unregistered skipped failure. :param backend: Registered implementation provider (``"torchcam"``, ``"captum"``, or ``"spacr"``), or an empty string for an unregistered skipped failure. :param params: Recorded call options, including the resolved layer and SmoothGrad flag for registry-dispatched methods; this is not a complete expansion of every effective default. :param notes: User-facing caveats, flat-map or single-logit context, or the reason a placeholder method failed. """ method: str map: np.ndarray target: int n_classes: int single_logit: bool predicted: int raw: Optional[np.ndarray] = None layer: Optional[str] = None family: str = "" backend: str = "" params: Dict[str, Any] = field(default_factory=dict) notes: List[str] = field(default_factory=list) @property
[docs] def shape(self) -> Tuple[int, int]: """Spatial shape of the map.""" return tuple(self.map.shape) # type: ignore[return-value]
[docs] def normalized(self) -> np.ndarray: """The map rescaled to ``[0, 1]``; all-zero when the map is flat. A flat map is exactly what a fully suppressed CAM produces, and the unguarded min-max rescale of one is ``0/0``. """ m = np.asarray(self.map, dtype=np.float64) lo = float(m.min()) rng = float(m.max()) - lo if not np.isfinite(rng) or rng <= 0: return np.zeros_like(m, dtype=np.float32) return ((m - lo) / rng).astype(np.float32)
[docs] def peak(self) -> Tuple[int, int]: """``(row, col)`` of the single highest-ranked pixel.""" idx = int(np.argmax(np.asarray(self.map))) return divmod(idx, int(self.map.shape[1]))
[docs] def is_flat(self) -> bool: """True when the map has no variation and therefore ranks nothing.""" m = np.asarray(self.map, dtype=np.float64) return not np.isfinite(m).all() or float(m.max() - m.min()) <= 0.0
def _finite_2d(values: torch.Tensor, size: Tuple[int, int]) -> np.ndarray: """Reduce an attribution tensor to a finite ``(H, W)`` float32 array. Channels are collapsed by summing the absolute per-channel attribution — the convention spaCR's ``saliency_image`` already used — and the result is resampled to the input's spatial size when the method worked at a coarser resolution. """ t = values.detach().float() while t.ndim > 3: t = t[0] if t.ndim == 3: t = t.abs().sum(dim=0) t = torch.nan_to_num(t, nan=0.0, posinf=0.0, neginf=0.0) if tuple(t.shape) != tuple(size): t = F.interpolate(t[None, None], size=size, mode="bilinear", align_corners=False)[0, 0] return t.cpu().numpy().astype(np.float32) _TORCHCAM_CLASSES = { "gradcam": "GradCAM", "gradcam_pp": "GradCAMpp", "scorecam": "ScoreCAM", "xgradcam": "XGradCAM", "layercam": "LayerCAM", } def _torchcam_cam(spec: "MethodSpec", wrapped: ClassScoreModel, x: torch.Tensor, target: int, layer: Optional[str], model_type: Optional[str], **kw) -> Tuple[np.ndarray, None, str, List[str]]: """Run one torchcam CAM extractor against the wrapped model. torchcam owns the maths for all five variants; this adapter only picks the target layer, keeps the extractor's hooks off the model afterwards, and resamples the CAM to the input resolution. The extractor is built on the *wrapped* model so ``class_idx`` means the same class for a single-logit head as for a C-logit one. """ try: import torchcam.methods as tcm except (ImportError, ModuleNotFoundError) as exc: raise AttributionError( f"{spec.name} requires the optional torchcam backend. Install it " "with `pip install 'spacr[attribution]'`, or choose eigencam, " "saliency, integrated_gradients, occlusion, or another " "non-torchcam attribution method." ) from exc layer_name, module = _spatial_target_layer(wrapped.model, layer, model_type) _check_spatial_activation(module, wrapped, x, layer_name, model_type, bool(kw.get("allow_pre_attention", False))) cls = getattr(tcm, _TORCHCAM_CLASSES[spec.name]) extractor_kw: Dict[str, Any] = {"input_shape": tuple(x.shape[1:])} if spec.name == "scorecam": extractor_kw["batch_size"] = int(kw.get("batch_size", 8)) with cls(wrapped, target_layer=module, **extractor_kw) as extractor: scores = wrapped(x) cams = extractor(int(target), scores) cam = cams[0] return (_finite_2d(cam, (int(x.shape[-2]), int(x.shape[-1]))), None, layer_name, []) def _eigen_cam(spec: "MethodSpec", wrapped: ClassScoreModel, x: torch.Tensor, target: int, layer: Optional[str], model_type: Optional[str], **kw) -> Tuple[np.ndarray, None, str, List[str]]: """Eigen-CAM: the first principal component of the target layer's activations. torchcam 0.4 ships every CAM variant except this one, so the eight lines of SVD are here. Eigen-CAM uses no gradients and no class index at all, which is worth knowing before reading one: **the map is identical for every class**, so it cannot be evidence that the model separated the classes. It is included because that same property makes it a useful control — a class-conditional method whose map matches Eigen-CAM is not being class-conditional. """ layer_name, module = _spatial_target_layer(wrapped.model, layer, model_type) _check_spatial_activation(module, wrapped, x, layer_name, model_type, bool(kw.get("allow_pre_attention", False))) captured: List[torch.Tensor] = [] def _hook(_m, _inp, out): """Capture the target layer's feature maps.""" captured.append(out if isinstance(out, torch.Tensor) else out[0]) handle = module.register_forward_hook(_hook) try: with torch.no_grad(): wrapped(x) finally: handle.remove() act = captured[0][0] c, h, w = act.shape flat = act.reshape(c, h * w).T flat = flat - flat.mean(dim=0, keepdim=True) try: _u, _s, vh = torch.linalg.svd(flat.double(), full_matrices=False) proj = (flat.double() @ vh[0]).reshape(h, w) except Exception: proj = act.mean(dim=0).double() if float(proj.sum()) < 0: proj = -proj proj = proj - proj.min() return (_finite_2d(proj.float(), (int(x.shape[-2]), int(x.shape[-1]))), None, layer_name, ["Eigen-CAM ignores the class and the gradients entirely: the same " "map is returned for every target, so it cannot show that the " "model distinguished the classes."]) def _captum_baseline(kind: Any, x: torch.Tensor) -> torch.Tensor: """Build the reference input a baseline-dependent method integrates from. :param kind: ``'zero'``, ``'mean'``, ``'blur'``, ``'uniform'``, a number, or a tensor broadcastable to ``x``. :returns: the baseline tensor. :raises AttributionError: for an unknown name. """ if isinstance(kind, torch.Tensor): return kind.to(x.dtype).expand_as(x).clone() if isinstance(kind, (int, float)) and not isinstance(kind, bool): return torch.full_like(x, float(kind)) name = str(kind or "zero").lower() if name in ("zero", "zeros", "black", "none"): return torch.zeros_like(x) if name in ("mean", "channel_mean"): return x.mean(dim=(-2, -1), keepdim=True).expand_as(x).clone() if name in ("blur", "blurred"): return _blur(x) if name in ("uniform", "random", "noise"): return torch.rand_like(x) * (x.max() - x.min()) + x.min() raise AttributionError( f"unknown baseline {kind!r}; use 'zero', 'mean', 'blur', 'uniform', a " f"number, or a tensor. The baseline is not cosmetic: integrated " f"gradients attributes the difference between the input and this " f"reference, so a different baseline is a different question.") def _blur(x: torch.Tensor, sigma: float = 5.0) -> torch.Tensor: """Gaussian-blur a batch with a separable depthwise convolution. The radius is clamped to the image: reflection padding wider than the dimension it pads raises, and spaCR's object crops are routinely smaller than the 3-sigma radius a 224-pixel default assumes. """ limit = max(1, min(int(x.shape[-2]), int(x.shape[-1])) - 1) radius = max(1, min(int(round(3.0 * float(sigma))), limit)) coords = torch.arange(-radius, radius + 1, dtype=x.dtype, device=x.device) kernel = torch.exp(-(coords ** 2) / (2.0 * float(sigma) ** 2)) kernel = kernel / kernel.sum() c = int(x.shape[1]) out = F.conv2d(F.pad(x, (radius, radius, 0, 0), mode="reflect"), kernel.view(1, 1, 1, -1).expand(c, 1, 1, -1), groups=c) out = F.conv2d(F.pad(out, (0, 0, radius, radius), mode="reflect"), kernel.view(1, 1, -1, 1).expand(c, 1, -1, 1), groups=c) return out def _captum_attribute(spec: "MethodSpec", wrapped: ClassScoreModel, x: torch.Tensor, target: int, layer: Optional[str], model_type: Optional[str], **kw ) -> Tuple[np.ndarray, np.ndarray, None, List[str]]: """Run one captum attributor against the wrapped model. captum owns the maths; this adapter supplies the arguments each method needs, converts spaCR's parameter names, and turns captum's signed per-channel output into the module's ``(H, W)`` contract without losing the signed form. """ ca = _import_captum_attr() notes: List[str] = [] inp = x.clone().requires_grad_(True) name = spec.name kwargs: Dict[str, Any] = {"target": int(target)} if name == "saliency": attributor: Any = ca.Saliency(wrapped) kwargs["abs"] = bool(kw.get("abs", True)) elif name == "integrated_gradients": attributor = ca.IntegratedGradients(wrapped) kwargs["baselines"] = _captum_baseline(kw.get("baseline", "zero"), x) kwargs["n_steps"] = int(kw.get("n_steps", kw.get("ig_steps", 50))) if kwargs["n_steps"] < 2: raise AttributionError( f"integrated gradients needs at least 2 steps to integrate " f"anything, got n_steps={kwargs['n_steps']}.") elif name == "guided_backprop": attributor = ca.GuidedBackprop(wrapped) notes.append( "Guided backprop is the method most often reported as failing the " "randomisation sanity check: its ReLU clamping recovers image edges " "almost independently of the weights. Run " "randomization_sanity_check() before reading anything into it.") elif name == "input_x_gradient": attributor = ca.InputXGradient(wrapped) elif name == "deeplift": attributor = ca.DeepLift(wrapped) kwargs["baselines"] = _captum_baseline(kw.get("baseline", "zero"), x) elif name == "gradient_shap": attributor = ca.GradientShap(wrapped) kwargs["baselines"] = _shap_baselines(x, kw) kwargs["n_samples"] = max(1, int(kw.get("shap_samples", 20))) kwargs["stdevs"] = float(_sigma_to_stdev(kw.get("shap_sigma", 0.09), x)) notes.append( f"GradientSHAP averaged {kwargs['n_samples']} noisy draws between " f"the image and {int(kwargs['baselines'].shape[0])} references.") elif name == "deeplift_shap": attributor = ca.DeepLiftShap(wrapped) kwargs["baselines"] = _shap_baselines(x, kw) elif name == "occlusion": attributor = ca.Occlusion(wrapped) window = int(kw.get("window", kw.get("occlusion_window", 8))) stride = int(kw.get("stride", kw.get("occlusion_stride", max(1, window // 2)))) c, h, w = int(x.shape[1]), int(x.shape[2]), int(x.shape[3]) window = max(1, min(window, h, w)) stride = max(1, min(stride, window)) kwargs["sliding_window_shapes"] = (c, window, window) kwargs["strides"] = (c, stride, stride) kwargs["baselines"] = _captum_baseline(kw.get("baseline", "zero"), x) kwargs["show_progress"] = False elif name == "feature_ablation": attributor = ca.FeatureAblation(wrapped) block = int(kw.get("block", kw.get("occlusion_window", 8))) kwargs["feature_mask"] = _block_mask(x, block) kwargs["baselines"] = _captum_baseline(kw.get("baseline", "zero"), x) kwargs["show_progress"] = False else: raise UnknownMethodError( f"{name!r} is registered as a captum method but this adapter has " f"no branch for it; the registry and the adapter disagree.") import warnings as _warnings try: with _warnings.catch_warnings(record=True) as caught: _warnings.simplefilter("always") if spec.smoothed: tunnel = ca.NoiseTunnel(attributor) kwargs["nt_type"] = "smoothgrad" kwargs["nt_samples"] = int(kw.get("n_samples", 25)) kwargs["stdevs"] = float( _sigma_to_stdev(kw.get("sigma", 0.15), x)) result = tunnel.attribute(inp, **kwargs) else: result = attributor.attribute(inp, **kwargs) except RuntimeError as exc: if "more than once" in str(exc) or "required for DeepLift" in str(exc): raise AttributionError( f"{name} cannot run on this model: it reuses one activation " f"module (typically a single nn.ReLU(inplace=True)) at several " f"points in the forward pass, and DeepLIFT needs a distinct " f"module per activation to attach its rescale rule to. " f"torchvision's ResNets are built this way. Use " f"'integrated_gradients' (the same axiomatic family, no such " f"requirement), 'input_x_gradient', or a perturbation method. " f"captum reported: {exc}") from exc raise for warning in caught: notes.append(f"captum warning: {warning.message}") raw = result.detach()[0].cpu().numpy().astype(np.float32) return (_finite_2d(result, (int(x.shape[-2]), int(x.shape[-1]))), raw, None, notes) def _block_mask(x: torch.Tensor, block: int) -> torch.Tensor: """Group pixels into ``block × block`` tiles for feature ablation. Ablating one pixel at a time on a 224² image is 50 176 forward passes and tells you almost nothing, because a single pixel changes no convolutional response enough to move the score. Tiles are the usable form. """ h, w = int(x.shape[-2]), int(x.shape[-1]) block = max(1, min(int(block), h, w)) rows = torch.arange(h) // block cols = torch.arange(w) // block n_cols = int(cols.max()) + 1 ids = (rows[:, None] * n_cols + cols[None, :]).long() return ids[None, None].expand(1, int(x.shape[1]), h, w).contiguous() def _sigma_to_stdev(sigma: float, x: torch.Tensor) -> float: """Convert a SmoothGrad noise fraction into an absolute standard deviation. ``sigma`` is a fraction of the input's dynamic range, matching the ``stdev_spread`` convention of the SmoothGrad paper and of :class:`spacr.deep_spacr.SmoothGrad`. An absolute standard deviation would mean something different for a ``[0, 1]`` image and a z-scored one. """ span = float(x.max() - x.min()) if span <= 0: span = 1.0 return max(float(sigma) * span, 1e-12) def _ask_for_attention_weights(blocks): """Make each MHA block return its weights. Returns an undo callable. The override is by keyword only. `nn.MultiheadAttention.forward` takes `need_weights` as its fifth positional parameter, and a caller that passed it positionally would have its argument silently replaced -- so those are left exactly as they are, and the existing "none returned attention weights" refusal still covers them. `average_attn_weights=False` because rollout fuses the heads itself: 'mean', 'max' and 'min' are the caller's choice, and a block that has already averaged offers only one of the three. """ originals = [] def _wrap(block): """Install the attention-weight adapter and return the original call.""" original = block.forward def forward(*args, **kwargs): """Request per-head weights unless a positional request owns them.""" if len(args) < 5 and "need_weights" not in kwargs: kwargs = dict(kwargs) kwargs["need_weights"] = True kwargs.setdefault("average_attn_weights", False) elif kwargs.get("need_weights") is False: kwargs = dict(kwargs) kwargs["need_weights"] = True kwargs.setdefault("average_attn_weights", False) return original(*args, **kwargs) block.forward = forward return original for block in blocks: originals.append((block, _wrap(block))) def restore(): """Restore every block's exact original bound forward method.""" for block, original in originals: block.forward = original return restore
[docs] def attention_rollout(model: nn.Module, image: Any, *, target: Optional[int] = None, head_fusion: str = "mean", discard_ratio: float = 0.0, model_type: Optional[str] = None) -> Attribution: """Attention rollout (Abnar & Zuidema 2020) for transformer backbones. A CAM needs a convolutional feature map. A pure transformer has none, so this is the substitute: the per-layer attention matrices are averaged over heads, mixed with the identity to account for the residual stream, row- normalised and multiplied together, giving how much each input token contributes to the class token. It reads spaCR's :class:`torch.nn.MultiheadAttention` blocks, which return their attention weights from ``forward``. Backbones whose attention is a fused kernel (timm's ViT via ``scaled_dot_product_attention``) expose no weights, and this raises rather than inventing a map. **Rollout is not class-conditional.** The result is the same for every ``target``: it describes where information flowed, not what the model concluded. It cannot show the model separated your classes; a gradient or perturbation method can. :param model: the transformer classifier. :param image: one image, ``(C, H, W)`` or ``(1, C, H, W)``. :param target: recorded on the result; does not change the map. :param head_fusion: ``'mean'``, ``'max'`` or ``'min'`` over attention heads. :param discard_ratio: fraction of the lowest attention weights zeroed per layer before rollout, which sharpens the map and is pure cosmetics. :param model_type: architecture name for the error messages. :returns: the :class:`Attribution`. :raises NoSpatialLayerError: when the model exposes no attention weights. """ x = _to_batch(image) wrapped = ClassScoreModel(model) target = _resolve_target(wrapped, x, target) predicted = _predicted_class(wrapped, x) blocks = [m for m in model.modules() if isinstance(m, nn.MultiheadAttention)] if not blocks: convs = conv_layer_names(model) raise NoSpatialLayerError( f"model_type={model_type or type(model).__name__!r} exposes no " f"torch.nn.MultiheadAttention block, so there are no attention " f"weights to roll out. " + (f"It does have convolutional layers ({convs[-1]!r} last), so a " f"CAM method applies instead." if convs else "It has no Conv2d layer either, so no CAM applies — use a " "gradient or perturbation method (saliency, " "integrated_gradients, occlusion), which need neither.") + " Backbones whose attention is a fused kernel never expose " "weights; no map is returned rather than a meaningless one.") captured: List[torch.Tensor] = [] def _hook(_m, _inp, out): """Capture the (B, L, S) attention weights an MHA block returns.""" if isinstance(out, (tuple, list)) and len(out) > 1 and \ isinstance(out[1], torch.Tensor): captured.append(out[1].detach()) handles = [b.register_forward_hook(_hook) for b in blocks] restore = _ask_for_attention_weights(blocks) try: with torch.no_grad(): wrapped(x) finally: restore() for h in handles: h.remove() if not captured: raise NoSpatialLayerError( f"model_type={model_type or type(model).__name__!r} has " f"MultiheadAttention blocks but none returned attention weights " f"(they were called with need_weights=False, or the fused kernel " f"path was taken). Rollout has nothing to roll; use a gradient or " f"perturbation method.") fuse = {"mean": lambda a: a.mean(dim=1), "max": lambda a: a.amax(dim=1), "min": lambda a: a.amin(dim=1)} if head_fusion not in fuse: raise AttributionError( f"head_fusion must be one of {sorted(fuse)}, got {head_fusion!r}.") rolled: Optional[torch.Tensor] = None for attn in captured: a = attn.double() if a.ndim == 4: a = fuse[head_fusion](a) a = a[0] if a.shape[0] != a.shape[1]: raise NoSpatialLayerError( f"an attention block returned a non-square {tuple(a.shape)} " f"matrix, so rollout (which composes square token-to-token " f"maps) does not apply to this architecture.") if discard_ratio > 0: k = int(a.numel() * float(discard_ratio)) if k > 0: flat = a.flatten() cut = flat.kthvalue(k).values a = torch.where(a <= cut, torch.zeros_like(a), a) a = a + torch.eye(a.shape[0], dtype=a.dtype) a = a / a.sum(dim=-1, keepdim=True).clamp_min(1e-12) rolled = a if rolled is None else a @ rolled n_tokens = int(rolled.shape[0]) grid = int(round(math.sqrt(n_tokens - 1))) if grid * grid == n_tokens - 1: weights = rolled[0, 1:] else: grid = int(round(math.sqrt(n_tokens))) if grid * grid != n_tokens: raise NoSpatialLayerError( f"{n_tokens} attention tokens do not form a square patch grid " f"(with or without a class token), so they cannot be laid back " f"out over the image.") weights = rolled.mean(dim=0) heat = weights.reshape(grid, grid).float() return Attribution( method="attention_rollout", map=_finite_2d(heat, (int(x.shape[-2]), int(x.shape[-1]))), target=int(target), n_classes=wrapped.n_classes, single_logit=bool(wrapped.single_logit), predicted=int(predicted), raw=None, layer=None, family="attention", backend="spacr", params={"head_fusion": head_fusion, "discard_ratio": discard_ratio}, notes=[f"Rolled out {len(captured)} attention blocks.", "Attention rollout is not class-conditional: the map is " "identical for every target, so it cannot show that the model " "distinguished your classes."])
def _attention_adapter(spec: "MethodSpec", wrapped: ClassScoreModel, x: torch.Tensor, target: int, layer: Optional[str], model_type: Optional[str], **kw): """Registry entry point for rollout, returning the adapter tuple.""" att = attention_rollout(wrapped.model, x, target=target, model_type=model_type, head_fusion=str(kw.get("head_fusion", "mean")), discard_ratio=float(kw.get("discard_ratio", 0.0))) return att.map, None, None, list(att.notes) def _layer_activation_and_gradient(wrapped: ClassScoreModel, x: torch.Tensor, module: nn.Module, target: int ) -> Tuple[torch.Tensor, torch.Tensor]: """The target layer's feature maps and the class score's gradient on them. The activation is copied at forward time, because a later in-place ReLU would otherwise overwrite the tensor that was hooked. The gradient is taken with :func:`torch.autograd.grad` against the input, so no parameter's ``.grad`` is touched. :returns: ``(activation, gradient)``, both ``(1, C, h, w)``. :raises AttributionError: when the layer never received a gradient. """ captured: Dict[str, torch.Tensor] = {} def _keep_gradient(grad): """Store the gradient that reaches the hooked feature map.""" captured["grad"] = grad.detach().clone() def _hook(_m, _inp, out): """Copy the feature maps and ask for their gradient.""" tensor = out if isinstance(out, torch.Tensor) else out[0] captured["act"] = tensor.detach().clone() if tensor.requires_grad: tensor.register_hook(_keep_gradient) handle = module.register_forward_hook(_hook) try: with torch.enable_grad(): inp = x.detach().clone().requires_grad_(True) score = wrapped(inp)[0, int(target)] torch.autograd.grad(score, inp, allow_unused=True) finally: handle.remove() if "act" not in captured or "grad" not in captured: raise AttributionError( "the target layer received no gradient from the class score, so a " "gradient-weighted CAM has nothing to weight. Check that the layer " "is on the path to the classifier head.") return captured["act"], captured["grad"] def _hires_cam(spec: "MethodSpec", wrapped: ClassScoreModel, x: torch.Tensor, target: int, layer: Optional[str], model_type: Optional[str], **kw) -> Tuple[np.ndarray, None, str, List[str]]: """HiRes-CAM (Draelos & Carin 2020): gradient times activation, element-wise. Grad-CAM first averages each channel's gradient over space and then weights the whole channel by that one number; HiRes-CAM keeps the gradient at every position. For a network that ends in global average pooling and a linear head the two coincide, and where they differ HiRes-CAM is the one that is provably tied to the score: summed over the map it is the first-order contribution of the layer to the class score, which Grad-CAM's channel-averaging does not guarantee. """ layer_name, module = _spatial_target_layer(wrapped.model, layer, model_type) _check_spatial_activation(module, wrapped, x, layer_name, model_type, bool(kw.get("allow_pre_attention", False))) act, grad = _layer_activation_and_gradient(wrapped, x, module, target) cam = torch.relu((act[0] * grad[0]).sum(dim=0)) return (_finite_2d(cam, (int(x.shape[-2]), int(x.shape[-1]))), None, layer_name, []) def _ablation_cam(spec: "MethodSpec", wrapped: ClassScoreModel, x: torch.Tensor, target: int, layer: Optional[str], model_type: Optional[str], **kw ) -> Tuple[np.ndarray, None, str, List[str]]: """Ablation-CAM (Desai & Ramaswamy 2020): weight each channel by the score it costs. Every channel of the target layer is zeroed in turn and the class score re-measured; the channel's weight is the fractional score drop. No gradient is used, so the map survives the saturated gradients that make Grad-CAM go quiet. It costs one forward pass per channel, batched ``batch_size`` at a time. """ layer_name, module = _spatial_target_layer(wrapped.model, layer, model_type) _check_spatial_activation(module, wrapped, x, layer_name, model_type, bool(kw.get("allow_pre_attention", False))) captured: List[torch.Tensor] = [] state: Dict[str, Any] = {"channels": None} def _ablate(_m, _inp, out): """Record the clean feature maps, or zero one channel per batch row.""" tensor = out if isinstance(out, torch.Tensor) else out[0] channels = state["channels"] if channels is None: captured.append(tensor.detach().clone()) return None tensor = tensor.clone() tensor[torch.arange(len(channels)), channels] = 0 if isinstance(out, torch.Tensor): return tensor return (tensor,) + tuple(out[1:]) handle = module.register_forward_hook(_ablate) try: with torch.no_grad(): base = float(wrapped(x)[0, int(target)]) act = captured[0][0] n_channels = int(act.shape[0]) drops = torch.zeros(n_channels, dtype=torch.float64) step = max(1, int(kw.get("batch_size", 32))) for start in range(0, n_channels, step): channels = torch.arange(start, min(n_channels, start + step)) state["channels"] = channels batch = x.expand(len(channels), *x.shape[1:]) scores = wrapped(batch)[:, int(target)].double().cpu() drops[channels] = base - scores finally: state["channels"] = None handle.remove() weights = (drops / (abs(base) + 1e-7)).to(act.dtype).to(act.device) cam = torch.relu((weights[:, None, None] * act).sum(dim=0)) return (_finite_2d(cam, (int(x.shape[-2]), int(x.shape[-1]))), None, layer_name, [f"Ablation-CAM re-scored the image {n_channels} times, once per " f"channel of {layer_name!r}."]) def _shap_baselines(x: torch.Tensor, kw: Dict[str, Any]) -> torch.Tensor: """The reference distribution GradientSHAP and DeepSHAP average over. SHAP values are defined against a distribution of references, not one image; a single black baseline turns GradientSHAP back into integrated gradients. The default draws on three cheap references built from the input itself: black, a blurred copy and the per-channel mean. """ names = kw.get("shap_baselines") or ("zero", "blur", "mean") if isinstance(names, str): names = [part.strip() for part in names.replace(",", " ").split() if part.strip()] refs = [_captum_baseline(name, x) for name in names] if len(refs) < 2: refs.append(_captum_baseline("blur", x)) return torch.cat(refs, dim=0) def _import_captum_attr(): """captum's ``attr`` package, or an install instruction.""" try: import captum.attr as ca except (ImportError, ModuleNotFoundError) as exc: raise AttributionError( "this attribution method needs the captum package, which is not " "installed. Install it with `pip install captum`, or choose a CAM " "method (eigencam, hirescam, ablation_cam), which need only torch." ) from exc return ca
[docs] def architecture_kind(model: Optional[nn.Module] = None, model_type: Optional[str] = None) -> str: """Classify a backbone into what the attribution families can use. :param model: the model itself, when it is loaded. :param model_type: the architecture name spaCR stores in settings (``'resnet50'``, ``'maxvit_t'``, ``'vit_b_16'`` ...). :returns: ``'vit'`` (global self-attention over a patch grid), ``'swin'`` (shifted-window attention), ``'hybrid'`` (MaxViT: convolutions after attention), ``'cnn'``, ``'other'`` (a loaded model with neither convolution nor attention) or ``'unknown'`` (nothing to go on). """ name = str(model_type or "").strip().lower() if name.startswith("vit"): return "vit" if name.startswith("swin"): return "swin" if name.startswith("maxvit"): return "hybrid" if model is not None: kinds = {type(m).__name__ for m in model.modules()} if any(k.startswith("MaxVit") for k in kinds): return "hybrid" if any(k.startswith("Swin") or k == "ShiftedWindowAttention" for k in kinds): return "swin" if any(isinstance(m, nn.MultiheadAttention) for m in model.modules()): return "vit" if conv_layer_names(model): return "cnn" return "other" if name: return "cnn" return "unknown"
def _reuses_relu_modules(model: Optional[nn.Module], model_type: Optional[str]) -> bool: """Whether the backbone calls one ReLU module at several points. torchvision's ResNet blocks do, and captum's DeepLIFT rules refuse such a model because they attach one rescale rule per activation module. """ name = str(model_type or "").strip().lower() if name.startswith(("resnet", "wide_resnet", "resnext")): return True if model is not None: return any(type(m).__name__ in ("BasicBlock", "Bottleneck") for m in model.modules()) return False
[docs] def method_applicability(method: str, *, model: Optional[nn.Module] = None, model_type: Optional[str] = None ) -> Tuple[bool, str]: """Whether a registered method can give a meaningful map for this backbone. Decided from the architecture and from which optional backends are installed, without running the model, so a settings form can grey the methods that do not apply and show the reason. :param method: a key of :data:`ATTRIBUTION_METHODS`. :param model: the loaded model, when there is one. :param model_type: the architecture name, when there is no model. :returns: ``(applies, reason)``; the reason is empty when it applies. :raises UnknownMethodError: for an unregistered method name. """ import importlib.util spec = ATTRIBUTION_METHODS.get(str(method)) if spec is None: raise UnknownMethodError( f"unknown attribution method {method!r}; registered: " f"{sorted(ATTRIBUTION_METHODS)}") if spec.backend == "torchcam" and importlib.util.find_spec("torchcam") is None: return False, ("needs the optional torchcam package: " "pip install 'spacr[attribution]'") if spec.backend == "captum" and importlib.util.find_spec("captum") is None: return False, "needs the captum package: pip install captum" kind = architecture_kind(model, model_type) if spec.family == "cam": if kind in ("vit", "swin"): return False, ( "a pure transformer's only convolution is the patch embedding, " "which runs before every attention block, so a CAM over it " "shows local image statistics; use chefer (ViT), saliency or a " "SHAP method") if kind == "other": return False, ("the model has no Conv2d feature map for a CAM to " "weight") if spec.family == "attention" and kind not in ("vit", "unknown"): if kind == "hybrid": return False, ( "MaxViT attends inside local windows and grids with a relative " "position bias and has no class token, so its attention cannot " "be composed into one token map; use hirescam or gradcam on " "its MBConv layers") if kind == "swin": return False, ( "Swin attends inside shifted local windows with no class token, " "so its attention cannot be composed into one token map; use " "saliency or a SHAP method") return False, ("the model has no self-attention blocks; attention " "methods apply to Vision Transformers only") if spec.name in ("deeplift", "deeplift_shap") and _reuses_relu_modules( model, model_type): return False, ( "ResNet-style blocks call one ReLU module twice, which captum's " "DeepLIFT rules refuse; use gradient_shap or integrated_gradients") return True, ""
[docs] def applicable_methods(model: Optional[nn.Module] = None, model_type: Optional[str] = None ) -> Dict[str, Tuple[bool, str]]: """:func:`method_applicability` for every registered method, by name.""" return {name: method_applicability(name, model=model, model_type=model_type) for name in sorted(ATTRIBUTION_METHODS)}
#: ``cam_type`` names that ``generate_activation_map`` serves with spaCR's own #: pre-registry generators rather than through :data:`ATTRIBUTION_METHODS`. LEGACY_CAM_TYPES: Tuple[str, ...] = ("gradcam", "gradcam_pp", "saliency_image", "saliency_channel") #: ``cam_type`` spellings for registry methods whose own name is taken by a #: legacy generator: the torchcam Grad-CAM and Grad-CAM++. CAM_TYPE_ALIASES: Dict[str, str] = { "torchcam_gradcam": "gradcam", "torchcam_gradcam_pp": "gradcam_pp", }
[docs] def resolve_cam_type(cam_type: str) -> Optional[str]: """The registry method a ``cam_type`` names, or None for a legacy one. :param cam_type: a ``cam_type`` setting value: a legacy generator, a :data:`CAM_TYPE_ALIASES` spelling or a key of :data:`ATTRIBUTION_METHODS`. :returns: the registry method name, or None for a legacy generator. :raises UnknownMethodError: for a name that is neither. """ name = str(cam_type) if name in CAM_TYPE_ALIASES: return CAM_TYPE_ALIASES[name] if name in LEGACY_CAM_TYPES: return None if name in ATTRIBUTION_METHODS: return name raise UnknownMethodError( f"unknown cam_type {name!r}; choose one of {list(cam_type_choices())}")
[docs] def cam_type_choices() -> Tuple[str, ...]: """Every ``cam_type`` the Activation Maps settings can name, in menu order.""" registry = [name for name in sorted(ATTRIBUTION_METHODS) if name not in CAM_TYPE_ALIASES.values()] return LEGACY_CAM_TYPES + tuple(CAM_TYPE_ALIASES) + tuple(registry)
[docs] def cam_type_applicability(cam_type: str, *, model: Optional[nn.Module] = None, model_type: Optional[str] = None ) -> Tuple[bool, str]: """:func:`method_applicability` for a ``cam_type`` setting value. The legacy Grad-CAM generators need a spatial layer like any CAM; the legacy saliency maps apply to every backbone. :param cam_type: the ``cam_type`` setting value, as :func:`resolve_cam_type` accepts it. :param model: the loaded model, when there is one. :param model_type: the architecture name, when there is no model. :returns: ``(applies, reason)``; the reason is empty when it applies. :raises UnknownMethodError: for a name :func:`resolve_cam_type` refuses. """ registry_name = resolve_cam_type(cam_type) if registry_name is not None: return method_applicability(registry_name, model=model, model_type=model_type) if str(cam_type).startswith("saliency"): return True, "" return method_applicability("eigencam", model=model, model_type=model_type)
[docs] def chefer_relevance(model: nn.Module, image: Any, *, target: Optional[int] = None, model_type: Optional[str] = None) -> Attribution: """Class-specific transformer relevance (Chefer, Gur & Wolf 2021). Rollout multiplies raw attention and so gives one map whatever class is asked about. This weights every head's attention by the class score's gradient on it, keeps the positive part, averages the heads and accumulates the result through the residual stream, ``R <- R + A_bar R`` from the identity. The class token's row of ``R`` is the relevance of each patch *for the target class*, so asking for the other class gives a different map. It reads :class:`torch.nn.MultiheadAttention` blocks (torchvision's ViT and spaCR's own), asking each for per-head weights the way :func:`attention_rollout` does. :param model: the transformer classifier. :param image: one image, ``(C, H, W)`` or ``(1, C, H, W)``. :param target: class to explain; defaults to the prediction. :param model_type: architecture name for the error messages. :returns: the :class:`Attribution`. :raises NoSpatialLayerError: when the model has no attention blocks whose weights sit on the gradient path, or its tokens do not form a grid. """ x = _to_batch(image) was_training = model.training model.eval() wrapped = ClassScoreModel(model) target = _resolve_target(wrapped, x, target) predicted = _predicted_class(wrapped, x) blocks = [m for m in model.modules() if isinstance(m, nn.MultiheadAttention)] if not blocks: _ok, reason = method_applicability("chefer", model=model, model_type=model_type) raise NoSpatialLayerError( f"model_type={model_type or type(model).__name__!r}: Chefer " f"relevance is not applicable — " f"{reason or 'no torch.nn.MultiheadAttention block to read'}.") captured: List[torch.Tensor] = [] def _hook(_m, _inp, out): """Keep the attention weights an MHA block returns, graph attached.""" if isinstance(out, (tuple, list)) and len(out) > 1 and \ isinstance(out[1], torch.Tensor): captured.append(out[1]) handles = [b.register_forward_hook(_hook) for b in blocks] restore = _ask_for_attention_weights(blocks) try: with torch.enable_grad(): inp = x.detach().clone().requires_grad_(True) score = wrapped(inp)[0, int(target)] sources = [] for attn in captured: base = attn._base use_base = (base is not None and base.requires_grad and base.numel() == attn.numel()) sources.append(base if use_base else attn) grads = (torch.autograd.grad(score, sources, allow_unused=True) if sources else ()) finally: restore() for handle in handles: handle.remove() if was_training: model.train() if not captured or any(g is None for g in grads): raise NoSpatialLayerError( f"model_type={model_type or type(model).__name__!r} has attention " f"blocks, but their weights are not on the gradient path to the " f"class score (a fused attention kernel), so there is nothing for " f"Chefer relevance to weight. Use saliency or a SHAP method.") relevance: Optional[torch.Tensor] = None for attn, grad in zip(captured, grads): a = attn.detach().double() g = grad.detach().double().reshape(attn.shape) if a.ndim == 3: a, g = a[:, None], g[:, None] cam = (g[0] * a[0]).clamp_min(0).mean(dim=0) if cam.shape[0] != cam.shape[1]: raise NoSpatialLayerError( f"an attention block returned a non-square {tuple(cam.shape)} " f"matrix, so token relevance cannot be propagated through it.") if relevance is None: relevance = torch.eye(cam.shape[0], dtype=cam.dtype) relevance = relevance + cam @ relevance n_tokens = int(relevance.shape[0]) grid = int(round(math.sqrt(n_tokens - 1))) if grid * grid == n_tokens - 1: weights = relevance[0, 1:] note = "Read from the class token's row of the relevance matrix." else: grid = int(round(math.sqrt(n_tokens))) if grid * grid != n_tokens: raise NoSpatialLayerError( f"{n_tokens} attention tokens do not form a square patch grid " f"(with or without a class token), so they cannot be laid back " f"out over the image.") weights = (relevance - torch.eye(n_tokens, dtype=relevance.dtype) ).mean(dim=0) note = ("No class token: relevance was averaged over every token's " "row, the analogue for a mean-pooled head.") heat = weights.reshape(grid, grid).float() result = Attribution( method="chefer", map=_finite_2d(heat, (int(x.shape[-2]), int(x.shape[-1]))), target=int(target), n_classes=wrapped.n_classes, single_logit=bool(wrapped.single_logit), predicted=int(predicted), raw=None, layer=None, family="attention", backend="spacr", params={}, notes=[f"Propagated relevance through {len(captured)} " f"attention blocks.", note]) return result
def _chefer_adapter(spec: "MethodSpec", wrapped: ClassScoreModel, x: torch.Tensor, target: int, layer: Optional[str], model_type: Optional[str], **kw): """Registry entry point for Chefer relevance, returning the adapter tuple.""" att = chefer_relevance(wrapped.model, x, target=target, model_type=model_type) return att.map, None, None, list(att.notes) @dataclass(frozen=True)
[docs] class MethodSpec: """One registered attribution method. :param name: registry key callers pass as ``method=``. :param family: method family—``"cam"``, ``"gradient"``, ``"shap"``, ``"perturbation"``, or ``"attention"``. :param backend: implementation provider—``"torchcam"``, ``"captum"``, or ``"spacr"``. :param fn: adapter callable that computes this method's attribution. :param needs_layer: whether the method requires a spatial target layer. :param smoothed: whether the Captum adapter wraps the base attributor in a SmoothGrad noise tunnel. :param description: concise user-facing explanation suitable for method selectors. """ name: str family: str backend: str fn: Callable[..., Any] needs_layer: bool = False smoothed: bool = False description: str = ""
[docs] def smoothgrad_variant(self) -> "MethodSpec": """This method with SmoothGrad averaging turned on.""" return MethodSpec(name=self.name, family=self.family, backend=self.backend, fn=self.fn, needs_layer=self.needs_layer, smoothed=True, description=self.description)
def _spec(name, family, backend, fn, needs_layer=False, description=""): """Build one registry entry.""" return MethodSpec(name=name, family=family, backend=backend, fn=fn, needs_layer=needs_layer, description=description) ATTRIBUTION_METHODS: Dict[str, MethodSpec] = { "gradcam": _spec("gradcam", "cam", "torchcam", _torchcam_cam, True, "gradient-weighted feature maps; the default"), "gradcam_pp": _spec("gradcam_pp", "cam", "torchcam", _torchcam_cam, True, "Grad-CAM++: higher-order weights, better for several " "instances of the same object"), "scorecam": _spec("scorecam", "cam", "torchcam", _torchcam_cam, True, "Score-CAM: gradient-free, one forward pass per channel " "— slow but immune to gradient saturation"), "xgradcam": _spec("xgradcam", "cam", "torchcam", _torchcam_cam, True, "XGrad-CAM: axiom-derived weights, close to Grad-CAM on " "ReLU CNNs"), "layercam": _spec("layercam", "cam", "torchcam", _torchcam_cam, True, "Layer-CAM: per-element weights, usable at earlier " "layers where Grad-CAM degenerates"), "eigencam": _spec("eigencam", "cam", "spacr", _eigen_cam, True, "Eigen-CAM: first principal component of the " "activations; class-agnostic, useful as a control"), "saliency": _spec("saliency", "gradient", "captum", _captum_attribute, False, "absolute input gradient; the original saliency " "map"), "integrated_gradients": _spec( "integrated_gradients", "gradient", "captum", _captum_attribute, False, "integrated gradients: path integral from a baseline, so it depends on " "the baseline you pick"), "guided_backprop": _spec( "guided_backprop", "gradient", "captum", _captum_attribute, False, "guided backpropagation; sharp, and the usual failure case of the " "randomisation sanity check"), "input_x_gradient": _spec( "input_x_gradient", "gradient", "captum", _captum_attribute, False, "input times gradient: a first-order Taylor term"), "deeplift": _spec("deeplift", "gradient", "captum", _captum_attribute, False, "DeepLIFT rescale rule against a baseline"), "occlusion": _spec("occlusion", "perturbation", "captum", _captum_attribute, False, "slide a blanking window over the image and record the " "score drop"), "feature_ablation": _spec( "feature_ablation", "perturbation", "captum", _captum_attribute, False, "blank one tile at a time and record the score drop"), "attention_rollout": _spec( "attention_rollout", "attention", "spacr", _attention_adapter, False, "attention rollout for transformer backbones with no convolution to " "hook; not class-conditional"), "hirescam": _spec("hirescam", "cam", "spacr", _hires_cam, True, "HiRes-CAM: gradient times activation at every position, " "the Grad-CAM variant that stays faithful to the score"), "ablation_cam": _spec("ablation_cam", "cam", "spacr", _ablation_cam, True, "Ablation-CAM: gradient-free, zeroes one channel at a " "time and weights it by the score it costs"), "gradient_shap": _spec( "gradient_shap", "shap", "captum", _captum_attribute, False, "GradientSHAP: expected gradients between the image and a set of " "reference images; SHAP values in pixel units"), "deeplift_shap": _spec( "deeplift_shap", "shap", "captum", _captum_attribute, False, "DeepSHAP: DeepLIFT averaged over a set of reference images"), "chefer": _spec( "chefer", "attention", "spacr", _chefer_adapter, False, "Chefer transformer relevance: gradient-weighted attention propagated " "through the blocks; class-specific, Vision Transformers only"), }
[docs] def list_methods(family: Optional[str] = None) -> List[str]: """Registered method names, optionally restricted to one family.""" return sorted(n for n, s in ATTRIBUTION_METHODS.items() if family is None or s.family == family)
[docs] def methods_by_family() -> Dict[str, List[str]]: """Method names grouped by family, families in a stable order.""" out: Dict[str, List[str]] = {} for name in sorted(ATTRIBUTION_METHODS): out.setdefault(ATTRIBUTION_METHODS[name].family, []).append(name) return out
[docs] def attribute(model: nn.Module, image: Any, method: str = "gradcam", *, target: Optional[int] = None, layer: Optional[str] = None, model_type: Optional[str] = None, **kw) -> Attribution: """Attribute one image with one method. Every registered method returns the same thing: a finite ``(H, W)`` map at the input's resolution, for the class ``target`` names, for either head shape. What differs is what the number means, which is why :class:`Attribution` carries the family and the notes. :param model: the trained classifier. Left untouched — the wrapper and the hooks are removed before this returns. :param image: one image, ``(H, W)``, ``(C, H, W)`` or ``(1, C, H, W)``. :param method: a key of :data:`ATTRIBUTION_METHODS`. :param target: class index to explain; defaults to the model's prediction. For a single-logit head, 0 and 1 are both valid and give opposite maps. :param layer: dotted target-layer name for the CAM family; defaults to the last convolution. :param model_type: architecture name, used only to make errors readable. :param kw: method-specific options — ``n_steps``/``baseline`` (integrated gradients, DeepLIFT), ``window``/``stride`` (occlusion), ``block`` (feature ablation), ``head_fusion``/``discard_ratio`` (rollout), ``n_samples``/``sigma`` when called through :func:`smoothgrad`. :returns: the :class:`Attribution`. :raises UnknownMethodError: for an unregistered method name. :raises NoSpatialLayerError: when a CAM is asked of a model with no convolutional feature map, or rollout of a model with no attention. :raises AttributionError: for a bad target, layer or baseline. """ spec = ATTRIBUTION_METHODS.get(str(method)) if spec is None: grouped = ", ".join(f"{fam}: {names}" for fam, names in methods_by_family().items()) raise UnknownMethodError( f"unknown attribution method {method!r}. Registered methods by " f"family — {grouped}.") return _attribute_with_spec(spec, model, image, target=target, layer=layer, model_type=model_type, **kw)
def _attribute_with_spec(spec: MethodSpec, model: nn.Module, image: Any, *, target: Optional[int] = None, layer: Optional[str] = None, model_type: Optional[str] = None, **kw) -> Attribution: """Shared body of :func:`attribute` and :func:`smoothgrad`.""" x = _to_batch(image) was_training = model.training model.eval() wrapped = ClassScoreModel(model) try: target = _resolve_target(wrapped, x, target) predicted = _predicted_class(wrapped, x) out = spec.fn(spec, wrapped, x, int(target), layer, model_type, **kw) finally: if was_training: model.train() amap, raw, used_layer, notes = out amap = np.asarray(amap, dtype=np.float32) if not np.isfinite(amap).all(): amap = np.nan_to_num(amap, nan=0.0, posinf=0.0, neginf=0.0) result = Attribution( method=spec.name, map=amap, target=int(target), n_classes=wrapped.n_classes, single_logit=bool(wrapped.single_logit), predicted=int(predicted), raw=raw, layer=used_layer, family=spec.family, backend=spec.backend, params={"layer": used_layer, "smoothgrad": spec.smoothed, **dict(kw)}, notes=list(notes)) if wrapped.single_logit: result.notes.append( "This model has a single-logit binary head; it was attributed " f"through the [-z, +z] two-class view, so target={target} means " f"class {target} and not 'the logit'.") if result.is_flat(): result.notes.append( "The map is completely flat, so it ranks no pixel above another. " "Nothing downstream — deletion, insertion, pointing game — can say " "anything about it. For a CAM this usually means the target layer " "collapsed to 1x1 or its ReLU suppressed everything.") return result
[docs] def smoothgrad(model: nn.Module, image: Any, base_method: str = "saliency", *, n_samples: int = 25, sigma: float = 0.15, target: Optional[int] = None, layer: Optional[str] = None, model_type: Optional[str] = None, seed: Optional[int] = None, **kw) -> Attribution: """SmoothGrad: average ``base_method`` over ``n_samples`` noisy copies. Gradient maps are visually noisy because the gradient of a ReLU network fluctuates sharply between neighbouring inputs. Averaging over Gaussian perturbations of the input suppresses that fluctuation as ``1/sqrt(n)`` while keeping the structure that survives perturbation. For the captum-backed methods this is captum's own ``NoiseTunnel``. The CAM family and rollout are not captum attributors, so for those the averaging is done here over re-runs of the adapter — same definition, applied to the map. Smoothing makes a map *look* better. It does not make it more faithful, and a smoothed map that still fails :func:`randomization_sanity_check` fails it just as badly. :param model: the trained classifier. :param image: one image. :param base_method: any key of :data:`ATTRIBUTION_METHODS`. :param n_samples: number of noisy copies to average. ``1`` is a single noisy sample, *not* the clean map — call :func:`attribute` for that. :param sigma: noise standard deviation as a fraction of the input's dynamic range (the SmoothGrad paper's ``stdev_spread``). :param target: class to explain; resolved once on the clean image so the noise cannot flip which class is being explained sample to sample. :param layer: CAM target layer. :param model_type: architecture name for the error messages. :param seed: torch seed, so a repeated call reproduces. :param kw: forwarded to the base method. :returns: the averaged :class:`Attribution`. :raises AttributionError: for ``n_samples`` below 1. """ spec = ATTRIBUTION_METHODS.get(str(base_method)) if spec is None: raise UnknownMethodError( f"unknown base method {base_method!r} for SmoothGrad; registered " f"methods: {list_methods()}.") n_samples = int(n_samples) if n_samples < 1: raise AttributionError( f"n_samples must be at least 1, got {n_samples}; SmoothGrad over " f"zero samples has nothing to average.") if seed is not None: torch.manual_seed(int(seed)) x = _to_batch(image) wrapped = ClassScoreModel(model) target = _resolve_target(wrapped, x, target) if spec.backend == "captum": result = _attribute_with_spec( spec.smoothgrad_variant(), model, x, target=target, layer=layer, model_type=model_type, n_samples=n_samples, sigma=sigma, **kw) result.notes.append( f"SmoothGrad via captum NoiseTunnel: {n_samples} samples at " f"sigma={sigma} of the input range.") return result stdev = _sigma_to_stdev(sigma, x) maps: List[np.ndarray] = [] template: Optional[Attribution] = None for _ in range(n_samples): noisy = x + torch.randn_like(x) * stdev one = _attribute_with_spec(spec, model, noisy, target=target, layer=layer, model_type=model_type, **kw) maps.append(one.map) template = one averaged = np.mean(np.stack(maps, axis=0), axis=0).astype(np.float32) assert template is not None template.map = averaged template.raw = None template.params.update({"smoothgrad": True, "n_samples": n_samples, "sigma": sigma}) template.notes.append( f"SmoothGrad by averaging {n_samples} runs of {spec.name} at " f"sigma={sigma} of the input range (this family is not a captum " f"attributor, so NoiseTunnel does not apply).") return template
[docs] def compare_methods(model: nn.Module, image: Any, methods: Sequence[str] = (), *, target: Optional[int] = None, layer: Optional[str] = None, model_type: Optional[str] = None, skip_failures: bool = True, **kw) -> List[Attribution]: """Attribute one image with several methods, for side-by-side reading. The deliverable is the panel plus :func:`method_agreement` over it, not any single map. Methods that cannot run on this architecture (a CAM on a pure transformer) are skipped with their reason recorded rather than aborting the comparison, unless ``skip_failures`` is off. :param model: the trained classifier. :param image: one image. :param methods: method names; defaults to one representative of each family that works on any architecture, plus Grad-CAM. :param target: class to explain; resolved once so every method explains the same class. :param layer: CAM target layer. :param model_type: architecture name for the error messages. :param skip_failures: record and skip a failing method instead of raising. :param kw: forwarded to every method. :returns: the attributions, in the order requested. Failures appear as :class:`Attribution` objects with a flat map and the error in ``notes`` only when ``skip_failures`` is True. """ names = list(methods) or ["gradcam", "saliency", "integrated_gradients", "occlusion"] x = _to_batch(image) wrapped = ClassScoreModel(model) target = _resolve_target(wrapped, x, target) out: List[Attribution] = [] for name in names: try: out.append(attribute(model, x, name, target=target, layer=layer, model_type=model_type, **kw)) except Exception as exc: if not skip_failures: raise spec = ATTRIBUTION_METHODS.get(name) out.append(Attribution( method=name, map=np.zeros((int(x.shape[-2]), int(x.shape[-1])), dtype=np.float32), target=int(target), n_classes=wrapped.n_classes, single_logit=bool(wrapped.single_logit), predicted=_predicted_class(wrapped, x), family=spec.family if spec else "", layer=layer, backend=spec.backend if spec else "", notes=[f"FAILED: {type(exc).__name__}: {exc}", "This map is all zeros because the method did not run; " "it is a placeholder, not a result."])) return out
@dataclass
[docs] class Curve: """A deletion or insertion curve and its area. :ivar kind: ``'deletion'`` or ``'insertion'``. :ivar fractions: fraction of pixels removed / inserted at each step. :ivar scores: the target class's probability at each step. :ivar auc: area under the curve, trapezoidal over ``fractions``. Bounded in ``[0, 1]`` because the scores are probabilities. :ivar baseline: what removed pixels were replaced with. :ivar target: the class whose probability was tracked. :ivar notes: caveats. """ kind: str fractions: np.ndarray scores: np.ndarray auc: float baseline: str target: int notes: List[str] = field(default_factory=list) @property
[docs] def higher_is_better(self) -> bool: """Whether a larger AUC is the better outcome for this curve's kind.""" return self.kind == "insertion"
@property
[docs] def drop(self) -> float: """How far the score fell (deletion) or rose (insertion), start to end.""" return float(self.scores[0] - self.scores[-1])
def _ranked_pixels(amap: np.ndarray) -> np.ndarray: """Flat pixel indices ordered by the map, highest first, ties by position.""" flat = np.asarray(amap, dtype=np.float64).ravel() return np.argsort(-flat, kind="stable") def _perturbation_curve(model: nn.Module, image: Any, amap: Any, kind: str, *, target: Optional[int] = None, n_steps: int = 20, baseline: Any = "blur") -> Curve: """Shared body of :func:`deletion_curve` and :func:`insertion_curve`.""" if kind not in ("deletion", "insertion"): raise AttributionError( f"kind must be 'deletion' or 'insertion', got {kind!r}.") n_steps = int(n_steps) if n_steps < 1: raise AttributionError( f"n_steps must be at least 1, got {n_steps}; a curve needs at " f"least one perturbed point besides the unperturbed one.") x = _to_batch(image) amap = np.asarray( amap.map if isinstance(amap, Attribution) else amap, dtype=np.float64) h, w = int(x.shape[-2]), int(x.shape[-1]) if amap.shape != (h, w): raise AttributionError( f"the attribution map is {amap.shape} but the image is {(h, w)}; " f"the curve removes pixels by rank, so they must line up.") wrapped = ClassScoreModel(model) was_training = model.training model.eval() try: target = _resolve_target(wrapped, x, target) ref = _captum_baseline(baseline, x) order = _ranked_pixels(amap) n_pixels = order.size counts = [int(round(n_pixels * i / n_steps)) for i in range(n_steps + 1)] fractions: List[float] = [] scores: List[float] = [] for k in counts: mask = torch.ones(n_pixels, dtype=x.dtype) if k: mask[torch.as_tensor(order[:k].copy(), dtype=torch.long)] = 0.0 mask = mask.reshape(1, 1, h, w) if kind == "deletion": probe = x * mask + ref * (1.0 - mask) else: probe = ref * mask + x * (1.0 - mask) probs = class_scores(wrapped, probe, probability=True) fractions.append(k / float(n_pixels)) scores.append(float(probs[0, int(target)])) finally: if was_training: model.train() frac = np.asarray(fractions, dtype=np.float64) sc = np.asarray(scores, dtype=np.float64) auc = float(_trapezoid(sc, frac)) if frac.size > 1 else float(sc[0]) label = (str(baseline) if isinstance(baseline, (str, int, float)) else "custom tensor") return Curve(kind=kind, fractions=frac, scores=sc, auc=auc, baseline=label, target=int(target), notes=[CRITERION_CAVEATS[f"{kind}_auc"]])
[docs] def deletion_curve(model: nn.Module, image: Any, amap: Any, *, target: Optional[int] = None, n_steps: int = 20, baseline: Any = "blur") -> Curve: """Remove the highest-ranked pixels first and track the class probability. A faithful map removes the pixels the model actually uses, so the probability collapses early and the area under the curve is small. **A flat deletion curve is the finding**: the map ranked pixels the model does not use, whatever the picture looked like. :param model: the trained classifier. :param image: the image the map explains. :param amap: an :class:`Attribution` or a raw ``(H, W)`` map. :param target: class whose probability is tracked; defaults to the prediction on the unperturbed image. :param n_steps: perturbation steps between 0 % and 100 % removed. :param baseline: what removed pixels become — ``'blur'`` (default, the least out-of-distribution), ``'zero'``, ``'mean'``, ``'uniform'``, a number or a tensor. :returns: the :class:`Curve`; lower ``auc`` is better. """ return _perturbation_curve(model, image, amap, "deletion", target=target, n_steps=n_steps, baseline=baseline)
[docs] def insertion_curve(model: nn.Module, image: Any, amap: Any, *, target: Optional[int] = None, n_steps: int = 20, baseline: Any = "blur") -> Curve: """Start from a blanked image and add the highest-ranked pixels first. The mirror of :func:`deletion_curve`, and it answers a different question: deletion asks whether the map found pixels that are *necessary*, insertion whether it found pixels that are *sufficient*. The two routinely rank methods differently, and that disagreement is information, not an error. :param model: the trained classifier. :param image: the image the map explains. :param amap: an :class:`Attribution` or a raw ``(H, W)`` map. :param target: class whose probability is tracked. :param n_steps: insertion steps between 0 % and 100 % inserted. :param baseline: what the not-yet-inserted pixels are. :returns: the :class:`Curve`; higher ``auc`` is better. """ return _perturbation_curve(model, image, amap, "insertion", target=target, n_steps=n_steps, baseline=baseline)
[docs] def faithfulness(model: nn.Module, image: Any, amap: Any, *, target: Optional[int] = None, n_steps: int = 20, baseline: Any = "blur", mask: Optional[Any] = None) -> Dict[str, Any]: """Every faithfulness number for one map, with the caveats attached. :param model: the trained classifier. :param image: the image the map explains. :param amap: an :class:`Attribution` or a raw ``(H, W)`` map. :param target: class to score. :param n_steps: steps for both curves. :param baseline: removal baseline for both curves. :param mask: optional boolean object mask enabling the pointing game. :returns: dict with ``deletion_auc``, ``insertion_auc``, ``deletion`` and ``insertion`` :class:`Curve` objects, ``pointing_game`` (or None), ``flat`` and ``notes``. """ dele = deletion_curve(model, image, amap, target=target, n_steps=n_steps, baseline=baseline) ins = insertion_curve(model, image, amap, target=target, n_steps=n_steps, baseline=baseline) raw_map = np.asarray(amap.map if isinstance(amap, Attribution) else amap) out: Dict[str, Any] = { "deletion": dele, "insertion": ins, "deletion_auc": dele.auc, "insertion_auc": ins.auc, "pointing_game": None if mask is None else pointing_game(raw_map, mask), "flat": bool(float(raw_map.max() - raw_map.min()) <= 0), "notes": [NOT_AN_EXPLANATION, CRITERION_CAVEATS["deletion_auc"], CRITERION_CAVEATS["insertion_auc"]], } if out["flat"]: out["notes"].append( "The map is flat, so both curves describe removing pixels in " "arbitrary order. Neither AUC means anything here.") if mask is not None: out["notes"].append(CRITERION_CAVEATS["pointing_game"]) return out
[docs] def pointing_game(amap: Any, mask: Any, *, tolerance: int = 0) -> float: """Does the map's brightest pixel land inside the object? spaCR already has the answer key: ``merged/*.npy`` stores the label-mask planes next to the image channels, so a boolean object mask is free. The game is deliberately crude — one pixel, hit or miss — because that is all it claims to measure. :param amap: an :class:`Attribution` or a ``(H, W)`` map. :param mask: object mask, same spatial shape. Any non-zero value is inside the object, so a spaCR integer label plane can be passed directly. :param tolerance: dilate the mask by this many pixels before testing, the allowance the original pointing-game protocol uses for maps computed at a coarser resolution than the image. :returns: 1.0 for a hit, 0.0 for a miss. :raises AttributionError: on a shape mismatch or an empty mask — an empty mask would score 0.0 and look like a method failure rather than a missing annotation. """ m = np.asarray(amap.map if isinstance(amap, Attribution) else amap, dtype=np.float64) obj = np.asarray(mask) if obj.ndim > 2: obj = obj.reshape(-1, *obj.shape[-2:]).any(axis=0) obj = obj != 0 if obj.shape != m.shape: raise AttributionError( f"the object mask is {obj.shape} but the attribution map is " f"{m.shape}; the pointing game compares one to the other pixel for " f"pixel.") if not obj.any(): raise AttributionError( "the object mask is empty, so there is nothing for the map to " "point at. A score of 0 here would say the method failed when in " "fact the annotation is missing.") if int(tolerance) > 0: pad = int(tolerance) grown = np.zeros_like(obj) for dy in range(-pad, pad + 1): for dx in range(-pad, pad + 1): grown |= np.roll(np.roll(obj, dy, axis=0), dx, axis=1) obj = grown row, col = divmod(int(np.argmax(m)), m.shape[1]) return 1.0 if bool(obj[row, col]) else 0.0
[docs] def pointing_game_rate(maps: Sequence[Any], masks: Sequence[Any], *, tolerance: int = 0) -> Dict[str, Any]: """Pointing-game hit rate over a set of images. :param maps: attributions or raw maps. :param masks: the matching object masks. :param tolerance: passed to :func:`pointing_game`. :returns: dict with ``rate``, ``hits``, ``n``, ``skipped`` (images whose mask was empty or mismatched, which are excluded rather than counted as misses) and ``notes``. :raises AttributionError: when the two sequences differ in length. """ if len(maps) != len(masks): raise AttributionError( f"got {len(maps)} maps and {len(masks)} masks; the pointing game " f"needs one mask per map.") hits = 0 scored = 0 skipped: List[str] = [] for i, (m, k) in enumerate(zip(maps, masks)): try: hits += int(pointing_game(m, k, tolerance=tolerance)) scored += 1 except AttributionError as exc: skipped.append(f"image {i}: {exc}") notes = [CRITERION_CAVEATS["pointing_game"]] if skipped: notes.append( f"{len(skipped)} of {len(maps)} images were excluded (not counted " f"as misses) because their mask could not be used: {skipped[:3]}") if scored == 0: notes.append( "No image could be scored, so the rate below is not a measurement.") return {"rate": (hits / scored) if scored else float("nan"), "hits": hits, "n": scored, "skipped": skipped, "notes": notes}
@dataclass
[docs] class SanityCheck: """Result of randomising the model's weights and attributing again. :ivar method: the method under test. :ivar mode: ``'cascading'`` (randomise layers from the output backwards, accumulating) or ``'independent'`` (one layer at a time, from a fresh copy). :ivar stages: ``(layer_name, similarity)`` after each randomisation step, in the order applied. :ivar final_similarity: similarity to the trained model's map once every parameterised layer has been randomised. This is the number that matters: at that point the model is noise. :ivar max_similarity: the largest similarity over all stages. :ivar threshold: the value ``final_similarity`` must fall below to pass. :ivar passed: whether the method's map changed when the weights did. :ivar metric: name of the similarity measure. :ivar notes: the verdict in words. """ method: str mode: str stages: List[Tuple[str, float]] final_similarity: float max_similarity: float threshold: float passed: bool metric: str = "spearman_abs" notes: List[str] = field(default_factory=list) @property
[docs] def gap(self) -> float: """``1 - final_similarity``, clipped to ``[0, 1]``: higher is better. This is the form the hyperparameter search ranks on, so that every criterion there points the same way. """ if not math.isfinite(self.final_similarity): return 0.0 return float(min(1.0, max(0.0, 1.0 - self.final_similarity)))
[docs] def verdict(self) -> str: """One sentence a user can act on.""" if self.passed: return (f"{self.method} PASSES the randomisation sanity check: " f"after randomising every layer its map correlates " f"{self.final_similarity:.2f} with the trained model's " f"(threshold {self.threshold:.2f}), so it depends on the " f"weights.") return (f"{self.method} FAILS the randomisation sanity check: with " f"every weight randomised its map still correlates " f"{self.final_similarity:.2f} with the trained model's " f"(threshold {self.threshold:.2f}). It is describing the image, " f"not the model — do not read it as an explanation of this " f"classifier.")
def _rank_correlation(a: np.ndarray, b: np.ndarray) -> float: """Spearman rank correlation between two maps' absolute values. Rank-based on purpose: attribution scales are arbitrary and differ between methods, so only the ordering of pixels is comparable. Returns NaN when either map is constant, because a constant map has no ordering to correlate. """ x = np.abs(np.asarray(a, dtype=np.float64)).ravel() y = np.abs(np.asarray(b, dtype=np.float64)).ravel() if x.size != y.size: raise AttributionError( f"cannot correlate maps of different sizes ({x.size} vs {y.size}).") if x.size < 2 or float(x.max() - x.min()) <= 0 or float(y.max() - y.min()) <= 0: return float("nan") try: from scipy.stats import spearmanr rho = float(spearmanr(x, y).statistic) except Exception: rx = np.argsort(np.argsort(x)).astype(np.float64) ry = np.argsort(np.argsort(y)).astype(np.float64) rho = float(np.corrcoef(rx, ry)[0, 1]) return rho if math.isfinite(rho) else float("nan") def _parameterised_modules(model: nn.Module) -> List[str]: """Names of modules owning parameters directly, in definition order.""" return [name for name, mod in model.named_modules() if name and any(True for _ in mod.parameters(recurse=False))] def _randomize_module(module: nn.Module, generator: torch.Generator) -> None: """Replace a module's own parameters with noise of the same scale. Matching the original scale matters: parameters drawn far outside the trained range saturate the network, every logit collapses to the same value, and the resulting map is flat for *every* method — which would make every method appear to pass. """ with torch.no_grad(): for param in module.parameters(recurse=False): std = float(param.detach().float().std()) if not math.isfinite(std) or std <= 0: std = 0.05 noise = torch.randn(param.shape, generator=generator, dtype=torch.float32).to(param.dtype) param.copy_(noise * std)
[docs] def randomization_sanity_check( model: nn.Module, image: Any, method: Union[str, Callable] = "gradcam", *, target: Optional[int] = None, layer: Optional[str] = None, model_type: Optional[str] = None, mode: str = "cascading", threshold: float = 0.5, seed: int = 0, max_stages: Optional[int] = None, attribute_fn: Optional[Callable[..., Any]] = None, **kw) -> SanityCheck: """Adebayo et al. 2018: does the map change when the weights are destroyed? Randomise the model's parameters layer by layer, from the output backwards, re-attribute at every stage, and correlate each map with the map from the trained model. A method that depends on what the model learned produces an uncorrelated map once the weights are noise. Several widely used methods do not: guided backprop and guided Grad-CAM are the canonical failures, and a method that fails this is an edge detector being read as an explanation. **This is the most informative check in the module.** A map that passes deletion and insertion but fails this one is describing the image; a method that fails here cannot be rescued by smoothing, a better colormap or a different target layer. :param model: the trained classifier. Deep-copied — the original is never modified. :param image: one image. :param method: a registered method name, or any callable when ``attribute_fn`` is not given. :param target: class to explain, resolved once against the trained model and held fixed, so a stage's map is not silently for a different class. :param layer: CAM target layer. :param model_type: architecture name for the error messages. :param mode: ``'cascading'`` (default, the paper's) or ``'independent'``. :param threshold: rank correlation below which the method passes. :param seed: RNG seed for the randomisation, so the check reproduces. :param max_stages: cap on the number of layers randomised, from the output backwards. The final stage always randomises everything regardless. :param attribute_fn: ``fn(model, image, target=...) -> map or Attribution``, for testing a method that is not in the registry. :param kw: forwarded to the attribution call. :returns: the :class:`SanityCheck`. :raises AttributionError: for an unknown mode, or a model with no parameters to randomise. """ if mode not in ("cascading", "independent"): raise AttributionError( f"mode must be 'cascading' or 'independent', got {mode!r}.") x = _to_batch(image) wrapped = ClassScoreModel(model) target = _resolve_target(wrapped, x, target) method_name = method if isinstance(method, str) else getattr( method, "__name__", "custom") def _run(m: nn.Module) -> np.ndarray: """Attribute ``image`` with the method under test on model ``m``.""" if attribute_fn is not None: out = attribute_fn(m, x, target=int(target)) elif callable(method) and not isinstance(method, str): out = method(m, x, target=int(target)) else: out = attribute(m, x, str(method), target=int(target), layer=layer, model_type=model_type, **kw) return np.asarray(out.map if isinstance(out, Attribution) else out, dtype=np.float64) reference = _run(model) names = _parameterised_modules(model) if not names: raise AttributionError( "the model has no parameterised layers, so there is nothing to " "randomise and the sanity check cannot be run.") order = list(reversed(names)) if max_stages is not None and int(max_stages) > 0: order = order[:int(max_stages)] generator = torch.Generator().manual_seed(int(seed)) stages: List[Tuple[str, float]] = [] cascade = copy.deepcopy(model) cascade.eval() for name in order: if mode == "cascading": probe = cascade else: probe = copy.deepcopy(model) probe.eval() _randomize_module(dict(probe.named_modules())[name], generator) stages.append((name, _rank_correlation(reference, _run(probe)))) full_generator = torch.Generator().manual_seed(int(seed) + 991) full = copy.deepcopy(model) full.eval() for name in reversed(names): _randomize_module(dict(full.named_modules())[name], full_generator) final = _rank_correlation(reference, _run(full)) if not stages or stages[-1][0] != names[0] or mode != "cascading": stages.append(("<all layers>", final)) finite = [s for _n, s in stages if math.isfinite(s)] max_sim = max(finite) if finite else float("nan") passed = bool(math.isfinite(final) and final < float(threshold)) notes = [CRITERION_CAVEATS["sanity_gap"]] if not math.isfinite(final): notes.append( "The randomised model produced a constant map, so no rank " "correlation exists and the check is inconclusive rather than " "passed. A flat map is not evidence of sensitivity to the weights.") passed = False result = SanityCheck(method=method_name, mode=mode, stages=stages, final_similarity=final, max_similarity=max_sim, threshold=float(threshold), passed=passed, notes=notes) result.notes.insert(0, result.verdict()) return result
@dataclass
[docs] class Agreement: """Pairwise rank correlation between several methods on the same image. :param methods: attribution-method names in matrix row and column order. :param matrix: symmetric Spearman rank-correlation matrix for those methods. :param mean: mean of the finite off-diagonal correlations. :param minimum: smallest finite off-diagonal correlation. :param pairs: ``(method_a, method_b, rho)`` comparisons ordered from greatest disagreement upward. :param notes: verdict and interpretive caveats accompanying the agreement result. """ methods: List[str] matrix: np.ndarray mean: float minimum: float pairs: List[Tuple[str, str, float]] = field(default_factory=list) notes: List[str] = field(default_factory=list)
[docs] def verdict(self) -> str: """One sentence, deliberately asymmetric — see the note in the source.""" if not math.isfinite(self.mean): return ("Agreement could not be computed: at least one map is " "constant and has no pixel ordering to correlate.") if self.minimum < 0.2: return (f"The methods DISAGREE (lowest pair rho={self.minimum:.2f}, " f"mean {self.mean:.2f}). They cannot all be describing what " f"this model uses, so no single map here should be trusted " f"on its own.") if self.mean > 0.7: return (f"The methods agree (mean rho={self.mean:.2f}). Agreement " f"is weak evidence: methods sharing a failure mode — the " f"gradient family shares several — agree with each other " f"while all being wrong. Check the randomisation sanity " f"check before reading agreement as confirmation.") return (f"The methods partly agree (mean rho={self.mean:.2f}, lowest " f"pair {self.minimum:.2f}). Read the panel, not the ranking.")
[docs] def method_agreement(attributions: Sequence[Any]) -> Agreement: """Rank correlation between several attribution maps of the same image. Agreement and disagreement are not symmetric evidence. Methods agreeing is weak: the gradient family shares failure modes, so its members agree with each other whether or not any of them is faithful. Methods disagreeing is strong: at most one of them can be right, so none should be quoted alone. :param attributions: :class:`Attribution` objects or raw ``(H, W)`` maps; at least two, all the same shape. :returns: the :class:`Agreement`. :raises AttributionError: for fewer than two maps or a shape mismatch. """ if len(attributions) < 2: raise AttributionError( f"agreement needs at least two maps to compare, got " f"{len(attributions)}.") names: List[str] = [] maps: List[np.ndarray] = [] for i, a in enumerate(attributions): if isinstance(a, Attribution): names.append(a.method) maps.append(np.asarray(a.map, dtype=np.float64)) else: names.append(f"map_{i}") maps.append(np.asarray(a, dtype=np.float64)) shapes = {m.shape for m in maps} if len(shapes) > 1: raise AttributionError( f"the maps have different shapes {sorted(shapes)}; agreement " f"compares them pixel for pixel.") n = len(maps) matrix = np.eye(n, dtype=np.float64) pairs: List[Tuple[str, str, float]] = [] for i in range(n): for j in range(i + 1, n): rho = _rank_correlation(maps[i], maps[j]) matrix[i, j] = matrix[j, i] = rho pairs.append((names[i], names[j], rho)) off = [p[2] for p in pairs if math.isfinite(p[2])] mean = float(np.mean(off)) if off else float("nan") minimum = float(np.min(off)) if off else float("nan") pairs.sort(key=lambda p: (math.inf if not math.isfinite(p[2]) else p[2])) result = Agreement(methods=names, matrix=matrix, mean=mean, minimum=minimum, pairs=pairs, notes=[NOT_AN_EXPLANATION]) result.notes.insert(0, result.verdict()) if len(off) < len(pairs): result.notes.append( f"{len(pairs) - len(off)} pairs could not be correlated because a " f"map is constant; they are excluded rather than scored as " f"disagreement.") return result
[docs] class AttributionMapGenerator: """Batch adapter with the interface ``generate_activation_map`` already uses. :class:`spacr.utils.GradCAMGenerator` and :class:`spacr.utils.SaliencyMapGenerator` expose ``compute_*_and_predictions(X)`` plus ``plot_activation_grid``. This offers the same two calls for every method in :data:`ATTRIBUTION_METHODS`, so the existing batch loop gains twelve methods without changing shape. :param model: the trained classifier. :param method: a key of :data:`ATTRIBUTION_METHODS`. :param target_layer: CAM target layer, or None for the last convolution. :param model_type: architecture name, used in the error messages. :param smoothgrad_samples: when above 1, each map is SmoothGrad-averaged. :param smoothgrad_sigma: SmoothGrad noise as a fraction of the input range. :param kw: forwarded to the method. :raises UnknownMethodError: if ``method`` is not registered in :data:`ATTRIBUTION_METHODS`. """ def __init__(self, model, method: str = "gradcam", target_layer: Optional[str] = None, model_type: Optional[str] = None, smoothgrad_samples: int = 0, smoothgrad_sigma: float = 0.15, **kw): """Validate the method name up front rather than mid-batch.""" if str(method) not in ATTRIBUTION_METHODS: raise UnknownMethodError( f"unknown attribution method {method!r}. Registered methods by " f"family — " + ", ".join( f"{fam}: {names}" for fam, names in methods_by_family().items())) self.model = model self.method = str(method) self.target_layer = target_layer self.model_type = model_type self.smoothgrad_samples = int(smoothgrad_samples or 0) self.smoothgrad_sigma = float(smoothgrad_sigma) self.kw = dict(kw) self.model.eval()
[docs] def compute_maps_and_predictions(self, X): """Attribute every image in a batch. :param X: batch tensor ``(N, C, H, W)``. :returns: ``(maps, predictions)`` — maps is ``(N, H, W)`` float32, predictions is ``(N,)`` long, correct for either head shape. """ maps: List[np.ndarray] = [] preds: List[int] = [] for i in range(int(X.shape[0])): one = X[i:i + 1] if self.smoothgrad_samples > 1: att = smoothgrad(self.model, one, self.method, n_samples=self.smoothgrad_samples, sigma=self.smoothgrad_sigma, layer=self.target_layer, model_type=self.model_type, **self.kw) else: att = attribute(self.model, one, self.method, layer=self.target_layer, model_type=self.model_type, **self.kw) maps.append(att.map) preds.append(att.predicted) return (torch.from_numpy(np.stack(maps, axis=0)), torch.tensor(preds, dtype=torch.long))
compute_gradcam_and_predictions = compute_maps_and_predictions compute_saliency_and_predictions = compute_maps_and_predictions
[docs] def plot_activation_grid(self, X, maps, predictions, overlay=True, normalize=False): """Render the batch grid, reusing spaCR's existing layout. :param X: the input batch. :param maps: the attribution maps. :param predictions: predicted class per image. :param overlay: draw the map over the image. :param normalize: percentile-stretch the image under the overlay. :returns: the matplotlib Figure. """ from .utils import SaliencyMapGenerator return SaliencyMapGenerator(self.model).plot_activation_grid( X, maps, predictions, overlay=overlay, normalize=normalize)
_CF_SIDE = 64 class _CounterfactualGenerator(nn.Module): """A small conditional autoencoder whose decoder is told which class to draw. The encoder maps a crop, resized to at most ``_CF_SIDE`` pixels a side, to a latent vector; the decoder draws it back from that vector and a class code (one weight per class, summing to one). Decoding the same latent with a different code is the counterfactual edit. """ def __init__(self, channels: int, side: int, n_classes: int, latent: int = 32, width: int = 16): """Build the encoder and decoder for ``channels``-channel crops.""" super().__init__() self.side = int(side) self.n_classes = int(n_classes) self.width = int(width) cells = 2 * self.width * (self.side // 4) ** 2 self.encoder = nn.Sequential( nn.Conv2d(channels, self.width, 3, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(self.width, 2 * self.width, 3, 2, 1), nn.LeakyReLU(0.2), nn.Flatten(), nn.Linear(cells, latent)) self.decoder_in = nn.Linear(latent + self.n_classes, cells) self.decoder = nn.Sequential( nn.LeakyReLU(0.2), nn.ConvTranspose2d(2 * self.width, self.width, 4, 2, 1), nn.LeakyReLU(0.2), nn.ConvTranspose2d(self.width, channels, 4, 2, 1)) def encode(self, x: torch.Tensor) -> torch.Tensor: """Latent vectors for crops already at the generator's side.""" return self.encoder(x) def decode(self, z: torch.Tensor, code: torch.Tensor) -> torch.Tensor: """Crops drawn from latents ``z`` under class codes ``code``.""" h = self.decoder_in(torch.cat([z, code], dim=1)) h = h.view(-1, 2 * self.width, self.side // 4, self.side // 4) return self.decoder(h) #: Generators ``counterfactual_generator`` can name. COUNTERFACTUAL_GENERATORS = ('autoencoder', 'diffusion') class _CfResBlock(nn.Module): """A residual block conditioned on a time-plus-class embedding.""" def __init__(self, cin: int, cout: int, emb: int): """Two normalised convolutions with the embedding added between them. :param cin: input channels. :param cout: output channels. :param emb: length of the time-plus-class embedding. """ super().__init__() self.norm1 = nn.GroupNorm(min(8, cin), cin) self.conv1 = nn.Conv2d(cin, cout, 3, padding=1) self.emb = nn.Linear(emb, cout) self.norm2 = nn.GroupNorm(min(8, cout), cout) self.conv2 = nn.Conv2d(cout, cout, 3, padding=1) self.skip = nn.Conv2d(cin, cout, 1) if cin != cout else nn.Identity() def forward(self, x: torch.Tensor, e: torch.Tensor) -> torch.Tensor: """The block's output for features ``x`` under embedding ``e``.""" h = self.conv1(F.silu(self.norm1(x))) h = h + self.emb(e)[:, :, None, None] h = self.conv2(F.silu(self.norm2(h))) return h + self.skip(x) class _CounterfactualDiffusion(nn.Module): """A class-conditional denoising diffusion model with the generator's API. A small U-Net predicts the noise in a crop at a diffusion step, told the step and a class code. The code enters through one learned vector per class, mixed by the code's weights, so a code between two classes is a blend of their embeddings and the class can be moved continuously; a learned null vector stands for "no class" and is what classifier-free guidance contrasts against. It offers the same ``encode`` / ``decode`` / ``side`` / ``n_classes`` as :class:`_CounterfactualGenerator`, so ``_cf_edit``, ``_counterfactual_sequences`` and the report use it unchanged: * ``encode`` runs deterministic DDIM inversion WITHOUT a class, from the crop to the noise level ``strength`` (a fraction of the schedule). What survives that far is the cell's layout; what the noise has erased is what a class can redraw. * ``decode`` runs guided DDIM sampling back to a crop under a class code. Unlike the autoencoder, nothing here is trained against the classifier, so a flip under the classifier is not something training optimised for: it has to come from the model having learned how the classes look. Crops are standardised per channel with the training mean and standard deviation, held as buffers so they travel with the weights. """ def __init__(self, channels: int, side: int, n_classes: int, width: int = 64, timesteps: int = 1000, strength: float = 0.95, sample_steps: int = 50, guidance_scale: float = 4.0): """Build the U-Net and the noise schedule. :param side: crop side the model works at; a multiple of 4. :param width: channels of the first U-Net level (doubled below). :param timesteps: length of the training noise schedule. :param strength: noise level ``encode`` inverts to, 0..1. The defaults (0.95, guidance 4, 50 steps, width 64) are the setting chosen on held-out BBBC014 crops: below about 0.8 too little is erased for the class to redraw anything and the classifier barely moves; at 1.0 the cell's own layout is lost and the phenotype overshoots the real hits. :param sample_steps: DDIM steps over the whole schedule; inversion and sampling use the share of them below ``strength``. :param guidance_scale: classifier-free guidance weight at sampling. """ super().__init__() self.side = int(side) self.n_classes = int(n_classes) self.channels = int(channels) self.width = int(width) self.timesteps = int(timesteps) self.strength = float(strength) self.sample_steps = int(sample_steps) self.guidance_scale = float(guidance_scale) w, emb = self.width, 4 * self.width self.time_mlp = nn.Sequential(nn.Linear(w, emb), nn.SiLU(), nn.Linear(emb, emb)) self.class_emb = nn.Parameter(torch.randn(self.n_classes, emb) * 0.02) self.null_emb = nn.Parameter(torch.zeros(emb)) self.inp = nn.Conv2d(channels, w, 3, padding=1) self.d1 = _CfResBlock(w, w, emb) self.d2 = _CfResBlock(w, 2 * w, emb) self.d3 = _CfResBlock(2 * w, 2 * w, emb) self.mid = _CfResBlock(2 * w, 2 * w, emb) self.u3 = _CfResBlock(4 * w, 2 * w, emb) self.u2 = _CfResBlock(4 * w, w, emb) self.u1 = _CfResBlock(2 * w, w, emb) self.out = nn.Sequential(nn.GroupNorm(min(8, w), w), nn.SiLU(), nn.Conv2d(w, channels, 3, padding=1)) steps = torch.arange(self.timesteps + 1, dtype=torch.float64) f = torch.cos((steps / self.timesteps + 0.008) / 1.008 * torch.pi / 2) ** 2 abar = (f / f[0]).clamp(1e-5, 1.0)[1:].float() self.register_buffer('abar', abar) self.register_buffer('mean', torch.zeros(1, channels, 1, 1)) self.register_buffer('std', torch.ones(1, channels, 1, 1)) def _time(self, t: torch.Tensor) -> torch.Tensor: """The sinusoidal embedding of diffusion steps ``t``, through the MLP.""" half = self.width // 2 freqs = torch.exp(-torch.log(torch.tensor(10000.0)) * torch.arange(half, device=t.device) / max(1, half - 1)) ang = t.float()[:, None] * freqs[None] return self.time_mlp(torch.cat([ang.sin(), ang.cos()], dim=1)) def _cond(self, code: Optional[torch.Tensor], n: int) -> torch.Tensor: """The class embedding of ``code``, or the null one for ``n`` crops.""" if code is None: return self.null_emb.expand(n, -1) return code.float() @ self.class_emb def forward(self, x: torch.Tensor, t: torch.Tensor, code: Optional[torch.Tensor], drop: Optional[torch.Tensor] = None) -> torch.Tensor: """Predicted noise in ``x`` at steps ``t`` under ``code`` (or none). :param drop: optional per-crop booleans; a True crop is given the null class instead of its code, which is how training teaches the unconditional model classifier-free guidance contrasts against. """ cond = self._cond(code, x.shape[0]) if drop is not None: cond = torch.where(drop[:, None], self.null_emb.expand_as(cond), cond) e = self._time(t) + cond h1 = self.d1(self.inp(x), e) h2 = self.d2(F.avg_pool2d(h1, 2), e) h3 = self.d3(F.avg_pool2d(h2, 2), e) h = self.mid(h3, e) h = self.u3(torch.cat([h, h3], 1), e) h = self.u2(torch.cat([F.interpolate(h, scale_factor=2.0), h2], 1), e) h = self.u1(torch.cat([F.interpolate(h, scale_factor=2.0), h1], 1), e) return self.out(h) def _schedule(self) -> List[int]: """The DDIM steps from 0 up to the ``strength`` noise level. A strength of zero is no step at all, so encode and decode return the crop unchanged. """ if self.strength <= 0: return [0] top = max(1, int(round(self.strength * (self.timesteps - 1)))) n = max(1, int(round(self.sample_steps * self.strength))) return sorted({int(round(v)) for v in np.linspace(0, top, n + 1)}) def _guided(self, x, t, code): """Classifier-free guided noise at step ``t``; unguided without a code.""" tt = torch.full((x.shape[0],), int(t), device=x.device, dtype=torch.long) uncond = self(x, tt, None) if code is None: return uncond cond = self(x, tt, code) return uncond + self.guidance_scale * (cond - uncond) def _ddim_step(self, x, eps, t_from, t_to): """One deterministic DDIM move of ``x`` from step ``t_from`` to ``t_to``.""" a0, a1 = self.abar[t_from], self.abar[t_to] x0 = (x - (1 - a0).sqrt() * eps) / a0.sqrt() return a1.sqrt() * x0 + (1 - a1).sqrt() * eps def encode(self, x: torch.Tensor) -> torch.Tensor: """Invert crops, without a class, to the ``strength`` noise level.""" z = (x - self.mean) / self.std sched = self._schedule() with torch.no_grad(): for t_from, t_to in zip(sched[:-1], sched[1:]): z = self._ddim_step(z, self._guided(z, t_from, None), t_from, t_to) return z def decode(self, z: torch.Tensor, code: torch.Tensor) -> torch.Tensor: """Sample crops back from inverted ``z`` under class ``code``.""" sched = self._schedule() x = z with torch.no_grad(): for t_from, t_to in zip(sched[::-1][:-1], sched[::-1][1:]): x = self._ddim_step(x, self._guided(x, t_from, code), t_from, t_to) return x * self.std + self.mean def _train_counterfactual_diffusion(model: nn.Module, crops: torch.Tensor, *, epochs: int = 30, batch_size: int = 64, lr: float = 2e-4, width: int = 64, timesteps: int = 1000, strength: float = 0.95, sample_steps: int = 50, guidance_scale: float = 4.0, drop_class: float = 0.15, seed: int = 0, device: Any = 'cpu', labels: Optional[torch.Tensor] = None, target_probs: Optional[torch.Tensor] = None, augment: bool = True): """Train a class-conditional diffusion generator on crops. The classifier only LABELS the crops (its predicted class, or the given condition ``labels``); it does not guide training, so the flip rate this generator later scores is not something it was optimised for. Each step noises a batch to random diffusion steps and fits the noise; the class is dropped with probability ``drop_class`` so the same network also learns the unconditional model classifier-free guidance needs. With ``augment`` crops are randomly flipped and rotated by multiples of 90 degrees, which is a symmetry of a centred single-cell crop. :returns: ``(generator, history)`` as :func:`_train_counterfactual_generator` gives them. The last history entry also holds ``reconstruction``: the mean squared error of decoding a crop's inversion under its own class, on up to 256 training crops, which is how far the round trip alone moves a crop. """ device = torch.device(device) torch.manual_seed(int(seed)) rng = torch.Generator().manual_seed(int(seed)) wrapped = ClassScoreModel(model).to(device).eval() crops = crops.float() labels, target_probs = _cf_codes(wrapped, crops.to(device), labels, target_probs) n_codes = target_probs.shape[0] side = _cf_side(crops.shape[-1]) small_all = _cf_resize(crops, side) gen = _CounterfactualDiffusion(crops.shape[1], side, n_codes, width=width, timesteps=timesteps, strength=strength, sample_steps=sample_steps, guidance_scale=guidance_scale) gen.mean.copy_(small_all.mean(dim=(0, 2, 3), keepdim=True)) gen.std.copy_(small_all.std(dim=(0, 2, 3), keepdim=True).clamp_min(1e-6)) gen = gen.to(device) optimiser = torch.optim.AdamW(gen.parameters(), lr=lr) history = [] n = small_all.shape[0] for _epoch in range(int(epochs)): gen.train() order = torch.randperm(n, generator=rng) total = 0.0 for start in range(0, n, int(batch_size)): idx = order[start:start + int(batch_size)] x = ((small_all[idx].to(device) - gen.mean) / gen.std) if augment: k = int(torch.randint(0, 4, (1,), generator=rng)) x = torch.rot90(x, k, dims=(2, 3)) if bool(torch.randint(0, 2, (1,), generator=rng)): x = x.flip(3) code = F.one_hot(labels[idx], n_codes).float().to(device) drop = (torch.rand(len(idx), generator=rng) < drop_class).to(device) t = torch.randint(0, gen.timesteps, (len(idx),), generator=rng).to(device) noise = torch.randn(x.shape, generator=rng).to(device) a = gen.abar[t][:, None, None, None] noisy = a.sqrt() * x + (1 - a).sqrt() * noise loss = F.mse_loss(gen(noisy, t, code, drop=drop), noise) optimiser.zero_grad() loss.backward() optimiser.step() total += loss.item() * len(idx) / n history.append({'denoising': total}) gen.eval() probe = small_all[:256].to(device) with torch.no_grad(): back = gen.decode(gen.encode(probe), F.one_hot(labels[:probe.shape[0]], n_codes).float().to(device)) if history: history[-1]['reconstruction'] = float(F.mse_loss(back, probe)) return gen, history def _cf_side(size: int) -> int: """The generator's working side for crops ``size`` pixels across.""" return max(8, (min(int(size), _CF_SIDE) // 4) * 4) def _cf_resize(x: torch.Tensor, side: int) -> torch.Tensor: """Resize a crop batch to ``side`` pixels a side, bilinearly.""" if tuple(x.shape[-2:]) == (side, side): return x return F.interpolate(x, size=(side, side), mode='bilinear', align_corners=False) def _cf_scores(wrapped: ClassScoreModel, x: torch.Tensor, batch_size: int = 64) -> torch.Tensor: """Softmax class probabilities for ``x``, batched, without gradients.""" out = [] with torch.no_grad(): for start in range(0, x.shape[0], batch_size): out.append(torch.softmax(wrapped(x[start:start + batch_size]), dim=-1)) return torch.cat(out, dim=0) def _cf_edit(generator: _CounterfactualGenerator, x: torch.Tensor, z: torch.Tensor, base: torch.Tensor, code: torch.Tensor) -> torch.Tensor: """``x`` plus the change the generator draws when the class code moves. Adding the difference between the decoding under ``code`` and the decoding under the crop's own class, rather than using the decoding itself, keeps every detail the small generator cannot reproduce: at the crop's own class the edit is exactly zero. """ delta = generator.decode(z, code) - base return x + F.interpolate(delta, size=tuple(x.shape[-2:]), mode='bilinear', align_corners=False) def _cf_other_class(labels: torch.Tensor, n_classes: int, rng: torch.Generator) -> torch.Tensor: """A random class different from each label.""" shift = torch.randint(1, n_classes, labels.shape, generator=rng) return (labels + shift) % n_classes def _train_counterfactual_generator(model: nn.Module, crops: torch.Tensor, *, epochs: int = 30, batch_size: int = 32, lr: float = 2e-3, guidance: float = 1.0, proximity: float = 1.0, latent: int = 32, seed: int = 0, device: Any = 'cpu', labels: Optional[torch.Tensor] = None, target_probs: Optional[torch.Tensor] = None): """Train a class-conditional generator on crops, guided by the classifier. The classifier is frozen and, by default, labels each crop with its own prediction, so the generator learns what the classifier separates, not what an annotator meant. Three terms are minimised per batch: reconstruction of the crop under its own code; the classifier's cross-entropy against the target code's class probabilities on the edited crop (see ``_cf_edit``); and the mean absolute size of the edit, which keeps counterfactuals close to the original. With ``labels`` and ``target_probs`` the codes are conditions instead of classes (for example the plate or well a crop came from): each crop has its condition's code, and an edit toward another condition is pushed to the classifier's mean class probabilities for that condition, so a control crop is drawn the way the classifier sees cells of a treated well. Without them the targets are the classes themselves (one-hot). Because the same classifier guides training and later scores the counterfactuals, a high flip rate alone can reflect an adversarial edit; read it with the edit size and the class-mean baseline that ``_counterfactual_report`` reports next to it. :param model: the trained classifier; its weights are not changed. :param crops: ``(N, C, H, W)`` float tensor of input-ready crops. :param epochs: passes over the crops. :param batch_size: crops per optimisation step. :param lr: Adam learning rate. :param guidance: weight of the classifier term. :param proximity: weight of the edit-size term. :param latent: latent vector length. :param seed: seed for initialisation, batching and target classes. :param device: torch device for training. :param labels: per-crop condition index, or ``None`` for the predicted class. :param target_probs: ``(n_conditions, n_classes)`` class probabilities each condition's edits are pushed toward; ``None`` for one-hot classes. :returns: ``(generator, history)``, the generator in eval mode on ``device`` and a list with one dict of mean losses per epoch. """ device = torch.device(device) torch.manual_seed(int(seed)) rng = torch.Generator().manual_seed(int(seed)) wrapped = ClassScoreModel(model).to(device).eval() crops = crops.float().to(device) labels, target_probs = _cf_codes(wrapped, crops, labels, target_probs) target_probs = target_probs.to(device) n_classes = target_probs.shape[0] side = _cf_side(crops.shape[-1]) generator = _CounterfactualGenerator(crops.shape[1], side, n_classes, latent=latent).to(device) optimiser = torch.optim.Adam(generator.parameters(), lr=lr) frozen = [p.requires_grad for p in wrapped.parameters()] for p in wrapped.parameters(): p.requires_grad_(False) history = [] try: for _epoch in range(int(epochs)): order = torch.randperm(crops.shape[0], generator=rng) totals = {'reconstruction': 0.0, 'classifier': 0.0, 'edit': 0.0} for start in range(0, len(order), int(batch_size)): idx = order[start:start + int(batch_size)] xb = crops[idx.to(device)] src = labels[idx] tgt = _cf_other_class(src, n_classes, rng) small = _cf_resize(xb, side) z = generator.encode(small) base = generator.decode(z, F.one_hot(src, n_classes).float().to(device)) recon = F.mse_loss(base, small) flipped = generator.decode(z, F.one_hot(tgt, n_classes).float().to(device)) edit = (flipped - base).abs().mean() edited = xb + F.interpolate(flipped - base, size=tuple(xb.shape[-2:]), mode='bilinear', align_corners=False) cls = F.cross_entropy(wrapped(edited), target_probs[tgt.to(device)]) loss = recon + guidance * cls + proximity * edit optimiser.zero_grad() loss.backward() optimiser.step() weight = len(idx) / len(order) totals['reconstruction'] += float(recon) * weight totals['classifier'] += float(cls) * weight totals['edit'] += float(edit) * weight history.append(totals) finally: for p, flag in zip(wrapped.parameters(), frozen): p.requires_grad_(flag) generator.eval() return generator, history def _cf_codes(wrapped: ClassScoreModel, crops: torch.Tensor, labels: Optional[torch.Tensor] = None, target_probs: Optional[torch.Tensor] = None): """Per-crop codes and each code's target class probabilities. Without ``labels`` a crop's code is its predicted class and the targets are one-hot; otherwise both are returned as given. """ if labels is None: labels = _cf_scores(wrapped, crops).argmax(dim=1).cpu() target_probs = torch.eye(wrapped.n_classes) return (torch.as_tensor(labels).long().cpu(), torch.as_tensor(target_probs).float()) def _cf_condition_of(name: str, key: str) -> str: """The condition a crop file named ``<plate>_<well>_...`` belongs to. ``key`` is ``plate``, ``well``, ``row`` or ``column``; rows and columns are read from the well (``B03`` is row ``r2``, column ``c3``). An empty string when the name has no such part. """ from .schema import parse_well parts = str(name).split('_') if key == 'plate': return parts[0] well = parts[1] if len(parts) > 1 else '' if key == 'well': return f'{parts[0]}_{well}' if well else '' try: row, column = parse_well(well) except Exception: return '' return row if key == 'row' else column def _spearman(a: np.ndarray, b: np.ndarray) -> float: """Spearman rank correlation of two vectors, NaN when one is constant.""" a, b = np.asarray(a, dtype=float), np.asarray(b, dtype=float) if a.size < 2 or np.ptp(a) == 0 or np.ptp(b) == 0: return float('nan') ra = np.argsort(np.argsort(a)).astype(float) rb = np.argsort(np.argsort(b)).astype(float) return float(np.corrcoef(ra, rb)[0, 1]) def _counterfactual_sequences(model: nn.Module, generator: _CounterfactualGenerator, crops: torch.Tensor, *, targets=None, steps: int = 7, keep: int = 0, monotone_tolerance: float = 0.02, device: Any = 'cpu', labels=None, target_probs=None, finals: Optional[list] = None): """Morph each crop toward another class and score every step. The code moves in ``steps`` equal steps from the crop's own code to its target, and every step is scored as the classifier's class probabilities projected on the target's, ``p . t / (t . t)``; for a class target that is the classifier's probability of the target class. With ``labels`` and ``target_probs`` the codes are conditions, as in ``_train_counterfactual_generator``. :param model: the classifier the generator was trained against. :param generator: from ``_train_counterfactual_generator``. :param crops: ``(N, C, H, W)`` crops. :param targets: per-crop target classes; ``None`` picks the next class after the predicted one (the other class for a binary model). :param steps: frames per sequence, including the unchanged crop. :param keep: how many sequences to return as images, from the first crop. :param monotone_tolerance: largest drop in the target score between consecutive frames that still counts as monotone. :param device: torch device. :param labels: per-crop condition index, or ``None``. :param target_probs: per-condition class probabilities, or ``None``. :param finals: when a list is given, every crop's last frame is appended to it (CPU tensors), which is what the realism score compares. :returns: ``(rows, frames)``: a list of one dict per crop (source and target class, target score at every step, Spearman correlation of score with step, whether the score is monotone, whether the final frame is predicted as the target, and the edit size as the mean absolute change over the crop's intensity range and the fraction of pixels changed by more than a tenth of it) and a ``(keep, steps, C, H, W)`` array. All steps are decoded in one batch. The same latent decoded under every interpolated code at once gives exactly the frames one decode per step would, in a single pass, which matters for the diffusion generator, where one decode is a whole guided sampling run. """ device = torch.device(device) wrapped = ClassScoreModel(model).to(device).eval() generator = generator.to(device).eval() crops = crops.float().to(device) n = generator.n_classes src, target_probs = _cf_codes(wrapped, crops, labels, target_probs) target_probs = target_probs.to(device) tgt = (src + 1) % n if targets is None else torch.as_tensor(targets).long() alphas = torch.linspace(0.0, 1.0, int(steps)) rows, frames = [], [] with torch.no_grad(): for i in range(crops.shape[0]): x = crops[i:i + 1] small = _cf_resize(x, generator.side) z = generator.encode(small) one_src = F.one_hot(src[i:i + 1], n).float().to(device) one_tgt = F.one_hot(tgt[i:i + 1], n).float().to(device) base = generator.decode(z, one_src) codes = torch.cat([(1 - a) * one_src + a * one_tgt for a in alphas]) seq = _cf_edit(generator, x.expand(len(alphas), -1, -1, -1), z.expand(len(alphas), *z.shape[1:]), base.expand(len(alphas), -1, -1, -1), codes) probs = torch.softmax(wrapped(seq), dim=-1) want = target_probs[int(tgt[i])] score = ((probs @ want) / (want @ want)).cpu().numpy() span = float(x.max() - x.min()) or 1.0 change = (seq[-1] - x[0]).abs() / span rows.append({ 'source_class': int(src[i]), 'target_class': int(tgt[i]), 'score_start': float(score[0]), 'score_end': float(score[-1]), 'score_path': ';'.join(f'{s:.4f}' for s in score), 'spearman': _spearman(alphas.numpy(), score), 'monotone': bool(np.all(np.diff(score) >= -monotone_tolerance)), 'flipped': bool(int(probs[-1].argmax()) == int(want.argmax())), 'edit_l1': float(change.mean()), 'changed_fraction': float((change > 0.1).float().mean()), }) if finals is not None: finals.append(seq[-1:].cpu()) if i < int(keep): frames.append(seq.cpu().numpy()) return rows, (np.stack(frames) if frames else np.zeros((0,))) def _class_mean_baseline(model: nn.Module, train: torch.Tensor, test: torch.Tensor, *, device: Any = 'cpu', train_labels=None, test_labels=None, target_probs=None, targets=None): """Flip rate and edit size of the naive counterfactual, a class-mean shift. Each test crop gets the difference between the mean training crop of its target class and of its own predicted class added to it. A generator that does no better than this has learned nothing beyond the average difference between the classes. With condition labels and ``target_probs`` the means are per condition and a flip is the target condition's most likely class. ``targets`` gives each test crop's target code; ``None`` uses the next code after its own. :returns: ``(flip_rate, median_edit_l1)``, NaN when a class has no training crops. """ device = torch.device(device) wrapped = ClassScoreModel(model).to(device).eval() train, test = train.float().to(device), test.float().to(device) train_labels, target_probs = _cf_codes(wrapped, train, train_labels, target_probs) src, _ = _cf_codes(wrapped, test, test_labels, target_probs) n = target_probs.shape[0] means = [] for c in range(n): members = train[(train_labels == c).to(device)] if members.shape[0] == 0: return float('nan'), float('nan') means.append(members.mean(dim=0)) tgt = (src + 1) % n if targets is None else torch.as_tensor(targets).long() shift = torch.stack([means[int(t)] - means[int(s)] for s, t in zip(src, tgt)]) moved = test + shift wanted = target_probs.argmax(dim=1)[tgt].to(device) flipped = (_cf_scores(wrapped, moved).argmax(dim=1) == wanted).float().mean() span = (test.amax(dim=(1, 2, 3)) - test.amin(dim=(1, 2, 3))).clamp_min(1e-8) l1 = shift.abs().mean(dim=(1, 2, 3)) / span return float(flipped), float(l1.median()) def _cf_features(x: torch.Tensor) -> np.ndarray: """Interpretable per-crop features for the realism score. Per channel: mean, standard deviation, the mean of a central disk over the mean of the ring around it (where a stain sits relative to the centred object -- for a translocation assay, the phenotype itself), and the mean gradient magnitude (texture and sharpness, which a blurry generator gets wrong first). Hand features rather than an ImageNet network's, because a microscopy crop is not an ImageNet picture and a distance in that space says little about cells. """ x = x.float() n, c, h, w = x.shape if n == 0: return np.zeros((0, 4 * c)) yy, xx = torch.meshgrid(torch.arange(h, dtype=torch.float32), torch.arange(w, dtype=torch.float32), indexing='ij') r = ((yy - (h - 1) / 2) ** 2 + (xx - (w - 1) / 2) ** 2).sqrt() disk = (r <= min(h, w) / 6).float() ring = ((r > min(h, w) / 6) & (r <= min(h, w) / 3)).float() flat = x.reshape(n, c, -1) feats = [flat.mean(-1), flat.std(-1)] feats.append((x * disk).sum((2, 3)) / disk.sum() / ((x * ring).sum((2, 3)) / ring.sum()).clamp_min(1e-6)) gy = (x[:, :, 1:, :] - x[:, :, :-1, :]).abs().mean((2, 3)) gx = (x[:, :, :, 1:] - x[:, :, :, :-1]).abs().mean((2, 3)) feats.append(gx + gy) return torch.cat(feats, dim=1).numpy().astype(np.float64) def _cf_frechet(a: np.ndarray, b: np.ndarray) -> float: """Frechet distance between Gaussians fitted to two feature sets. The FID formula on :func:`_cf_features`; features are standardised by ``b`` first so no feature dominates by its units. NaN with fewer than two rows in either set. """ from scipy import linalg if len(a) < 2 or len(b) < 2: return float('nan') mu, sd = b.mean(0), b.std(0) + 1e-8 a, b = (a - mu) / sd, (b - mu) / sd ma, mb = a.mean(0), b.mean(0) ca, cb = np.cov(a, rowvar=False), np.cov(b, rowvar=False) root = np.real(linalg.sqrtm(ca @ cb)) return float(((ma - mb) ** 2).sum() + np.trace(ca + cb - 2 * root)) def _cf_target_index(target: Any, levels: Sequence[str], n_classes: int) -> int: """The code index of a chosen counterfactual target. With condition ``levels`` the target is one of them by name; otherwise it is a class index of the classifier, ``0`` to ``n_classes - 1``. :raises ValueError: when the target is not among them. """ text = str(target).strip() if levels: if text not in levels: raise ValueError(f'target condition {text!r} is not among the ' f'training conditions {list(levels)}') return list(levels).index(text) try: index = int(float(text)) except ValueError: index = -1 if not 0 <= index < int(n_classes): raise ValueError(f'target class {text!r} is not a class index from ' f'0 to {int(n_classes) - 1}') return index def _counterfactual_report(model: nn.Module, crops: torch.Tensor, *, names: Optional[Sequence[str]] = None, epochs: int = 30, steps: int = 7, holdout: float = 0.25, show: int = 6, seed: int = 0, device: Any = 'cpu', out_dir: Optional[str] = None, conditions: Optional[Sequence[str]] = None, target: Any = None, generator: str = 'autoencoder', generator_options: Optional[Dict[str, Any]] = None): """Train a counterfactual generator on crops and score it on held-out ones. The crops are split, seeded, into a training part and a held-out part. The generator is trained on the first; every held-out crop is morphed toward the other class and scored. The summary reports the held-out flip rate, how often the classifier's target score rises monotonically along the sequence and its mean Spearman correlation with the step, the edit size, and the same flip rate and edit size for a class-mean shift, so the generator is judged against the naive answer. :param model: the trained classifier. :param crops: ``(N, C, H, W)`` input-ready crops, at least four. :param names: optional name per crop, written with its row. :param epochs: generator training epochs. :param steps: frames per counterfactual sequence. :param holdout: fraction of crops held out for scoring. :param show: sequences drawn in the figure. :param seed: seed for the split and the training. :param device: torch device. :param out_dir: when given, ``counterfactual_cells.csv``, ``counterfactual_summary.csv`` and a ``counterfactual_sequences`` figure are written there. :param conditions: optional condition per crop (for example its well). Given, the generator morphs toward conditions rather than classes: each condition's target is the classifier's mean class probabilities over its training crops, every held-out crop moves to the next condition in sorted order, and the rows name both conditions. :param target: optional target every held-out crop morphs toward: a condition name with ``conditions``, otherwise a class index. Held-out crops already in the target are left out. ``None`` or an empty string morphs each crop to the next class or condition. :param generator: one of :data:`COUNTERFACTUAL_GENERATORS`: ``'autoencoder'`` (the classifier-guided conditional autoencoder) or ``'diffusion'`` (a class-conditional diffusion model, not trained against the classifier; see :class:`_CounterfactualDiffusion`). With ``out_dir`` the diffusion weights are saved there as ``counterfactual_diffusion.pt``. :param generator_options: extra keyword arguments for the generator's trainer, for example ``width`` or ``strength`` for diffusion. :returns: ``(summary, rows, frames)``: a dict of summary metrics, the held-out rows of ``_counterfactual_sequences`` and their frames. :raises ValueError: for fewer than four crops, fewer than two conditions among the training crops, a target that is not a class or training condition, or no held-out crop outside the target. """ crops = torch.as_tensor(crops).float() n = crops.shape[0] if n < 4: raise ValueError(f'counterfactuals need at least 4 crops; got {n}') order = torch.randperm(n, generator=torch.Generator().manual_seed(int(seed))) n_test = min(n - 2, max(2, int(round(n * float(holdout))))) test_idx, train_idx = order[:n_test], order[n_test:] codes = {'labels': None, 'target_probs': None} levels: List[str] = [] if conditions is not None: conditions = [str(c) for c in conditions] levels = sorted({conditions[i] for i in train_idx.tolist()}) if len(levels) < 2: raise ValueError('counterfactuals by condition need at least 2 ' f'conditions among the training crops; got {levels}') index = {c: k for k, c in enumerate(levels)} keep = [i for i in test_idx.tolist() if conditions[i] in index] if not keep: raise ValueError('no held-out crop belongs to a training condition') test_idx = torch.as_tensor(keep, dtype=torch.long) labels = torch.as_tensor([index.get(c, 0) for c in conditions]) wrapped = ClassScoreModel(model).to(torch.device(device)).eval() probs = _cf_scores(wrapped, crops[train_idx].to(torch.device(device))).cpu() target_probs = torch.stack([probs[labels[train_idx] == k].mean(dim=0) for k in range(len(levels))]) codes = {'labels': labels, 'target_probs': target_probs} goal = None if target is not None and str(target).strip(): wrapped = ClassScoreModel(model).to(torch.device(device)).eval() scores = _cf_scores(wrapped, crops[test_idx].to(torch.device(device))) goal = _cf_target_index(target, levels, scores.shape[1]) own = (scores.argmax(dim=1).cpu() if codes['labels'] is None else codes['labels'][test_idx]) test_idx = test_idx[own != goal] if not len(test_idx): raise ValueError('every held-out crop is already in the target ' f'{target!r}') kind = str(generator or 'autoencoder').strip().lower() if kind not in COUNTERFACTUAL_GENERATORS: raise ValueError(f'counterfactual generator {generator!r} is not one of ' f'{list(COUNTERFACTUAL_GENERATORS)}') trainer = (_train_counterfactual_diffusion if kind == 'diffusion' else _train_counterfactual_generator) generator, history = trainer( model, crops[train_idx], epochs=epochs, seed=seed, device=device, labels=None if codes['labels'] is None else codes['labels'][train_idx], target_probs=codes['target_probs'], **dict(generator_options or {})) goals = None if goal is None else [goal] * len(test_idx) finals: list = [] rows, frames = _counterfactual_sequences( model, generator, crops[test_idx], targets=goals, steps=steps, keep=show, device=device, labels=None if codes['labels'] is None else codes['labels'][test_idx], target_probs=codes['target_probs'], finals=finals) realism, realism_source = _cf_realism( model, crops[train_idx], crops[test_idx], finals, rows, device=device, train_labels=None if codes['labels'] is None else codes['labels'][train_idx]) names = list(names) if names is not None else [str(i) for i in range(n)] for row, i in zip(rows, test_idx.tolist()): row['name'] = names[i] if levels: row['source_condition'] = levels[row['source_class']] row['target_condition'] = levels[row['target_class']] base_flip, base_l1 = _class_mean_baseline( model, crops[train_idx], crops[test_idx], device=device, train_labels=None if codes['labels'] is None else codes['labels'][train_idx], test_labels=None if codes['labels'] is None else codes['labels'][test_idx], target_probs=codes['target_probs'], targets=goals) spear = np.array([r['spearman'] for r in rows], dtype=float) summary = { 'train_crops': int(len(train_idx)), 'heldout_crops': int(len(test_idx)), 'epochs': int(epochs), 'steps': int(steps), 'reconstruction_mse': history[-1]['reconstruction'] if history else float('nan'), 'flip_rate': float(np.mean([r['flipped'] for r in rows])), 'monotone_fraction': float(np.mean([r['monotone'] for r in rows])), 'mean_spearman': float(np.nanmean(spear)) if np.isfinite(spear).any() else float('nan'), 'median_edit_l1': float(np.median([r['edit_l1'] for r in rows])), 'median_changed_fraction': float(np.median([r['changed_fraction'] for r in rows])), 'baseline_flip_rate': base_flip, 'baseline_median_edit_l1': base_l1, 'conditions': len(levels), 'target': '' if goal is None else (levels[goal] if levels else str(goal)), 'generator': kind, 'realism_fd': realism, 'realism_fd_unedited': realism_source, } if out_dir: _write_counterfactual_outputs(out_dir, summary, rows, frames) if kind == 'diffusion': import os torch.save({'state_dict': generator.state_dict(), 'config': {'channels': generator.channels, 'side': generator.side, 'n_classes': generator.n_classes, 'width': generator.width, 'timesteps': generator.timesteps, 'strength': generator.strength, 'sample_steps': generator.sample_steps, 'guidance_scale': generator.guidance_scale}, 'levels': list(levels)}, os.path.join(out_dir, 'counterfactual_diffusion.pt')) return summary, rows, frames def _cf_realism(model: nn.Module, train: torch.Tensor, test: torch.Tensor, finals: List[torch.Tensor], rows: List[Dict[str, Any]], *, device: Any = 'cpu', train_labels=None): """How close counterfactuals look to real crops of their target. For each target code, the Frechet distance (:func:`_cf_frechet` on :func:`_cf_features`) between the counterfactuals aimed at it and the real training crops that carry it, weighted by how many counterfactuals aimed there; and the same for the UNEDITED source crops, which is the distance the edit started from. A useful edit brings the first well below the second. :returns: ``(realism_fd, realism_fd_unedited)``; NaN when a target has too few real crops to compare against. """ if not finals: return float('nan'), float('nan') wrapped = ClassScoreModel(model).to(torch.device(device)).eval() codes = (_cf_scores(wrapped, train.float().to(torch.device(device))).argmax(1).cpu() if train_labels is None else torch.as_tensor(train_labels).long()) edited = torch.cat(finals, dim=0) targets = np.array([r['target_class'] for r in rows]) total = weight = 0.0 total_src = 0.0 for code in np.unique(targets): real = train[(codes == int(code)).cpu()] pick = targets == code d = _cf_frechet(_cf_features(edited[pick]), _cf_features(real)) d0 = _cf_frechet(_cf_features(test.float().cpu()[pick]), _cf_features(real)) if not (np.isfinite(d) and np.isfinite(d0)): return float('nan'), float('nan') total += d * pick.sum() total_src += d0 * pick.sum() weight += pick.sum() return float(total / weight), float(total_src / weight) def _write_counterfactual_outputs(out_dir: str, summary: Dict[str, Any], rows: List[Dict[str, Any]], frames: np.ndarray) -> None: """Write the counterfactual tables and the sequence figure to ``out_dir``.""" import os import pandas as pd from .tabular import write_table from .plot import save_figure os.makedirs(out_dir, exist_ok=True) write_table(pd.DataFrame(rows), os.path.join(out_dir, 'counterfactual_cells.csv')) write_table(pd.DataFrame([summary]), os.path.join(out_dir, 'counterfactual_summary.csv')) if frames.ndim != 5 or not len(frames): return np.save(os.path.join(out_dir, 'counterfactual_frames.npy'), frames.astype(np.float32)) import matplotlib.pyplot as plt from .figures.style import figure_style, theme_target n_rows, n_steps = frames.shape[0], frames.shape[1] with figure_style(theme_target()): fig, axes = plt.subplots(n_rows, n_steps, squeeze=False, figsize=(1.4 * n_steps, 1.5 * n_rows)) from .figures.bundle import _register_figure_data _register_figure_data(fig, None, kind="montage", title="Counterfactual paths") for r in range(n_rows): seq = frames[r] lo, hi = float(seq[0].min()), float(seq[0].max()) path = [float(s) for s in rows[r]['score_path'].split(';')] for k in range(n_steps): img = seq[k] img = img[0] if img.shape[0] != 3 else np.moveaxis(img, 0, -1) img = np.clip((img - lo) / ((hi - lo) or 1.0), 0, 1) ax = axes[r][k] ax.imshow(img, cmap=None if img.ndim == 3 else 'gray') ax.set_xticks([]) ax.set_yticks([]) ax.set_title(f'p={path[k]:.2f}', fontsize=7) source = rows[r].get('source_condition', rows[r]['source_class']) target = rows[r].get('target_condition', rows[r]['target_class']) axes[r][0].set_ylabel(f"{source}→{target}", fontsize=7) fig.suptitle('Counterfactual sequences (classifier score for the target class)', fontsize=8) save_figure(fig, os.path.join(out_dir, 'counterfactual_sequences.pdf'), close=True)