Repository navigation
Ignore index - #28
Open
jonscheunemann wants to merge 2 commits into
Open
Ignore index#28jonscheunemann wants to merge 2 commits into
jonscheunemann wants to merge 2 commits into
Conversation
perspic has never had a concept of masked targets: both SamplewiseCalculatorOpacus and SamplewiseCalculatorFunctorch normalized with n_elements = targets.numel(), and the chi_net Hutchinson/exact projection ran over the whole model output. That is exactly right when every target is real (e.g. packed language-model batches, or ordinary classification), which is why it went unnoticed for this long -- it's a missing feature, not a latent bug. It breaks for a padded batch whose loss uses `ignore_index` (e.g. torch's own CrossEntropyLoss convention, -100): n_elements over-counts the padded positions, and the network gradient projection includes sensitivity at positions whose target doesn't exist, both silently. Add an `ignore_index` constructor argument (default -100, matching nn.CrossEntropyLoss; pass None to disable) to both calculators. compute() resolves a shared mask via the new SamplewiseCalculator.resolve_target_mask(), used for both (a) n_elements (real-position count instead of targets.numel()) and (b) the network-gradient projection, which now multiplies by the mask before summing/backprop (Opacus: `(out * v * mask).sum()` in the Hutchinson path, `out * mask` before indexing in the exact path; functorch: masks the exact per-sample Jacobian before squaring). The Rademacher draws themselves are untouched, and resolve_target_mask returns None -- not an all-True tensor -- whenever a batch has no ignore_index value (or masking is disabled, or targets are float/complex), so the masking multiply is skipped entirely and an unmasked batch takes the exact pre-existing code path. Both calculators' decisions are made by the one shared, static resolve_target_mask so they can't disagree about what n_elements means -- see below for why that agreement matters. Verified before writing any of this: `origin/main` (0171e4f) has no `ignore_index` anywhere, so this is new ground, not a re-fix of Konsti's PR #19 (which taught these calculators about N-D outputs but never masking). Also verified: CouplingCalculator.calculate() computes grad_norm_squared / (chi_loss * chi_net), and chi_loss is multiplied by n_elements while chi_net is divided by it, so n_elements cancels exactly in that product -- confirmed numerically to 1e-17 in TestCouplingCancellation. Changing what n_elements means (real-token count instead of padded length) therefore cannot move chi_pos/chi_coup by itself; only the now-masked chi_net projection can, and only through positions that were never real to begin with. Tests (tests/unit/test_samplewise{,_opacus,_functorch}.py): - resolve_target_mask/broadcast_mask unit tests (no-op cases, dtype guards, custom ignore_index). - Unmasked-batch tests reproducing the pre-feature formula by hand and asserting torch.equal (not allclose) against compute()'s output, for both backends. - V=5/T=4 exact-mode test (Opacus) and its functorch equivalent: a right-padded, -100-masked batch's per-sample chi_net must equal the values computed by running each sample's own unpadded (shorter) sequence through the model individually. Uses an Embedding->Linear model applied per-position (no cross-position mixing), so a masked position provably cannot influence a real position's gradient. Measured: opacus exact mode matches to a max relative error of ~8e-8 (float32 summation-order noise); functorch matches exactly (0.0 diff, since it computes the full Jacobian analytically instead of iterating per output dimension). - Opacus approximate (Hutchinson) mode converges to the same unpadded ground truth within 15% at n=200, averaged over 8 draws. Cherry-picked from ws5-ignore-index (c071499) onto main. The N-D output guard tests from the batch-accumulation branch are not included. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com> Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
analyzer() gains an `ignore_index` argument. By default it is read from the wrapped module's `criterion.ignore_index` (falling back to -100 when the criterion has none); an explicit int overrides it and None disables masking. The value is passed to both sample-wise calculator engines and is not forwarded to the wrapped module's __init__. resolve_target_mask now warns when every target position is ignore_index, since chi_net/chi_loss are then NaN/0 for that batch. Docstrings note that targets must line up with the output's leading axes and that criteria which shift labels internally must be given already-shifted targets. New tests: - compute() on a right-padded, -100-masked batch with a mean-reduced CrossEntropyLoss matches the unpadded ground truth for both backends (covers mask passing through compute(), n_elements = mask.sum(), and the masked chi_loss) - Opacus (exact) and functorch agree on the same masked batch - 2-D classification: ignored samples get zero per-sample chi_net - analyzer ignore_index resolution and forwarding - the all-ignored warning Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Contributor
There was a problem hiding this comment.
Copilot review overview
🔵 Needs a closer look
The masking changes span multiple calculation backends and require final human review.
Review effort: Lite
Findings: None
What changed in this PR
Adds ignore_index support for masked targets across sample-wise calculators and analyzer configuration.
Changes:
- Added shared target-mask resolution and broadcasting helpers.
- Updated Functorch and Opacus masking and normalization.
- Added analyzer configuration and comprehensive tests.
| File | Description |
|---|---|
tests/unit/test_samplewise.py |
Helper and coupling tests |
tests/unit/test_samplewise_opacus.py |
Opacus masking tests |
tests/unit/test_samplewise_functorch.py |
Functorch masking tests |
tests/unit/test_analyzer.py |
Analyzer configuration tests |
perspic/calculator/samplewise.py |
Shared masking utilities |
perspic/calculator/samplewise_opacus.py |
Opacus masking and normalization |
perspic/calculator/samplewise_functorch.py |
Functorch masking and normalization |
perspic/analyzer.py |
Analyzer ignore_index configuration |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
Member
|
Looks like quite a bit of a backend addition. From a first look this seems quite good! |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Support
ignore_index-masked targets in sample-wise calculatorsTarget positions equal to
ignore_index(e.g.-100padding) are now left out of the χ_net projection and then_elementsnormalization. This matches hownn.CrossEntropyLosshandles them.Changes
perspic/calculator/samplewise.pyresolve_target_mask(targets, ignore_index)returns(mask, n_elements). Both backends call it, so they always use the same element count. That shared count is what lets the normalization cancel inCouplingCalculator.maskisNonewhen masking is disabled, when targets are float/complex, or when the batch contains no ignored positions. Unmasked batches therefore compute exactly the same result as before.broadcast_mask(mask, ndim)adds trailing singleton axes so the mask broadcasts against gradient and output tensors. mask: (B, T, 1) -> output: (B, T, V)samplewise_functorch.py/samplewise_opacus.pyignore_index=-100constructor argument.compute()resolves the mask once and normalizes by the maskedn_elements.g**2before the reduction.out * vin Hutchinson mode, and to the flattened output before each backward pass in exact mode.perspic/analyzer.pyignore_indexargument onanalyzer(). If omitted, it is read fromcriterion.ignore_index, falling back to-100. An int overrides that;Nonedisables masking.Tests
ignore_index, the all-ignored warning, and broadcasting.ignore_indexis not passed on to the wrapped module's__init__.Limitations
(B, T)targets for(B, T, V)logits. A(B, V, T)layout is not supported.targets[t] = input_ids[t+1], last/pad = -100, same lengthT), with a non-shifting criterion. Otherwise the mask is off by one position.