Skip to content

Ignore index - #28

Open
jonscheunemann wants to merge 2 commits into
mainfrom
ignore-index
Open

jonscheunemann wants to merge 2 commits into
mainfrom
ignore-index

Conversation

@jonscheunemann

Copy link
Copy Markdown
Collaborator

Support ignore_index-masked targets in sample-wise calculators

Target positions equal to ignore_index (e.g. -100 padding) are now left out of the χ_net projection and the n_elements normalization. This matches how nn.CrossEntropyLoss handles them.

Changes

perspic/calculator/samplewise.py

  • resolve_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 in CouplingCalculator.
    • mask is None when 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.
    • It warns if every position in a batch is ignored.
  • 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.py

  • New ignore_index=-100 constructor argument.
  • compute() resolves the mask once and normalizes by the masked n_elements.
  • Functorch: the mask is applied to the per-sample g**2 before the reduction.
  • Opacus: the mask is applied to out * v in Hutchinson mode, and to the flattened output before each backward pass in exact mode.

perspic/analyzer.py

  • New ignore_index argument on analyzer(). If omitted, it is read from criterion.ignore_index, falling back to -100. An int overrides that; None disables masking.

Tests

  • Unit tests for both helpers: the no-mask cases, a custom ignore_index, the all-ignored warning, and broadcasting.
  • A coupling test checking that χ_coup does not change with the normalization factor.
  • Both backends: unmasked batches give the same results as before; padded and masked batches give the same per-sample χ_net as the unpadded batch; ignored classification samples get a network norm of zero; Opacus approximate mode converges to the unpadded values.
  • Analyzer: how the value is chosen, and that ignore_index is not passed on to the wrapped module's __init__.

Limitations

  • Targets must line up with the leading axes of the logits: (B, T) targets for (B, T, V) logits. A (B, V, T) layout is not supported.
  • Causal-LM targets must be shifted in the collate (targets[t] = input_ids[t+1], last/pad = -100, same length T), with a non-shifting criterion. Otherwise the mask is off by one position.

jscheunemann and others added 2 commits October 2, 2026 20:11
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>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@KonstiNik

Copy link
Copy Markdown
Member

Looks like quite a bit of a backend addition.

From a first look this seems quite good!
How much of the failure modes do you think these tests cover?
Can we measure the reduction in computations if we have a lot of masking?

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants