Skip to content

fix: preserve attention masks in Inspect HF forwards - #1852

Open
emerardd wants to merge 1 commit into
TransformerLensOrg:devfrom
emerardd:fix/inspect-attention-mask
Open

emerardd wants to merge 1 commit into
TransformerLensOrg:devfrom
emerardd:fix/inspect-attention-mask

Conversation

@emerardd

@emerardd emerardd commented Oct 4, 2026

Copy link
Copy Markdown
Contributor

Description

Fixes #1851.

Inspect RemoteBridge accepted attention_mask, but the Inspect driver omitted it from the request and the HF provider ran an unmasked forward. Masking only the final loss could therefore score logits computed from the wrong context, and cached activations were affected as well.

This change validates and serializes single-sequence binary padding masks, forwards them to the HF-backed tl_bridge provider, and derives padding-aware position IDs only where the model supports ordinary 2-D positions. Model-owned position derivation (including OPT's mask-consuming embeddings and mRoPE) is left intact. No-mask and all-ones-mask behavior is preserved. The provider's completion and logprobs use the last attended token for right-padded inputs.

The Inspect tl_bridge_vllm and vllm-lens profiles reject supplied masks explicitly instead of silently ignoring them; the direct tl_bridge_vllm capture entry point also rejects masks. This does not change the separate native vLLM driver's padding support. The boot_inspect API documentation now states the mask contract and provider limits.

The offline integration regression boots the real Inspect model envelope, driver, and HF provider from locally saved tiny GPT-2, Llama, and OPT models. It covers left/right padding, raw HF parity, unpadded-control logits and residual caches, masked implicit/explicit-label loss, attention patterns and interventions, all-ones masks, binary mask dtypes, malformed masks at both public/provider boundaries, and last-attended-token completion/logprobs. No model downloads or mocked model loads are needed.

Type of change

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

Validation

  • Before the fix, the real tiny GPT-2 provider regression failed for both left and right padding; the all-ones control passed.
  • After the fix, the affected Inspect surface passed: 146 passed, using tests/integration/model_bridge/test_inspect_attention_mask.py, tests/unit/model_bridge/test_inspect_driver.py, tests/unit/model_bridge/test_inspect_vllm_provider.py, and tests/unit/model_bridge/sources/test_inspect_provider_model_class.py.
  • Changed-file Makefile-equivalent pycln/isort/Black checks passed.
  • Full mypy . passed: 388 source files.
  • git diff --check passed.

Validation used the frozen lockfile with the inspect extra. The complete unit/integration/acceptance suite was not run locally; CI will cover the broader surface. The selected run emitted existing SWIG deprecation warnings and the expected GPT-2 fused-QKV capability warning. No live vLLM/GPU run was performed; the new unsupported-mask rejection is tested before a vLLM import or engine call.

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

The unit-test checkbox is left unchecked because only the affected unit-test surface was run, not the complete unit suite.

@jlarson4 jlarson4 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thanks for this quick fix @emerardd! Just a couple small comments in relation to code hygiene

return bool(identity_ok and causal_ok)


def _accepts_mask_positions(model: torch.nn.Module) -> bool:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This mirrors TransformerBridge._accepts_derived_position_ids, but it has already drifted. The local gate accepts a forward(**kwargs) model and this one doesn't, so such a model would get different positions remotely vs locally. Can both call sites use a shared helper in order to stay in sync and reuse the existing gate tests?

profile = profiles.TLBridgeProfile(
supported_kinds=kinds,
provides_sequence_logits=psl,
supports_attention_mask=provider == "tl_bridge",

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The other capabilities in this block are read off the provider API, but mask support is keyed on the provider name here and again in profiles.for_provider. Would it be possible for the provider declare it the way it declares provides_sequence_logits?

This branch has not been deployed

No deployments
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