Fix Gemma 3/3n KV cache to honor per-batch end_index - #765
Open
Ayush7614 wants to merge 1 commit into
Open
Conversation
Gemma 4 already updates cache slots with per-example indices, but Gemma 3 and Gemma 3n still used end_index[0], so divergent fill lengths in a batch silently corrupted every row. Align both paths with the Gemma 4 scatter update and add regression tests.
|
Thanks for your pull request! It looks like this may be your first contribution to a Google open source project. Before we can look at your pull request, you'll need to sign a Contributor License Agreement (CLA). View this failed invocation of the CLA check for more information. For the most up to date status, view the checks section at the bottom of the pull request. |
Ayush7614
force-pushed
the
feat/per-batch-kv-cache-gemma3
branch
from
July 28, 2026 21:35
d14faf4 to
b811f90
Compare
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.
Summary
end_index[0], so every batch row was written at the same offset whenever fill lengths diverged.batch_indices+ modulo indices) intogm/nn/_modules.pyandgm/nn/gemma3n/_modules.py.% cache_size(matching Gemma 4).end_indexvalues and assert each row updates its own slot (and that the old shared-slot corruption no longer happens).This is independent of existing rolling-cache / OOM issues and does not claim any open issue.
Test plan
end_index=[1,4]: row0 writes slot1, row1 writes slot4; row1 slot1 unchangedend_indexstill updates both rows at the shared offsetpytest gemma/gm/nn/_modules_test.py::test_attention_cache_uses_per_batch_end_index gemma/gm/nn/gemma3n/_modules_test.py::test_attention_cache_uses_per_batch_end_index