Repository navigation
fix(bridge): update the KV cache in MPT's attention reconstruction - #1874
Merged
Merged
Conversation
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.
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
With the KV cache on,
TransformerBridge.generategives different tokens from Hugging Face for MPT models, while generation without the cache is correct.MPTALiBiAttentionBridge._reconstruct_attentionnever calls_update_kv_cache, unlikeBloomAttentionBridgeand 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_cachealready reads MPT'slayer_pastkwarg (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 withposition_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
generateof the Hugging Face model (fp32, CPU):devhf-internal-testing/tiny-random-MptForCausalLM, cachedMptForCausalLMwith wide weights, cachedI 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_modulereplaces the submodules of the model it is given, so the Hugging Face reference has to be generated before bridging or on adeepcopy, which the new test does.Fixes # (no issue, found by the check above)
Type of change
Tests
TestMPTCachedGeneration::test_cached_generation_matches_hfintest_mpt_adapter.pybuilds a tiny randomMptForCausalLM, generates with an untouched copy, then with the bridge with and without the cache, and requires all three to be equal. It fails ondevand 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/unitthat matchesmpt or alibi or bloom or generate or kv_cache or cachepasses (707 passed, 10 skipped). I did not run the wholetests/unitsuite for this change since it only touches the MPT attention bridge.black,isort,pyclnandmypyare clean on the changed files.Checklist: