Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 31 additions & 2 deletions perspic/analyzer.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,9 @@
from perspic.logger import LogarithmicWindowSchedule
from perspic.utils import BatchStatSnapshot

# Sentinel meaning "take ignore_index from the wrapped module's criterion".
_FROM_CRITERION = object()


def analyzer(
lightning_module: pl.LightningModule,
Expand All @@ -22,6 +25,7 @@ def analyzer(
analyze_every: Optional[int] = None,
analysis_schedule: Optional[LogarithmicWindowSchedule] = None,
cross_response: bool = False,
ignore_index: Optional[int] = _FROM_CRITERION,
**model_kwargs,
):
"""Factory function that wraps a LightningModule with analysis capabilities.
Expand Down Expand Up @@ -55,6 +59,19 @@ def analyzer(
cross_response: If True, enables cross-batch response analysis and assumes
the training batch is a dict with 'train' and 'measure' keys.
Defaults to False.
ignore_index: Target value marking positions that contribute no loss
(e.g. padding), forwarded to the sample-wise calculators so masked
positions are excluded from the ``chi_net`` projection and the
element count. By default it is read from the wrapped module's
``criterion.ignore_index`` (falling back to -100, the
``nn.CrossEntropyLoss`` convention, if the criterion has none). An
explicit int overrides the criterion; None disables masking
(the behaviour before ignore_index support existed).
Targets passed to the criterion must line up with the logits'
leading axes (e.g. (B, T) targets vs (B, T, V) logits). A criterion
that shifts labels internally (HF-style causal LM, ``logits[:, :-1]``
vs ``labels[:, 1:]``) must be given already-shifted targets by the
time perspic sees them, otherwise the mask is off by one position.
**model_kwargs: Additional keyword arguments passed to the
LightningModule constructor.

Expand Down Expand Up @@ -117,6 +134,7 @@ def __init__(
analyze_every=analyze_every,
analysis_schedule=analysis_schedule,
cross_response=cross_response,
ignore_index=ignore_index,
**model_kwargs,
):
super().__init__(**model_kwargs)
Expand Down Expand Up @@ -145,11 +163,21 @@ def __init__(
if analyze_every is not None and analyze_every < 1:
raise ValueError("analyze_every must be a positive integer")

# The criterion only exists once the wrapped __init__ has run.
if ignore_index is _FROM_CRITERION:
ignore_index = getattr(
getattr(self, "criterion", None), "ignore_index", -100
)

if sample_wise_engine == "functorch":
self.sample_calc = SamplewiseCalculatorFunctorch()
self.sample_calc = SamplewiseCalculatorFunctorch(
ignore_index=ignore_index
)
elif sample_wise_engine == "opacus":
self.sample_calc = SamplewiseCalculatorOpacus(
strict=opacus_strict, approximate_with_n=opacus_approximate_with_n
strict=opacus_strict,
approximate_with_n=opacus_approximate_with_n,
ignore_index=ignore_index,
)

# Initialize the linearizer
Expand Down Expand Up @@ -443,5 +471,6 @@ def _after_training_step(self, batch, batch_idx, output):
analyze_every=analyze_every,
analysis_schedule=analysis_schedule,
cross_response=cross_response,
ignore_index=ignore_index,
**model_kwargs,
)
97 changes: 95 additions & 2 deletions perspic/calculator/samplewise.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

import warnings
from abc import ABC, abstractmethod
from typing import Callable, Dict
from typing import Callable, Dict, Optional, Tuple

import torch
import torch.nn as nn
Expand Down Expand Up @@ -41,6 +41,91 @@ def _warn_if_batchnorm_training(model: nn.Module) -> None:
)
return # Only warn once

@staticmethod
def resolve_target_mask(
targets: torch.Tensor, ignore_index: Optional[int]
) -> Tuple[Optional[torch.Tensor], torch.Tensor]:
"""Determine which target positions are real vs. ignored.

Shared by both calculator backends so "should this batch be masked"
is decided in exactly one place -- the Opacus and functorch
implementations must never disagree about it, since `compute()`
divides `batch_grad_norms_network` and multiplies
`batch_grad_norms_loss` by the *same* `n_elements`, and only that
agreement makes the two normalizations cancel in
``CouplingCalculator`` (chi_loss * chi_net is independent of
`n_elements` -- see `coupling.py`).

Args:
targets: Target tensor of shape (batch, ...), e.g. (B,) for plain
classification or (B, T) for a sequence model's per-position
class indices.
ignore_index: Target value marking a position that contributes no
loss and should be excluded from the sample-wise projection
(mirrors `nn.CrossEntropyLoss`'s `ignore_index`, default
-100). Pass `None` to disable masking entirely regardless of
the batch's contents. `targets` must line up with the model
output's leading axes ((B, T) targets for (B, T, V) output;
a (B, V, T) output layout is not supported). Criteria that
shift labels internally (HF-style causal LM) must be given
already-shifted targets here, otherwise the mask is off by
one position.

Returns:
A `(mask, n_elements)` tuple. `mask` is `None` when no masking
applies -- `ignore_index` is `None`, `targets` is float/complex
(one-hot/soft targets; see `compute()`'s `normalize=False` for
those), or the batch simply contains no `ignore_index` value.
Returning `None` rather than an all-True tensor is what keeps an
unmasked batch's computation bitwise identical to before this
feature existed -- callers must skip the masking multiply
entirely when `mask is None`, not multiply by an all-True mask.
When masking does apply, `mask` is a boolean tensor shaped like
`targets` (True = real, scored position). `n_elements` is
`targets.numel()` when `mask is None`, else `mask.sum()`.

Warns:
UserWarning: If every position is `ignore_index`. The loss then
has no real positions, so chi_net/chi_loss for this batch are
NaN/0 (the return values are unchanged).
"""
if ignore_index is None:
return None, targets.numel()
if targets.is_floating_point() or targets.is_complex():
return None, targets.numel()
is_ignored = targets.eq(ignore_index)
if not bool(is_ignored.any()):
return None, targets.numel()
mask = ~is_ignored
n_elements = mask.sum()
if n_elements == 0:
warnings.warn(
"Every target position equals ignore_index="
f"{ignore_index}; the loss has no real positions, so "
"chi_net/chi_loss will be NaN/0 for this batch.",
UserWarning,
stacklevel=2,
)
return mask, n_elements

@staticmethod
def broadcast_mask(mask: torch.Tensor, ndim: int) -> torch.Tensor:
"""Append trailing singleton axes so `mask` broadcasts against a
tensor of `ndim` dimensions.

`mask` covers every leading axis a label tensor has (e.g. (B, T) for
a sequence model's per-position targets); the tensor it must
broadcast against carries additional trailing axes the label doesn't
have -- at minimum the loss-reduction ("class"/"vocab") axis, and for
functorch's per-sample-gradient tensors, the parameter axes on top of
that. This is a plain reshape (`unsqueeze`), not a copy; broadcasting
multiplication (`*`) handles the rest without materializing the
expanded size.
"""
while mask.dim() < ndim:
mask = mask.unsqueeze(-1)
return mask

@staticmethod
def compute_cross_metrics(
sample_wise_metrics_self: Dict[str, torch.Tensor],
Expand Down Expand Up @@ -99,7 +184,10 @@ def compute(
@staticmethod
@abstractmethod
def _compute_per_sample_gradient_norm_network(
model: nn.Module, inputs: torch.Tensor, reduce: bool = True
model: nn.Module,
inputs: torch.Tensor,
reduce: bool = True,
mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Compute per-sample gradient norms for network parameters.

Expand All @@ -108,6 +196,11 @@ def _compute_per_sample_gradient_norm_network(
inputs: Input tensor batch of shape (batch_size, ...).
reduce: If True, sum over batch dimension. If False, return
per-sample squared norms.
mask: Optional boolean tensor shaped like a label tensor (e.g.
(batch, seq_len) for a sequence model), True at positions to
include in the projection. See `resolve_target_mask`/
`broadcast_mask`. `None` (the default) computes over every
output position, exactly as before this parameter existed.

Returns:
If reduce=True: Scalar tensor (sum of squared gradient norms).
Expand Down
62 changes: 57 additions & 5 deletions perspic/calculator/samplewise_functorch.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import Callable, Dict
from typing import Callable, Dict, Optional

import torch
import torch.func as func
Expand All @@ -13,12 +13,30 @@ class SamplewiseCalculatorFunctorch(SamplewiseCalculator):
This implementation uses PyTorch's functorch (torch.func) for efficient
per-sample gradient computation via vectorized Jacobian calculations.

Args:
ignore_index: Target value marking a position that contributes no loss
and should be excluded from both the network-gradient projection
and the `n_elements` normalization -- mirrors
`nn.CrossEntropyLoss`'s `ignore_index`. Defaults to -100. Pass
`None` to disable masking entirely (e.g. if -100 is a legitimate
target value in your data). Masking only activates when `targets`
is integer-typed and actually contains `ignore_index`; an
unmasked batch's computation is bitwise identical to before this
parameter existed. Targets must line up with the model output's
leading axes ((B, T) vs (B, T, V); (B, V, T) is not supported),
and criteria that shift labels internally (HF-style causal LM)
must be given already-shifted targets, otherwise the mask is off
by one. See `SamplewiseCalculator.resolve_target_mask`.

Note:
For models with BatchNorm, wrap calls with `BatchStatSnapshot` context
manager to freeze running statistics for correct sample-wise
gradient computation.
"""

def __init__(self, ignore_index: int | None = -100):
self.ignore_index = ignore_index

def compute(
self,
model: nn.Module,
Expand All @@ -43,9 +61,12 @@ def compute(
Returns:
Dictionary with 'batch_grad_norms_network' and 'batch_grad_norms_loss'.
"""
mask, n_elements = SamplewiseCalculatorFunctorch.resolve_target_mask(
targets, self.ignore_index
)
batch_grad_norms_network = (
SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_network(
model, inputs
model, inputs, mask=mask
)
)
batch_grad_norms_loss = (
Expand All @@ -56,7 +77,6 @@ def compute(

# Optionally normalize the results
if normalize:
n_elements = targets.numel()
batch_grad_norms_network /= n_elements
batch_grad_norms_loss *= n_elements

Expand Down Expand Up @@ -94,14 +114,24 @@ def model_fn(params, buffers, x):

@staticmethod
def _compute_per_sample_gradient_norm_network(
model: nn.Module, inputs: torch.Tensor, reduce: bool = True
model: nn.Module,
inputs: torch.Tensor,
reduce: bool = True,
mask: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""Compute per-sample gradient norms for network parameters (∇_θ f).

Args:
model: The neural network model.
inputs: Input tensor batch of shape (batch_size, ...).
reduce: If True, sum over batch. If False, return per-sample norms.
mask: Optional boolean tensor shaped like a label tensor (e.g.
(batch, seq_len)), True at output positions to include in the
projection. Positions where mask is False contribute zero to
every sample's squared gradient norm -- see
`SamplewiseCalculator.resolve_target_mask`/`broadcast_mask`.
`None` (the default) computes over every output position,
exactly as before this parameter existed.

Returns:
If reduce=True: Scalar (sum of squared gradient norms).
Expand All @@ -124,9 +154,31 @@ def _compute_per_sample_gradient_norm_network(
assert v.shape[0] == inputs.shape[0]
# Assert that the v.shape[1:] matches the shape of the parameter
assert v.shape[-len(params[k].shape) :] == params[k].shape

# `mask` is shaped like the label tensor (e.g. (B, T)), but each `g`
# here carries the extra size-1 axis `inputs.unsqueeze(1)` introduced
# above (jacrev differentiates the model's full output, "fake batch
# of 1" included), ahead of the output-shape and parameter-shape
# axes -- so unsqueeze at position 1 to line the mask up with that
# axis before appending the trailing (output tail + parameter) axes
# broadcast_mask adds. None (no `-100` in this batch, or masking
# disabled) skips this and every op below is unchanged from before
# this parameter existed.
mask_for_grads = mask.unsqueeze(1) if mask is not None else None

# Compute per-sample gradient magnitude (L2 norm)
per_sample_grad_magnitudes = torch.stack(
[(g**2).sum(dim=tuple(range(1, g.ndim))) for g in per_sample_grads.values()]
[
(
g**2
* SamplewiseCalculatorFunctorch.broadcast_mask(
mask_for_grads, g.dim()
).to(dtype=g.dtype)
if mask_for_grads is not None
else g**2
).sum(dim=tuple(range(1, g.ndim)))
for g in per_sample_grads.values()
]
).sum(
dim=0
) # Sum across parameters
Expand Down
Loading
Loading