Skip to content

Fix QuantizedKVCache abort / silent corruption on non-float32 activations - #23399

Open
tritsystem wants to merge 1 commit into
pytorch:mainfrom
tritsystem:fix-quantized-kv-cache-non-fp32-activation-crash
Open

tritsystem wants to merge 1 commit into
pytorch:mainfrom
tritsystem:fix-quantized-kv-cache-non-fp32-activation-crash

Conversation

@tritsystem

@tritsystem tritsystem commented Oct 3, 2026 •

Copy link
Copy Markdown

Summary

QuantizedKVCache._update_and_return_float_values() (in
examples/models/llama/source_transformation/custom_kv_cache.py) dequantizes
its int8 cache back to a hardcoded self.cache_fp_type (always
torch.float32, set once in __init__ and never parameterized by the
caller's actual compute dtype), then writes the live k_val/v_val
straight into that float32 buffer. On-device LLM inference overwhelmingly
runs attention in bfloat16 or float16, so k_val/v_val are routinely
not float32 here, and neither update path tolerates that:

  • use_custom_update_cache_op=True calls torch.ops.llama.update_cache,
    a raw memcpy (extension/llm/custom_ops/op_update_cache.cpp) keyed on
    value.element_size() == cache.element_size() (byte width), not on
    scalar_type(). A different-byte-size mismatch (e.g. bfloat16/float16
    value into this float32 cache) fails that check via ET_CHECK_MSG, which
    calls a FATAL log handler that aborts the whole process — not a
    catchable Python exception. I additionally verified, directly against the
    op in isolation, that a same-byte-size mismatch (e.g. a bfloat16 value
    into a float16 cache) passes the element_size() check and gets memcpy'd
    as raw bits, silently reinterpreting the value under the wrong dtype —
    producing a numerically wrong cache entry with no error at all. This repo
    isn't exposed to that second case today (its only non-float32 cache dtype
    for this path is... none — cache_fp_type is always float32), but it shows
    the op's own contract is byte-size-based, not dtype-based, which is worth
    knowing about independent of this specific class.
  • use_custom_update_cache_op=False falls back to k_out[:, input_pos] = k_val, which is advanced/fancy indexing (input_pos is a tensor) and
    lowers to index_put_. On current PyTorch this raises RuntimeError: Index put requires the source and destination dtypes match rather than
    promoting, unlike plain slice assignment.

Fix

Cast k_val/v_val to k_out.dtype/v_out.dtype once, up front, before
either update path runs. This matches the semantics the fallback path
already (attempts to) provide — the latest activation value overwrites the
cache at the cache's float precision — and fixes both paths uniformly rather
than only the custom-op one.

Test plan

Reproduced the crash directly against QuantizedKVCache.update() with
AffineAsymmetric + use_custom_update_cache_op=True + a bfloat16 k_val/
v_val: confirmed fatal process abort (full traceback through
torch.ops.llama.update_cache, not a catchable exception). The
non-custom-op path raises a RuntimeError for the same input instead. Both
are present identically on unmodified main. The existing
test_quantized_kv_cache.py only ever exercises torch.float32 k/v, which
is why neither path was caught before.

Added 4 regression tests to test_quantized_kv_cache.py
(test_update_{bfloat16,float16}_activation[_use_custom_op]) that exercise
QuantizedKVCache.from_float() (the class's own factory, exactly as
production callers use it) with bfloat16/float16 k/v against both update
paths, and cross-check the result against an un-quantized reference
KVCache at the same dtype.

Red/green confirmed: reverting only the source fix causes the custom-op
variants to abort the whole pytest process, and the non-custom-op variants
to raise RuntimeError. With the fix, all 4 new tests plus the 4
pre-existing tests in this file pass (8/8), and the sibling
test_sdpa_with_quantized_kv_cache.py / test_quantized_sdpa.py files also
pass (13/13 total, no regressions).

Ran flake8 (repo's .flake8 config) clean on both changed files. ufmt
flags one pre-existing, unrelated import-spacing difference present
identically on unmodified main (a local usort/ufmt version mismatch on
my machine, not something this change introduces).

Verification limitation

Verification used the executorch==1.5.1 prebuilt wheel (which ships the
compiled custom_ops_aot_lib extension) against torch==2.14.0+cpu, since
building from source was out of scope for my environment. I did not test the
DA8W4/CUDA/XPU/NPU-specific KV cache variants in this file, or
StaticQuantizedKVCache/CalibratedQuantizedKVCache, which are separate
classes from the one this PR fixes.

cc @nil-is-all

… activations

QuantizedKVCache._update_and_return_float_values() dequantizes its int8
cache back to a hardcoded self.cache_fp_type (always torch.float32,
set once in __init__ and never parameterized by the caller's actual
compute dtype), then writes the live k_val/v_val straight into that
float32 buffer. On-device LLM inference overwhelmingly runs attention
in bfloat16 or float16, so k_val/v_val are routinely NOT float32 here,
and neither the custom-op nor the fallback update path tolerates that:

- use_custom_update_cache_op=True calls torch.ops.llama.update_cache,
  a raw memcpy keyed on value.element_size() == cache.element_size()
  (byte width), not on scalar_type(). A different-byte-size mismatch
  (e.g. bfloat16/float16 value into this float32 cache) fails that
  check via ET_CHECK_MSG, which calls a FATAL log handler that aborts
  the whole process -- not a catchable Python exception. Verified
  directly against the op (extension/llm/custom_ops/op_update_cache.cpp)
  with a same-byte-size pair (bfloat16 value into a float16 cache):
  element_size() matches, the check passes, and the raw bits are
  memcpy'd and silently reinterpreted under the wrong dtype, producing
  a numerically wrong cache entry with no error at all.
- use_custom_update_cache_op=False falls back to `k_out[:, input_pos]
  = k_val`, which is advanced/fancy indexing (input_pos is a tensor)
  and lowers to index_put_. On current PyTorch this raises
  "RuntimeError: Index put requires the source and destination dtypes
  match" rather than promoting, unlike plain slice assignment.

Reproduced the crash directly against QuantizedKVCache.update() with
AffineAsymmetric + use_custom_update_cache_op=True + a bfloat16 k_val/
v_val: fatal process abort, confirmed to kill the whole interpreter
(not an exception). The non-custom-op path raises a RuntimeError for
the same input instead. Both are present identically on unmodified
main; checked existing tests (test_quantized_kv_cache.py) only ever
exercise torch.float32 k/v, which is why neither path was caught.

Fix: cast k_val/v_val to k_out.dtype/v_out.dtype once, up front, before
either update path runs. This matches the semantics the fallback path
already (attempts to) provide -- the latest activation value overwrites
the cache at the cache's float precision -- and fixes both paths
uniformly rather than only the custom-op one.

Added 4 regression tests to test_quantized_kv_cache.py
(test_update_{bfloat16,float16}_activation[_use_custom_op]) that
exercise QuantizedKVCache.from_float() (the class's own factory,
exactly as production callers use it) with bfloat16/float16 k/v
against both update paths, and cross-check the result against an
un-quantized reference KVCache at the same dtype.

Red/green confirmed: reverting only the source fix causes the custom-op
variants to abort the whole pytest process (verified directly, with a
full traceback through torch.ops.llama.update_cache) and the
non-custom-op variants to raise RuntimeError. With the fix, all 4 new
tests plus the 4 pre-existing tests in this file pass (8/8), as do the
sibling test_sdpa_with_quantized_kv_cache.py and test_quantized_sdpa.py
files (13/13 total, no regressions). Ran flake8 (repo's .flake8 config)
clean on both changed files; ufmt flags one pre-existing, unrelated
import-spacing difference present identically on unmodified main
(a local usort/ufmt version mismatch, not something this change
introduces).

Verification used the executorch 1.5.1 prebuilt wheel (which ships the
compiled custom_ops_aot_lib extension) against torch==2.14.0+cpu, since
building from source was out of scope for this environment. Did not
test the DA8W4/CUDA/XPU/NPU-specific KV cache variants in this file, or
StaticQuantizedKVCache / CalibratedQuantizedKVCache, which were out of
scope for this specific crash.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
@pytorch-bot

pytorch-bot Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

🔗 Helpful Links

🧪 See artifacts and rendered test results at hud.pytorch.org/pr/pytorch/executorch/23399

Note: Links to docs will display an error until the docs builds have been completed.

❌ 20 Awaiting Approval, 1 New Failure

As of commit c188d6f with merge base e229c58 (image):

AWAITING APPROVAL - The following workflows need approval before CI can run:

NEW FAILURE - The following job has failed:

  • Cadence Build & Test / Resolve CI docker image / resolve (gh)
    ##[error]Refusing to check out fork pull request code from a 'pull_request_target' workflow. This workflow runs with the base repository's GITHUB_TOKEN, secrets, default-branch cache scope, and runner access. Fetching and executing a fork's code in that trusted context commonly leads to "pwn request" vulnerabilities. To opt in, review the risks at https://gh.io/securely-using-pull_request_target and set 'allow-unsafe-pr-checkout: true' on the actions/checkout step.

This comment was automatically generated by Dr. CI and updates every 15 minutes.

@meta-cla meta-cla Bot added the CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. label Oct 3, 2026
@github-actions

github-actions Bot commented Oct 3, 2026

Copy link
Copy Markdown

This PR needs a release notes: label

If your change should be included in the release notes (i.e. would users of this library care about this change?), please use a label starting with release notes:. This helps us keep track and include your important work in the next release notes.

To add a label, you can comment to pytorchbot, for example
@pytorchbot label "release notes: none"

For more information, see
https://github.com/pytorch/pytorch/wiki/PyTorch-AutoLabel-Bot#why-categorize-for-release-notes-and-how-does-it-work.

@executorch-triage executorch-triage Bot added the community: contribution PRs coming from community (excluding hardware partners) label Oct 3, 2026

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

CLA Signed This label is managed by the Facebook bot. Authors need to sign the CLA before a PR can be reviewed. community: contribution PRs coming from community (excluding hardware partners)

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant