Repository navigation
Refuse remove_batch_dim for batch > 1 on every caching path - #1860
Closed
raashish1601 wants to merge 1 commit into
Closed
raashish1601 wants to merge 1 commit into
raashish1601 wants to merge 1 commit into
Conversation
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. |
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.
Description
remove_batch_dim=Truewith a batch larger than one now fails the same way on every caching path. Before this, only theActivationCachepath refused it. The plain-dict path ofrun_with_cachekept the batch dim, andget_caching_hooks/add_caching_hookscached only the first example.TransformerBridge.get_caching_hooks(andadd_caching_hooks, which delegates to it):save_hookasserts the leading dim is 1 before dropping it.TransformerBridge.run_with_cache(..., return_cache_object=False): the dict path now goes throughActivationCache.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_hooksandget_caching_hooks: same assert as the bridgesave_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
Checklist:
New tests are in
tests/unit/test_remove_batch_dim_batch_size.pyand use aboot_nativemodel and a smallHookedRootModule, so they need no downloads. Four of them fail ondevwithout the fix. I ran the unit suite (pytest tests/unit -m "not slow"): everything passes except three tests that fail the same way ondevwithout this change on my Windows machine (test_model_structure_dochits a cp1252 decode error reading the docs, and twotest_sparse_probingtolerance tests stop online_search). The format checks (pycln, isort, black) and mypy pass on the changed files.