Skip to content

Refuse remove_batch_dim for batch > 1 on every caching path - #1860

Closed
raashish1601 wants to merge 1 commit into
TransformerLensOrg:devfrom
raashish1601:fix/1858-remove-batch-dim-batch-size
Closed

raashish1601 wants to merge 1 commit into
TransformerLensOrg:devfrom
raashish1601:fix/1858-remove-batch-dim-batch-size

Conversation

@raashish1601

Copy link
Copy Markdown

Description

remove_batch_dim=True with a batch larger than one now fails the same way on every caching path. Before this, only the ActivationCache path refused it. The plain-dict path of run_with_cache kept the batch dim, and get_caching_hooks / add_caching_hooks cached only the first example.

  • TransformerBridge.get_caching_hooks (and add_caching_hooks, which delegates to it): save_hook asserts the leading dim is 1 before dropping it.
  • TransformerBridge.run_with_cache(..., return_cache_object=False): the dict path now goes through ActivationCache.remove_batch_dim(), so it uses the same batch-size check and only squeezes entries whose leading dim is 1, as before. The dict is still edited in place and returned as a dict.
  • HookedRootModule.add_caching_hooks and get_caching_hooks: same assert as the bridge save_hook, since the silent [0] came from there.

The assert message matches the existing one in ActivationCache.remove_batch_dim ("Cannot remove batch dimension from cache with batch size N").

Fixes #1858

Type of change

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

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

New tests are in tests/unit/test_remove_batch_dim_batch_size.py and use a boot_native model and a small HookedRootModule, so they need no downloads. Four of them fail on dev without the fix. I ran the unit suite (pytest tests/unit -m "not slow"): everything passes except three tests that fail the same way on dev without this change on my Windows machine (test_model_structure_doc hits a cp1252 decode error reading the docs, and two test_sparse_probing tolerance tests stop on line_search). The format checks (pycln, isort, black) and mypy pass on the changed files.

@jlarson4

jlarson4 commented Oct 6, 2026

Copy link
Copy Markdown
Collaborator

Hi @raashish1601! Thank you for putting this together. Unfortunately, this issue is both assigned to a different contributor, and it has been determined that we want to take the solution in a different direction. My apologies.

@jlarson4 jlarson4 closed this Oct 6, 2026
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