Skip to content

fix(bridge): update the KV cache in MPT's attention reconstruction - #1874

Merged
jlarson4 merged 1 commit into
TransformerLensOrg:devfrom
Mudassiruddin7:fix-mpt-kv-cache
Oct 9, 2026
Merged

jlarson4 merged 1 commit into
TransformerLensOrg:devfrom
Mudassiruddin7:fix-mpt-kv-cache

Conversation

@Mudassiruddin7

Copy link
Copy Markdown
Contributor

Description

With the KV cache on, TransformerBridge.generate gives different tokens from Hugging Face for MPT models, while generation without the cache is correct. MPTALiBiAttentionBridge._reconstruct_attention never calls _update_kv_cache, unlike BloomAttentionBridge and the generic joint-QKV bridge. So each cached decoding step builds K and V from the new token alone, and the model does not see the prompt. The prompt pass itself is complete, so the first generated tokens can still be right, and later ones drift.

The fix is one call in mpt_alibi_attention.py, right after the heads are split and before the scores are computed (k, v = self._update_kv_cache(k, v, **kwargs)). _update_kv_cache already reads MPT's layer_past kwarg (its comment lists MPT) and returns K and V unchanged when no cache is passed, so the uncached path is untouched. The ALiBi bias is sliced with position_bias[:, :, -kv_len:] from the score width, so it follows the longer K once the cache is in.

Greedy generation, 8 new tokens, against an untouched generate of the Hugging Face model (fp32, CPU):

dev this PR
hf-internal-testing/tiny-random-MptForCausalLM, cached 8 of 26 ids differ identical
same, uncached identical identical
tiny random MptForCausalLM with wide weights, cached first two new tokens equal, the third differs identical

I found this with a parity battery over the supported architectures that also compares cached and uncached generation. In the same run the full forward pass of MPT matched Hugging Face (max diff 2.4e-06), so only the cached decoding path was affected.

A note for anyone who writes similar checks: build_bridge_from_module replaces the submodules of the model it is given, so the Hugging Face reference has to be generated before bridging or on a deepcopy, which the new test does.

Fixes # (no issue, found by the check above)

Type of change

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

Tests

TestMPTCachedGeneration::test_cached_generation_matches_hf in test_mpt_adapter.py builds a tiny random MptForCausalLM, generates with an untouched copy, then with the bridge with and without the cache, and requires all three to be equal. It fails on dev and passes with this change.

Run locally (Windows, Python 3.12, torch 2.14.1 CPU, transformers 5.19.0): the MPT, ALiBi and Bloom test files pass (61 tests), and everything in tests/unit that matches mpt or alibi or bloom or generate or kv_cache or cache passes (707 passed, 10 skipped). I did not run the whole tests/unit suite for this change since it only touches the MPT attention bridge. black, isort, pycln and mypy are clean on the changed files.

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 (the related subset, see above)
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

MPTALiBiAttentionBridge._reconstruct_attention never called _update_kv_cache, so each cached decoding step built K and V from the new token alone and the model ignored the prompt. Generation with the cache differed from Hugging Face (hf-internal-testing/tiny-random-MptForCausalLM: 8 of 26 ids), while generation without it matched. Call _update_kv_cache after the heads are split, as the Bloom and joint-QKV bridges do.
Copilot AI balanced review requested due to automatic review settings October 9, 2026 14:19

Copilot AI 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.

Copilot was unable to review this pull request because the user who requested the review has reached their quota limit.

@jlarson4
jlarson4 merged commit db88af5 into TransformerLensOrg:dev Oct 9, 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