Skip to content

Fix per-layer head result cache validation - #1843

Merged
jlarson4 merged 3 commits into
TransformerLensOrg:devfrom
yuanwuyuan9:fix/head-results-cache-validation
Oct 2, 2026
Merged

jlarson4 merged 3 commits into
TransformerLensOrg:devfrom
yuanwuyuan9:fix/head-results-cache-validation

Conversation

@yuanwuyuan9

@yuanwuyuan9 yuanwuyuan9 commented Oct 1, 2026 •

Copy link
Copy Markdown

Description

compute_head_results() currently uses layer 0 to decide whether results are cached for every layer and treats valid batchless results as stale. This can leave missing or merged results untouched, overwrite existing results, or raise KeyError during repeated calls, including through stack_head_results().

Validate each layer independently, using has_batch_dim to distinguish valid 4D and 3D per-head results and checking the head count. Preserve valid cached tensors and recompute only invalid or missing results from cached hook_z and the model's W_O. Retain the existing warning when all layers are already cached. The multiplication and reduction are unchanged; no new dependencies are required.

Fixes #1842

Type of change

  • Bug fix (non-breaking change which fixes an issue)

Validation

The existing LN/RMS unit fixture covers per-layer cache states, incorrect head counts, warning behavior, preservation of cached edits, repeated calls through stack_head_results(), and missing required activations. The new acceptance test sits next to the existing compute-head-results test and uses a separate distilgpt2 Bridge to exercise forward-produced batchless results without manual cache edits.

Local validation on Linux, Python 3.12.3 and PyTorch 2.11.0+cu130, with CUDA disabled:

Validation Result
Expanded review regression checks on the initial fix 8 failed, 24 passed, 56 deselected, 2 warnings
Final cache unit test file, including all 32 head-result regression cases 88 passed, 2 warnings
New batchless forward acceptance test (results-only and full caching) 1 passed, 42 deselected, 2 warnings
All related acceptance tests 12 passed, 31 deselected, 2 warnings
Full unit suite 6820 passed, 61 skipped, 57 deselected, 3 xfailed, 306 warnings; no reruns
mypy Success, 387 source files
Format and diff checks Ran make format, make check-format, and git diff --check

The final full unit run used OMP_NUM_THREADS=1 MKL_NUM_THREADS=1. Earlier default-thread runs required retries in the GPT-2 component benchmark. With retries disabled, both the unmodified 02a7f5a0 baseline and the initial fix failed on blocks.11.attn with identical reported differences; both passed with single-thread settings. The numerical root cause remains undetermined.

Checklist:

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

Comment thread tests/unit/test_activation_cache.py
Comment thread transformer_lens/ActivationCache.py Outdated
@jlarson4

jlarson4 commented Oct 1, 2026

Copy link
Copy Markdown
Collaborator

Looks great @yuanwuyuan9, thank you for getting to those comments so swiftly! I will merge as soon as it passes CI

@yuanwuyuan9

Copy link
Copy Markdown
Author

Hi @jlarson4, thanks for reviewing!
The conflict arose because #1844 added tests at the same location in test_activation_cache.py. I’ve merged the latest dev and preserved both sets of tests. All 185 tests in the two affected test files pass, along with formatting and mypy checks.

@jlarson4

jlarson4 commented Oct 2, 2026

Copy link
Copy Markdown
Collaborator

Excellent, thanks for resolving this conflicts @yuanwuyuan9

@jlarson4
jlarson4 merged commit 6affbe9 into TransformerLensOrg:dev Oct 2, 2026
27 checks passed
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.

2 participants