diff --git a/perspic/analyzer.py b/perspic/analyzer.py index 1c842e9..56a2138 100644 --- a/perspic/analyzer.py +++ b/perspic/analyzer.py @@ -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, @@ -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. @@ -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. @@ -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) @@ -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 @@ -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, ) diff --git a/perspic/calculator/samplewise.py b/perspic/calculator/samplewise.py index e49d57c..5131411 100644 --- a/perspic/calculator/samplewise.py +++ b/perspic/calculator/samplewise.py @@ -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 @@ -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], @@ -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. @@ -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). diff --git a/perspic/calculator/samplewise_functorch.py b/perspic/calculator/samplewise_functorch.py index 3509a16..9752a10 100644 --- a/perspic/calculator/samplewise_functorch.py +++ b/perspic/calculator/samplewise_functorch.py @@ -1,4 +1,4 @@ -from typing import Callable, Dict +from typing import Callable, Dict, Optional import torch import torch.func as func @@ -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, @@ -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 = ( @@ -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 @@ -94,7 +114,10 @@ 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). @@ -102,6 +125,13 @@ def _compute_per_sample_gradient_norm_network( 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). @@ -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 diff --git a/perspic/calculator/samplewise_opacus.py b/perspic/calculator/samplewise_opacus.py index 90ae9c9..042766b 100644 --- a/perspic/calculator/samplewise_opacus.py +++ b/perspic/calculator/samplewise_opacus.py @@ -1,6 +1,6 @@ """Sample-wise gradient norm calculator using Opacus with ghost clipping.""" -from typing import Callable, Dict +from typing import Callable, Dict, Optional import torch import torch.nn as nn @@ -194,15 +194,34 @@ class SamplewiseCalculatorOpacus(SamplewiseCalculator): random projections instead of iterating over all output dimensions. This provides a faster but approximate computation. Defaults to None (exact computation). + 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, similar to the functorch calculator. """ - def __init__(self, strict: bool = False, approximate_with_n: int | None = None): + def __init__( + self, + strict: bool = False, + approximate_with_n: int | None = None, + ignore_index: int | None = -100, + ): self.strict = strict self.approximate_with_n = approximate_with_n + self.ignore_index = ignore_index def compute( self, @@ -228,12 +247,16 @@ def compute( Returns: Dictionary with 'batch_grad_norms_network' and 'batch_grad_norms_loss'. """ + mask, n_elements = SamplewiseCalculatorOpacus.resolve_target_mask( + targets, self.ignore_index + ) batch_grad_norms_network = ( SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( model, inputs, strict=self.strict, approximate_with_n=self.approximate_with_n, + mask=mask, ) ) batch_grad_norms_loss = ( @@ -244,7 +267,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 @@ -261,6 +283,7 @@ def _compute_per_sample_gradient_norm_network( reduce: bool = True, strict: bool = False, approximate_with_n: int | None = None, + mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Compute per-sample gradient norms for network parameters (∇_θ f). @@ -276,6 +299,13 @@ def _compute_per_sample_gradient_norm_network( approximate_with_n: If not None, the sample-wise gradients will not be computed for each output dimension. Instead, we will apply n low-dimensional projections to estimate the sum of output dimensions. + 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). @@ -298,6 +328,18 @@ def _compute_per_sample_gradient_norm_network( model, strict=strict, loss_reduction="sum" ) + # Mask is resolved once, outside both loops -- it doesn't change + # per-iteration, and (per resolve_target_mask's contract) is None + # whenever no position actually needs masking, so the `out * mask` + # multiply below is skipped entirely and every op that follows + # (including the Rademacher draws) is byte-for-byte the same + # sequence of operations as before this parameter existed. + mask_full = None + if mask is not None: + mask_full = SamplewiseCalculatorOpacus.broadcast_mask( + mask, sample_out.dim() + ).to(dtype=sample_out.dtype) + if approximate_with_n is not None: # Implementation of Hutchinson's trace estimator # Each iteration requires a fresh forward pass because Opacus @@ -313,7 +355,10 @@ def _compute_per_sample_gradient_norm_network( gs_model.zero_grad() out = gs_model(inputs) - projected = (out * v).sum() + if mask_full is not None: + projected = (out * v * mask_full).sum() + else: + projected = (out * v).sum() projected.backward() total_sq_norms += gs_model.get_norm_sample() ** 2 @@ -325,6 +370,21 @@ def _compute_per_sample_gradient_norm_network( # Non-2D outputs are reshaped to (B, -1) so indexing is uniform. n_output_dims = sample_out[0].numel() needs_reshape = sample_out.dim() != 2 + # Flatten the mask the same way `out` is about to be + # flattened, so `mask_flat[:, dim]` lines up with + # `out[:, dim]`. Note: `sample_out` was computed on + # `inputs[:1]` (batch size 1, purely for shape probing), while + # `mask_full`'s batch axis is the real batch size -- expand + # against `inputs.shape[0]` + the rest of `sample_out`'s + # shape, never against `sample_out` itself. This is a view, + # not a copy, until reshape needs one, since a size-1 axis + # can't itself be reshaped into the vocab/class axis's size. + mask_flat = None + if mask_full is not None: + full_batch_shape = (inputs.shape[0], *sample_out.shape[1:]) + mask_flat = mask_full.expand(full_batch_shape) + if needs_reshape: + mask_flat = mask_flat.reshape(mask_flat.shape[0], -1) for dim in range(n_output_dims): _reset_opacus_state(model) @@ -332,6 +392,8 @@ def _compute_per_sample_gradient_norm_network( out = gs_model(inputs) if needs_reshape: out = out.reshape(out.shape[0], -1) + if mask_flat is not None: + out = out * mask_flat out[:, dim].sum().backward() total_sq_norms += gs_model.get_norm_sample() ** 2 diff --git a/tests/unit/test_analyzer.py b/tests/unit/test_analyzer.py index 6046012..a77e792 100644 --- a/tests/unit/test_analyzer.py +++ b/tests/unit/test_analyzer.py @@ -170,6 +170,87 @@ def test_log_metrics_flag_false(self, simple_lightning_module): assert model.log_metrics is False +def _make_module_with_criterion(criterion, init_log=None): + """Create a minimal LightningModule using the given criterion.""" + + class CriterionModule(pl.LightningModule): + def __init__(self, **kwargs): + super().__init__() + if init_log is not None: + init_log.update(kwargs) + self.model = nn.Linear(10, 2) + self.criterion = criterion + + def forward(self, x): + return self.model(x) + + def training_step(self, batch, batch_idx): + x, y = batch + return self.criterion(self(x), y) + + def configure_optimizers(self): + return torch.optim.Adam(self.parameters(), lr=0.001) + + return CriterionModule + + +class TestAnalyzerIgnoreIndex: + """Test resolution and forwarding of the ignore_index option.""" + + @pytest.mark.parametrize("engine", ["opacus", "functorch"]) + def test_taken_from_criterion(self, engine): + """ignore_index defaults to the criterion's ignore_index.""" + module = _make_module_with_criterion(nn.CrossEntropyLoss(ignore_index=0)) + model = analyzer(module, sample_wise_engine=engine) + + assert model.sample_calc.ignore_index == 0 + + @pytest.mark.parametrize("engine", ["opacus", "functorch"]) + def test_default_without_criterion_attribute(self, engine): + """A criterion without ignore_index falls back to -100.""" + module = _make_module_with_criterion(nn.MSELoss()) + model = analyzer(module, sample_wise_engine=engine) + + assert model.sample_calc.ignore_index == -100 + + @pytest.mark.parametrize("engine", ["opacus", "functorch"]) + def test_explicit_value_overrides_criterion(self, engine): + """An explicit int takes precedence over the criterion.""" + module = _make_module_with_criterion(nn.CrossEntropyLoss(ignore_index=0)) + model = analyzer(module, sample_wise_engine=engine, ignore_index=5) + + assert model.sample_calc.ignore_index == 5 + + @pytest.mark.parametrize("engine", ["opacus", "functorch"]) + def test_explicit_none_disables_masking(self, engine): + """An explicit None disables masking even if the criterion has one.""" + module = _make_module_with_criterion(nn.CrossEntropyLoss(ignore_index=0)) + model = analyzer(module, sample_wise_engine=engine, ignore_index=None) + + assert model.sample_calc.ignore_index is None + + def test_both_engines_receive_value(self): + """Both engines get a calculator of the right type with the value.""" + module = _make_module_with_criterion(nn.CrossEntropyLoss(ignore_index=3)) + + opacus_model = analyzer(module, sample_wise_engine="opacus") + functorch_model = analyzer(module, sample_wise_engine="functorch") + + assert isinstance(opacus_model.sample_calc, SamplewiseCalculatorOpacus) + assert isinstance(functorch_model.sample_calc, SamplewiseCalculatorFunctorch) + assert opacus_model.sample_calc.ignore_index == 3 + assert functorch_model.sample_calc.ignore_index == 3 + + @pytest.mark.parametrize("ignore_index", [None, 5]) + def test_not_forwarded_to_wrapped_init(self, ignore_index): + """ignore_index must not leak into the wrapped module's __init__.""" + init_log = {} + module = _make_module_with_criterion(nn.CrossEntropyLoss(), init_log=init_log) + analyzer(module, ignore_index=ignore_index, extra=1) + + assert init_log == {"extra": 1} + + class TestAnalyzerInitialization: """Test Analyzer class initialization.""" diff --git a/tests/unit/test_samplewise.py b/tests/unit/test_samplewise.py index 4704e20..d14929a 100644 --- a/tests/unit/test_samplewise.py +++ b/tests/unit/test_samplewise.py @@ -6,6 +6,7 @@ import torch import torch.nn as nn +from perspic.calculator.samplewise import SamplewiseCalculator from perspic.calculator.samplewise_functorch import SamplewiseCalculatorFunctorch from perspic.calculator.samplewise_opacus import SamplewiseCalculatorOpacus from perspic.utils import BatchStatSnapshot @@ -161,3 +162,151 @@ def test_no_warning_when_batchnorm_in_eval_mode(self, calculator): bn_warnings = [x for x in w if "BatchNorm" in str(x.message)] assert len(bn_warnings) == 0 + + +class TestResolveTargetMask: + """Tests for `SamplewiseCalculator.resolve_target_mask`, the single place + both calculator backends decide whether/how to mask -- they must never + disagree, since `compute()` uses the same `n_elements` to normalize both + `batch_grad_norms_network` and `batch_grad_norms_loss`, and only that + agreement is what makes CouplingCalculator's chi_pos independent of + `n_elements` (see `coupling.py` and `TestCouplingCancellation` below). + """ + + def test_ignore_index_none_disables_masking(self): + """Passing ignore_index=None must skip masking even if -100 is + present in targets -- the documented escape hatch for data where + -100 is a legitimate label.""" + targets = torch.tensor([-100, 1, 2]) + mask, n_elements = SamplewiseCalculator.resolve_target_mask( + targets, ignore_index=None + ) + assert mask is None + assert n_elements == targets.numel() + + def test_no_ignore_index_present_returns_none(self): + """A batch that happens to contain no -100 must resolve to mask=None + (not an all-True tensor) -- this is what lets callers skip the + masking multiply entirely and stay bitwise identical to before this + feature existed.""" + targets = torch.tensor([[0, 1, 2], [3, 0, 1]]) + mask, n_elements = SamplewiseCalculator.resolve_target_mask( + targets, ignore_index=-100 + ) + assert mask is None + assert n_elements == targets.numel() + + def test_masks_ignore_index_positions(self): + """A batch containing -100 must produce a boolean mask (True = real + position) shaped like targets, and n_elements = number of real + positions, not targets.numel().""" + targets = torch.tensor([[5, -100, 2, -100], [1, 2, 3, 4]]) + mask, n_elements = SamplewiseCalculator.resolve_target_mask( + targets, ignore_index=-100 + ) + assert mask is not None + assert mask.dtype == torch.bool + assert mask.shape == targets.shape + expected = torch.tensor([[True, False, True, False], [True, True, True, True]]) + assert torch.equal(mask, expected) + assert n_elements == expected.sum() + assert n_elements.item() == 6 # 8 positions total, 2 masked out + + def test_float_targets_never_masked(self): + """One-hot / soft (float) targets must never be masked, even if a + value happens to equal -100.0 -- compute()'s normalize=False is the + documented path for those, per the existing docstring.""" + targets = torch.tensor([[-100.0, 1.0], [0.3, 0.7]]) + mask, n_elements = SamplewiseCalculator.resolve_target_mask( + targets, ignore_index=-100 + ) + assert mask is None + assert n_elements == targets.numel() + + def test_custom_ignore_index(self): + """A non-default ignore_index value must be respected.""" + targets = torch.tensor([0, 1, -1, 2]) + mask, n_elements = SamplewiseCalculator.resolve_target_mask( + targets, ignore_index=-1 + ) + assert mask is not None + assert torch.equal(mask, torch.tensor([True, True, False, True])) + assert n_elements.item() == 3 + + def test_all_ignored_batch_warns(self): + """A batch where every position is ignore_index has no real + positions: warn, but leave the return values unchanged.""" + targets = torch.full((2, 3), -100) + with pytest.warns(UserWarning, match="ignore_index=-100"): + mask, n_elements = SamplewiseCalculator.resolve_target_mask( + targets, ignore_index=-100 + ) + assert mask is not None + assert not mask.any() + assert n_elements.item() == 0 + + def test_partially_ignored_batch_does_not_warn(self): + targets = torch.tensor([[-100, -100], [1, -100]]) + with warnings.catch_warnings(): + warnings.simplefilter("error") + SamplewiseCalculator.resolve_target_mask(targets, ignore_index=-100) + + +class TestBroadcastMask: + """Tests for `SamplewiseCalculator.broadcast_mask`.""" + + def test_appends_trailing_singleton_dims(self): + mask = torch.ones(2, 3, dtype=torch.bool) + out = SamplewiseCalculator.broadcast_mask(mask, ndim=4) + assert out.shape == (2, 3, 1, 1) + + def test_noop_when_already_correct_ndim(self): + mask = torch.ones(2, 3, dtype=torch.bool) + out = SamplewiseCalculator.broadcast_mask(mask, ndim=2) + assert out.shape == (2, 3) + + def test_result_broadcasts_against_target_shape(self): + """The whole point: multiplying against a same-batch, larger tensor + must work via ordinary broadcasting once reshaped.""" + mask = torch.tensor([[True, False], [True, True]]) # (2, 2) + out = torch.ones(2, 2, 5) # e.g. (batch, seq_len, vocab) + mask_b = SamplewiseCalculator.broadcast_mask(mask, out.dim()) + result = out * mask_b + assert result.shape == out.shape + assert torch.all(result[0, 1] == 0) # masked position zeroed + assert torch.all(result[0, 0] == 1) # real position untouched + + +class TestCouplingCancellation: + """Verifies the claim that CouplingCalculator's chi_pos (chi_coup) is + independent of the `n_elements` normalization factor -- chi_loss is + multiplied by it and chi_net divided by it, so it cancels in their + product. This is what makes it safe for `resolve_target_mask` to change + what `n_elements` means (targets.numel() -> count of real positions) + without perturbing chi_pos through the normalization alone; only the + *masked* chi_net projection (a different number now, not just rescaled) + can still move chi_pos. + """ + + def test_chi_coup_invariant_to_normalization_factor(self): + from perspic.calculator.coupling import CouplingCalculator + + raw_net = 3.7 # chi_net before normalization (an arbitrary example) + raw_loss = 5.2 # chi_loss before normalization + delta_loss = -2.0 + calc = CouplingCalculator() + + # Two different normalization factors, e.g. N=20 (targets.numel()) + # vs. M=11 (masked real-position count) -- values chosen to echo + # the ~44.6% pad fraction measured for SimpleStories (N/M ~= 1.8). + n_big, n_small = 20.0, 11.0 + + coup_n = calc.calculate( + delta_loss=delta_loss, chi_loss=raw_loss * n_big, chi_net=raw_net / n_big + ) + coup_m = calc.calculate( + delta_loss=delta_loss, + chi_loss=raw_loss * n_small, + chi_net=raw_net / n_small, + ) + assert coup_n == pytest.approx(coup_m, rel=1e-12) diff --git a/tests/unit/test_samplewise_functorch.py b/tests/unit/test_samplewise_functorch.py index b756c91..1764aff 100644 --- a/tests/unit/test_samplewise_functorch.py +++ b/tests/unit/test_samplewise_functorch.py @@ -326,3 +326,221 @@ def test_reduce_parameter_loss(self): # Sum of per-sample should equal reduced assert torch.allclose(loss_grad_norms_reduced, loss_grad_norms_per_sample.sum()) + + +class Toy3DModel(nn.Module): + """MLP whose output is reshaped to (batch, seq_len, vocab_size) -- see + the identically-named model in test_samplewise_opacus.py.""" + + def __init__(self, input_dim=10, n_hidden=10, seq_len=4, vocab_size=6): + super().__init__() + self.seq_len = seq_len + self.vocab_size = vocab_size + self.fc1 = nn.Linear(input_dim, n_hidden) + self.fc2 = nn.Linear(n_hidden, n_hidden) + self.fc3 = nn.Linear(n_hidden, seq_len * vocab_size) + + def forward(self, x): + x = torch.relu(self.fc1(x)) + x = torch.relu(self.fc2(x)) + out = self.fc3(x) + return out.reshape(x.shape[0], self.seq_len, self.vocab_size) + + +class IndependentPositionModel(nn.Module): + """Embedding -> Linear applied per-position: output position t depends + ONLY on token t's own id, with no mixing across positions -- see the + identically-named model in test_samplewise_opacus.py for why this is + what makes the padded-vs-unpadded equality test below valid.""" + + def __init__(self, vocab_size: int, embed_dim: int, n_classes: int): + super().__init__() + self.embed = nn.Embedding(vocab_size, embed_dim) + self.head = nn.Linear(embed_dim, n_classes) + + def forward(self, x): + return self.head(self.embed(x)) + + +class TestIgnoreIndexMasking: + """Tests for perspic's new support for `-100`-masked targets (WS5a) on + the functorch backend -- mirrors test_samplewise_opacus.py's + TestIgnoreIndexMasking. Functorch has no approximate mode (it always + computes the full Jacobian via jacrev), so there is only one "exact + mode" style equality to check. + """ + + def test_constructor_default_ignore_index(self): + calc = SamplewiseCalculatorFunctorch() + assert calc.ignore_index == -100 + + def test_constructor_accepts_ignore_index_override(self): + calc = SamplewiseCalculatorFunctorch(ignore_index=None) + assert calc.ignore_index is None + + def test_unmasked_batch_compute_matches_pre_feature_formula(self): + """No -100 anywhere in targets -> compute() must reproduce the + pre-feature formula (n_elements = targets.numel(), no masking + anywhere) exactly -- functorch has no randomness anywhere in this + path, so any deviation would mean the new masking logic changed + behavior for callers who never asked for it.""" + torch.manual_seed(5) + model = Toy3DModel(input_dim=6, n_hidden=6, seq_len=3, vocab_size=4) + X = torch.randn(5, 6) + y = torch.randint(0, 4, (5, 3)) # no -100 present + + def loss_fn(outputs, targets): + return nn.functional.cross_entropy( + outputs.reshape(-1, outputs.shape[-1]), + targets.reshape(-1), + reduction="sum", + ) + + pre_feature_net = ( + SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_network( + model, X, reduce=True + ) + ) + pre_feature_loss = ( + SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_loss( + model, loss_fn, X, y, reduce=True + ) + ) + n_elements_old = y.numel() + expected_net = pre_feature_net / n_elements_old + expected_loss = pre_feature_loss * n_elements_old + + calc = SamplewiseCalculatorFunctorch() # default ignore_index=-100 + result = calc.compute(model, loss_fn, X, y, normalize=True) + + assert torch.equal(result["batch_grad_norms_network"], expected_net) + assert torch.equal(result["batch_grad_norms_loss"], expected_loss) + + def test_padded_masked_matches_unpadded_per_sample(self): + """V=5, T=4: a padded-and-masked batch's per-sample chi_net values + must equal the values computed by running each sample's own + (shorter, unpadded) sequence through the model individually.""" + torch.manual_seed(11) + vocab_size, embed_dim, n_classes, seq_len = 5, 4, 5, 4 + model = IndependentPositionModel(vocab_size, embed_dim, n_classes) + + real_lens = [2, 3] + x_padded = torch.randint(0, vocab_size, (2, seq_len)) + y_padded = torch.full((2, seq_len), -100, dtype=torch.long) + for i, rl in enumerate(real_lens): + y_padded[i, :rl] = torch.randint(0, n_classes, (rl,)) + + mask, n_elements = SamplewiseCalculatorFunctorch.resolve_target_mask( + y_padded, ignore_index=-100 + ) + assert mask is not None + assert n_elements.item() == sum(real_lens) + + padded_per_sample = ( + SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_network( + model, x_padded, reduce=False, mask=mask + ) + ) + + unpadded_per_sample = torch.stack( + [ + SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_network( + model, x_padded[i, :rl].unsqueeze(0), reduce=True + ) + for i, rl in enumerate(real_lens) + ] + ) + + assert torch.allclose( + padded_per_sample, unpadded_per_sample, atol=1e-5, rtol=1e-4 + ) + + @staticmethod + def _masked_batch(): + """Right-padded (B=3, T=5) batch with real lengths [2, 5, 3].""" + torch.manual_seed(21) + vocab_size, embed_dim, n_classes, seq_len = 6, 4, 6, 5 + model = IndependentPositionModel(vocab_size, embed_dim, n_classes) + real_lens = [2, 5, 3] + x = torch.randint(0, vocab_size, (3, seq_len)) + y = torch.full((3, seq_len), -100, dtype=torch.long) + for i, rl in enumerate(real_lens): + y[i, :rl] = torch.randint(0, n_classes, (rl,)) + criterion = nn.CrossEntropyLoss(ignore_index=-100) # mean reduction + + def loss_fn(outputs, targets): + return criterion( + outputs.reshape(-1, outputs.shape[-1]), targets.reshape(-1) + ) + + return model, x, y, real_lens, loss_fn + + def test_compute_masked_batch_matches_unpadded_ground_truth(self): + """compute() on a right-padded, -100-masked batch (mean-reduction + CrossEntropyLoss) must equal ground truth built from unpadded data. + + compute() returns batch-summed values. With N real tokens: + network = sum_i ||grad f(x_i[:len_i])||^2 / N (masked positions + contribute nothing, so each unpadded sequence's norm is the target), + loss = N * sum_b ||dL/dlogits_b||^2, where L is the full masked mean + loss and dL/dlogits is zero at ignored positions. + """ + model, x, y, real_lens, loss_fn = self._masked_batch() + n_real = sum(real_lens) + + expected_net = ( + sum( + SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_network( + model, x[i, :rl].unsqueeze(0), reduce=True + ) + for i, rl in enumerate(real_lens) + ) + / n_real + ) + + logits = model(x).detach().requires_grad_(True) + (dlogits,) = torch.autograd.grad(loss_fn(logits, y), logits) + pad_mask = torch.arange(x.shape[1])[None, :] < torch.tensor(real_lens)[:, None] + assert torch.all(dlogits[~pad_mask] == 0) + expected_loss = (dlogits[pad_mask] ** 2).sum() * n_real + + result = SamplewiseCalculatorFunctorch().compute( + model, loss_fn, x, y, normalize=True + ) + + assert torch.allclose( + result["batch_grad_norms_network"], expected_net, atol=1e-6, rtol=1e-4 + ) + assert torch.allclose( + result["batch_grad_norms_loss"], expected_loss, atol=1e-6, rtol=1e-4 + ) + + def test_classification_ignored_samples_have_zero_network_norm(self): + """(B,) targets with some -100 and (B, C) output: ignored samples get + exactly zero per-sample network norm; the rest match the unmasked + per-sample values.""" + torch.manual_seed(31) + model = MLP(output_dim=5) + X = torch.randn(6, 10) + y = torch.randint(0, 5, (6,)) + y[[1, 4]] = -100 + + mask, n_elements = SamplewiseCalculatorFunctorch.resolve_target_mask( + y, ignore_index=-100 + ) + assert mask is not None + assert n_elements.item() == 4 + + masked = ( + SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_network( + model, X, reduce=False, mask=mask + ) + ) + unmasked = ( + SamplewiseCalculatorFunctorch._compute_per_sample_gradient_norm_network( + model, X, reduce=False + ) + ) + + assert torch.equal(masked[~mask], torch.zeros(2)) + assert torch.allclose(masked[mask], unmasked[mask], atol=1e-6, rtol=1e-4) diff --git a/tests/unit/test_samplewise_opacus.py b/tests/unit/test_samplewise_opacus.py index f320697..5991d13 100644 --- a/tests/unit/test_samplewise_opacus.py +++ b/tests/unit/test_samplewise_opacus.py @@ -511,3 +511,291 @@ def test_approximate_reduce_false_returns_per_sample(self): model, inputs, reduce=False, approximate_with_n=3 ) assert result.shape == (batch_size,) + + +class Toy3DModel(nn.Module): + """MLP whose output is reshaped to (batch, seq_len, vocab_size). + + Used to test N-D (rank > 2) output support in the per-sample gradient + calculators -- e.g. a causal LM's (batch, seq_len, vocab_size) logits, + where shape[-1] (vocab_size) and axis-1 (seq_len) are different axes, + unlike the 2-D (batch, num_classes) case every other model in this file + exercises. + """ + + def __init__(self, input_dim=10, n_hidden=10, seq_len=4, vocab_size=6): + super().__init__() + self.seq_len = seq_len + self.vocab_size = vocab_size + self.fc1 = nn.Linear(input_dim, n_hidden) + self.fc2 = nn.Linear(n_hidden, n_hidden) + self.fc3 = nn.Linear(n_hidden, seq_len * vocab_size) + + def forward(self, x): + x = torch.relu(self.fc1(x)) + x = torch.relu(self.fc2(x)) + out = self.fc3(x) + return out.reshape(x.shape[0], self.seq_len, self.vocab_size) + + +class IndependentPositionModel(nn.Module): + """Embedding -> Linear applied per-position: output position t depends + ONLY on token t's own id, with no mixing across positions (unlike + Toy3DModel above, whose fully-connected layers mix every input feature + into every output position). That independence is what makes right- + padding safe to test against truncated unpadded sequences below -- + masking out a padded position cannot leak into any real position's + gradient, because there is no path for it to leak through. + """ + + def __init__(self, vocab_size: int, embed_dim: int, n_classes: int): + super().__init__() + self.embed = nn.Embedding(vocab_size, embed_dim) + self.head = nn.Linear(embed_dim, n_classes) + + def forward(self, x): + return self.head(self.embed(x)) + + +class TestIgnoreIndexMasking: + """Tests for perspic's new support for `-100`-masked targets (WS5a). + + perspic previously had no concept of `ignore_index`: the chi_net + Hutchinson/exact projection ran over the whole output including padded + positions, and `n_elements = targets.numel()` counted them too. These + tests cover the two correctness properties the fix must have: (a) a + batch with no masked targets must be untouched -- bitwise identical to + the pre-feature computation; (b) a padded-and-masked batch's per-sample + values must equal the per-sample values computed on the corresponding + unpadded (shorter) sequences. + """ + + def test_constructor_default_ignore_index(self): + calc = SamplewiseCalculatorOpacus() + assert calc.ignore_index == -100 + + def test_constructor_accepts_ignore_index_override(self): + calc = SamplewiseCalculatorOpacus(ignore_index=None) + assert calc.ignore_index is None + + def test_unmasked_batch_compute_matches_pre_feature_formula(self): + """No -100 anywhere in targets -> compute() must reproduce the + pre-feature formula (n_elements = targets.numel(), no masking + anywhere) exactly, not just approximately. Both computations reuse + the same model/inputs and exact mode has no randomness, so any + deviation would mean the new masking logic changed behavior for + callers who never asked for it.""" + torch.manual_seed(5) + model = Toy3DModel(input_dim=6, n_hidden=6, seq_len=3, vocab_size=4) + X = torch.randn(5, 6) + y = torch.randint(0, 4, (5, 3)) # no -100 present + + def loss_fn(outputs, targets): + return nn.functional.cross_entropy( + outputs.reshape(-1, outputs.shape[-1]), + targets.reshape(-1), + reduction="sum", + ) + + with BatchStatSnapshot(model, X): + pre_feature_net = ( + SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, X, reduce=True + ) + ) + pre_feature_loss = ( + SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_loss( + model, loss_fn, X, y, reduce=True + ) + ) + n_elements_old = y.numel() + expected_net = pre_feature_net / n_elements_old + expected_loss = pre_feature_loss * n_elements_old + + with BatchStatSnapshot(model, X): + calc = SamplewiseCalculatorOpacus() # default ignore_index=-100 + result = calc.compute(model, loss_fn, X, y, normalize=True) + + assert torch.equal(result["batch_grad_norms_network"], expected_net) + assert torch.equal(result["batch_grad_norms_loss"], expected_loss) + + def test_exact_mode_padded_masked_matches_unpadded_per_sample(self): + """V=5, T=4: a padded-and-masked batch's per-sample chi_net values + must equal the values computed by running each sample's own + (shorter, unpadded) sequence through the model individually. This is + the equality the whole feature exists to guarantee.""" + torch.manual_seed(11) + vocab_size, embed_dim, n_classes, seq_len = 5, 4, 5, 4 + model = IndependentPositionModel(vocab_size, embed_dim, n_classes) + + real_lens = [2, 3] + x_padded = torch.randint(0, vocab_size, (2, seq_len)) + y_padded = torch.full((2, seq_len), -100, dtype=torch.long) + for i, rl in enumerate(real_lens): + y_padded[i, :rl] = torch.randint(0, n_classes, (rl,)) + + mask, n_elements = SamplewiseCalculatorOpacus.resolve_target_mask( + y_padded, ignore_index=-100 + ) + assert mask is not None + assert n_elements.item() == sum(real_lens) + + padded_per_sample = ( + SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, x_padded, reduce=False, mask=mask + ) + ) + + unpadded_per_sample = torch.stack( + [ + SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, x_padded[i, :rl].unsqueeze(0), reduce=True + ) + for i, rl in enumerate(real_lens) + ] + ) + + assert torch.allclose( + padded_per_sample, unpadded_per_sample, atol=1e-5, rtol=1e-4 + ) + + def test_approximate_mode_padded_masked_converges_to_unpadded(self): + """Same equality as the exact-mode test above, but through the + Hutchinson approximate path (what every real llama run actually + uses, since a language-model output is far too large for exact + mode) -- averaged over enough draws to beat the estimator's own + variance.""" + torch.manual_seed(13) + vocab_size, embed_dim, n_classes, seq_len = 6, 5, 6, 5 + model = IndependentPositionModel(vocab_size, embed_dim, n_classes) + + real_lens = [2, 4] + x_padded = torch.randint(0, vocab_size, (2, seq_len)) + y_padded = torch.full((2, seq_len), -100, dtype=torch.long) + for i, rl in enumerate(real_lens): + y_padded[i, :rl] = torch.randint(0, n_classes, (rl,)) + + mask, _ = SamplewiseCalculatorOpacus.resolve_target_mask( + y_padded, ignore_index=-100 + ) + + exact_unpadded = torch.stack( + [ + SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, x_padded[i, :rl].unsqueeze(0), reduce=True + ) + for i, rl in enumerate(real_lens) + ] + ) + + approx_runs = torch.stack( + [ + SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, x_padded, reduce=False, mask=mask, approximate_with_n=200 + ) + for _ in range(8) + ] + ) + approx_mean = approx_runs.mean(dim=0) + + rel_error = (approx_mean - exact_unpadded).abs() / exact_unpadded.abs() + assert torch.all(rel_error < 0.15) + + @staticmethod + def _masked_batch(): + """Right-padded (B=3, T=5) batch with real lengths [2, 5, 3].""" + torch.manual_seed(21) + vocab_size, embed_dim, n_classes, seq_len = 6, 4, 6, 5 + model = IndependentPositionModel(vocab_size, embed_dim, n_classes) + real_lens = [2, 5, 3] + x = torch.randint(0, vocab_size, (3, seq_len)) + y = torch.full((3, seq_len), -100, dtype=torch.long) + for i, rl in enumerate(real_lens): + y[i, :rl] = torch.randint(0, n_classes, (rl,)) + criterion = nn.CrossEntropyLoss(ignore_index=-100) # mean reduction + + def loss_fn(outputs, targets): + return criterion( + outputs.reshape(-1, outputs.shape[-1]), targets.reshape(-1) + ) + + return model, x, y, real_lens, loss_fn + + def test_compute_masked_batch_matches_unpadded_ground_truth(self): + """compute() on a right-padded, -100-masked batch (mean-reduction + CrossEntropyLoss) must equal ground truth built from unpadded data. + + compute() returns batch-summed values. With N real tokens: + network = sum_i ||grad f(x_i[:len_i])||^2 / N (masked positions + contribute nothing, so each unpadded sequence's norm is the target), + loss = N * sum_b ||dL/dlogits_b||^2, where L is the full masked mean + loss and dL/dlogits is zero at ignored positions. + """ + model, x, y, real_lens, loss_fn = self._masked_batch() + n_real = sum(real_lens) + + expected_net = ( + sum( + SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, x[i, :rl].unsqueeze(0), reduce=True + ) + for i, rl in enumerate(real_lens) + ) + / n_real + ) + + logits = model(x).detach().requires_grad_(True) + (dlogits,) = torch.autograd.grad(loss_fn(logits, y), logits) + pad_mask = torch.arange(x.shape[1])[None, :] < torch.tensor(real_lens)[:, None] + assert torch.all(dlogits[~pad_mask] == 0) + expected_loss = (dlogits[pad_mask] ** 2).sum() * n_real + + result = SamplewiseCalculatorOpacus().compute( + model, loss_fn, x, y, normalize=True + ) + + assert torch.allclose( + result["batch_grad_norms_network"], expected_net, atol=1e-6, rtol=1e-4 + ) + assert torch.allclose( + result["batch_grad_norms_loss"], expected_loss, atol=1e-6, rtol=1e-4 + ) + + def test_classification_ignored_samples_have_zero_network_norm(self): + """(B,) targets with some -100 and (B, C) output: ignored samples get + exactly zero per-sample network norm; the rest match the unmasked + per-sample values.""" + torch.manual_seed(31) + model = MLP(output_dim=5) + X = torch.randn(6, 10) + y = torch.randint(0, 5, (6,)) + y[[1, 4]] = -100 + + mask, n_elements = SamplewiseCalculatorOpacus.resolve_target_mask( + y, ignore_index=-100 + ) + assert mask is not None + assert n_elements.item() == 4 + + masked = SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, X, reduce=False, mask=mask + ) + unmasked = SamplewiseCalculatorOpacus._compute_per_sample_gradient_norm_network( + model, X, reduce=False + ) + + assert torch.equal(masked[~mask], torch.zeros(2)) + assert torch.allclose(masked[mask], unmasked[mask], atol=1e-6, rtol=1e-4) + + def test_backends_agree_on_masked_batch(self): + """Exact-mode Opacus and functorch compute() must agree on a masked + batch for both keys.""" + model, x, y, _, loss_fn = self._masked_batch() + + opacus_result = SamplewiseCalculatorOpacus().compute(model, loss_fn, x, y) + functorch_result = SamplewiseCalculatorFunctorch().compute(model, loss_fn, x, y) + + for key in ("batch_grad_norms_network", "batch_grad_norms_loss"): + assert torch.allclose( + opacus_result[key], functorch_result[key], atol=1e-6, rtol=1e-4 + ), key