Skip to content

fix(typing): correct annotations that current beartype rejects - #1838

Merged
jlarson4 merged 1 commit into
TransformerLensOrg:devfrom
YHC66:fix-annotations-newer-beartype
Sep 30, 2026
Merged

jlarson4 merged 1 commit into
TransformerLensOrg:devfrom
YHC66:fix-annotations-newer-beartype

Conversation

@YHC66

@YHC66 YHC66 commented Sep 29, 2026

Copy link
Copy Markdown
Contributor

Description

pyproject.toml allows beartype>=0.14.1, but uv.lock still pins 0.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 on dev fail with BeartypeCallHintParamViolation / BeartypeCallHintReturnViolation. Two of the three causes are wrong annotations in library code:

  • TLWorkerExtension.tl_read_batched_captures returns encode_tensor() dicts, not tensors. It is now annotated Dict[str, Dict[str, Any]], like tl_read_captures.
  • select_displacement_matched_control_token is called by run_causal_swap_trial with decomposition.support, which is a tensor, but active_support was annotated Container[int]. It now accepts Union[Container[int], torch.Tensor]. The membership check behaves the same for both.
  • test_resolve_state_dict_key_dense_mlp_fallback passed ints as state-dict values. It now passes tensors.

No runtime behaviour changes, and uv.lock is 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.py
Environment Before After
beartype 0.22.9, tests/unit 12 beartype failures 0
locked env (beartype 0.14.1), tests/unit pass 6934 passed, 0 failed

(With transformers 5.17 two TestNemotronHStatefulCache tests also fail on a MagicMock config. That is unrelated to beartype and not addressed here.)

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

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 koriyoshi2041 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.

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.

@jlarson4

Copy link
Copy Markdown
Collaborator

Looks good! Thanks @YHC66

@jlarson4
jlarson4 merged commit 3912f92 into TransformerLensOrg:dev Sep 30, 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.

3 participants