Fix QuantizedKVCache abort / silent corruption on non-float32 activations - #23399
Open
tritsystem wants to merge 1 commit into
Open
tritsystem wants to merge 1 commit into
tritsystem wants to merge 1 commit into
Conversation
… 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>
🔗 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 FailureAs of commit c188d6f with merge base e229c58 ( AWAITING APPROVAL - The following workflows need approval before CI can run:
NEW FAILURE - The following job has failed:
This comment was automatically generated by Dr. CI and updates every 15 minutes. |
This PR needs a
|
This branch has not been deployed
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
QuantizedKVCache._update_and_return_float_values()(inexamples/models/llama/source_transformation/custom_kv_cache.py) dequantizesits int8 cache back to a hardcoded
self.cache_fp_type(alwaystorch.float32, set once in__init__and never parameterized by thecaller's actual compute dtype), then writes the live
k_val/v_valstraight into that float32 buffer. On-device LLM inference overwhelmingly
runs attention in bfloat16 or float16, so
k_val/v_valare routinelynot float32 here, and neither update path tolerates that:
use_custom_update_cache_op=Truecallstorch.ops.llama.update_cache,a raw
memcpy(extension/llm/custom_ops/op_update_cache.cpp) keyed onvalue.element_size() == cache.element_size()(byte width), not onscalar_type(). A different-byte-size mismatch (e.g. bfloat16/float16value into this float32 cache) fails that check via
ET_CHECK_MSG, whichcalls a
FATALlog handler that aborts the whole process — not acatchable 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 getsmemcpy'das 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_typeis always float32), but it showsthe 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=Falsefalls back tok_out[:, input_pos] = k_val, which is advanced/fancy indexing (input_posis a tensor) andlowers to
index_put_. On current PyTorch this raisesRuntimeError: Index put requires the source and destination dtypes matchrather thanpromoting, unlike plain slice assignment.
Fix
Cast
k_val/v_valtok_out.dtype/v_out.dtypeonce, up front, beforeeither 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()withAffineAsymmetric+use_custom_update_cache_op=True+ a bfloat16k_val/v_val: confirmed fatal process abort (full traceback throughtorch.ops.llama.update_cache, not a catchable exception). Thenon-custom-op path raises a
RuntimeErrorfor the same input instead. Bothare present identically on unmodified
main. The existingtest_quantized_kv_cache.pyonly ever exercisestorch.float32k/v, whichis 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 exerciseQuantizedKVCache.from_float()(the class's own factory, exactly asproduction callers use it) with bfloat16/float16 k/v against both update
paths, and cross-check the result against an un-quantized reference
KVCacheat 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 4pre-existing tests in this file pass (8/8), and the sibling
test_sdpa_with_quantized_kv_cache.py/test_quantized_sdpa.pyfiles alsopass (13/13 total, no regressions).
Ran
flake8(repo's.flake8config) clean on both changed files.ufmtflags one pre-existing, unrelated import-spacing difference present
identically on unmodified
main(a localusort/ufmtversion mismatch onmy machine, not something this change introduces).
Verification limitation
Verification used the
executorch==1.5.1prebuilt wheel (which ships thecompiled
custom_ops_aot_libextension) againsttorch==2.14.0+cpu, sincebuilding 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 separateclasses from the one this PR fixes.
cc @nil-is-all