Default flash attention retains redundant BF16 KV projection history

Open
#1,029 0 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Assessment

Difficulty
4/5
Estimated time
3-5 days
Newbie friendliness
68/100
Issue type
Bug
Clarity
Clearly specified
Activity status
Active
Tech stack
cpp

Research direction

Read gemma/attention.cc lines 208-322 and gemma/kv_cache.cc lines 275-297 first to trace the projection, transpose, and cache allocations. Reproduce the default-flash Gemma 3 workload with the stated prompt and sequence settings, then inspect cache extents and peak RSS. Done means the redundant KV history is not retained while context length, BF16 rounding, and model behavior remain unchanged.

Written by the indexing model from the issue text.

Description

Summary

Default flash attention retains two BF16 representations of KV history.
The sequence-major projection buffer remains allocated and populated after
its keys and values have been transposed into the buffers attention reads.
This increases resident memory as context length grows.

Affected path

  • Backend: default --attention_impl flash.
  • Confirmed with Gemma 3 270M, 1B, and 4B text inference.
  • Measured baseline: ffc1abc05abdf11d875d25d36c8859553ddf2641.
  • The retained-history path is also present on current dev (b68def1).
  • T5 and DeepSeek use their legacy buffers differently.

What happens

  1. ComputeQKV writes BF16 projections into sequence-major kv_cache.
  2. It applies normalization and positional encoding with BF16 rounding.
  3. It transposes the results into k_cache and v_cache.
  4. FlashAttention consumes those transposed buffers.
  5. The original projection buffer still retains every layer and position.

References:

Reproduction and observed cost

Run default-flash Gemma 3 270M with a 32,736-token text prompt,
sequence capacity 32,768, prefill batch 4,096, and 16 decode tokens.
Inspect the allocated cache extents and peak process RSS during inference.

The projection buffer occupies 578 MiB including row padding.
The transposed K/V buffers occupy another 576 MiB.
Measured peak process RSS is 1,778.59 MiB on this workload.
This is active retained data, not merely unused virtual address space.

Environment: Linux, Intel i5-12400F, six pinned threads,
Release AVX2/Haswell build, no oneDNN, and approximately 15.5 GiB RAM.

Expected behavior

Projection intermediates should not retain a second complete KV history
after attention's persistent representation has been produced.
Required context, existing BF16 rounding, and model behavior must be preserved.

Dominant language
C++
Stars
7k
Forks
660
Avg merge
20h 43m
Merged PRs (30d)
33

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

More from google/gemma.cpp

All issues in google/gemma.cpp

Similar issues

More C++ issues

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.