fix(typing): correct annotations that current beartype rejects - #1838
Merged
jlarson4 merged 1 commit intoSep 30, 2026
Merged
Conversation
pyproject allows beartype>=0.14.1, but uv.lock still pins 0.14.1. Newer beartype (0.22.x, what a fresh install resolves to) checks dict values and container items, and the jaxtyping test hook then rejects three call sites: - TLWorkerExtension.tl_read_batched_captures returns encode_tensor() dicts, not tensors; annotate it like tl_read_captures. - select_displacement_matched_control_token is called with the decomposition's support tensor; accept a tensor for active_support. - test_resolve_state_dict_key_dense_mlp_fallback passed ints as state-dict values; use tensors.
koriyoshi2041
left a comment
Contributor
There was a problem hiding this comment.
Reviewed at 707bc75f. The three annotation changes match the runtime values: batched captures contain encoded tensor payloads, decomposition.support is a tensor accepted by the membership path, and the state-dict fixture values are tensors. I also ran the affected control-selection, batched worker-extension, and dense-MLP fallback tests under beartype 0.22.5: 5 passed. The hosted 3.10–3.12, type, coverage, and notebook matrix is green.
Collaborator
|
Looks good! Thanks @YHC66 |
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
pyproject.tomlallowsbeartype>=0.14.1, butuv.lockstill pins0.14.1, so CI runs the old version. A fresh install resolves to beartype 0.22.x, which checks dict values and container items, and with the test-time jaxtyping hook (--jaxtyping-packages=transformer_lens,beartype.beartype) 12 unit tests ondevfail withBeartypeCallHintParamViolation/BeartypeCallHintReturnViolation. Two of the three causes are wrong annotations in library code:TLWorkerExtension.tl_read_batched_capturesreturnsencode_tensor()dicts, not tensors. It is now annotatedDict[str, Dict[str, Any]], liketl_read_captures.select_displacement_matched_control_tokenis called byrun_causal_swap_trialwithdecomposition.support, which is a tensor, butactive_supportwas annotatedContainer[int]. It now acceptsUnion[Container[int], torch.Tensor]. The membership check behaves the same for both.test_resolve_state_dict_key_dense_mlp_fallbackpassed ints as state-dict values. It now passes tensors.No runtime behaviour changes, and
uv.lockis left alone.Repro on
dev(687e3f5):uv pip install "beartype==0.22.9" uv run pytest tests/unit/test_weight_processing.py tests/unit/model_bridge/test_vllm_worker_extension.py tests/unit/tools/test_jacobian_lens_causal_swap_benchmark_trials.pytests/unittests/unit(With transformers 5.17 two
TestNemotronHStatefulCachetests also fail on aMagicMockconfig. That is unrelated to beartype and not addressed here.)Type of change
Checklist: