From c188d6f290e02f388f511cdcfb89b5e6774e605a Mon Sep 17 00:00:00 2001 From: tritsystem Date: Fri, 2 Oct 2026 23:00:01 -0700 Subject: [PATCH] Fix QuantizedKVCache process abort / silent corruption on non-float32 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 --- .../source_transformation/custom_kv_cache.py | 20 +++++ .../test_quantized_kv_cache.py | 77 +++++++++++++++++++ 2 files changed, 97 insertions(+) diff --git a/examples/models/llama/source_transformation/custom_kv_cache.py b/examples/models/llama/source_transformation/custom_kv_cache.py index 71cfe33d753..598a4fc84ac 100644 --- a/examples/models/llama/source_transformation/custom_kv_cache.py +++ b/examples/models/llama/source_transformation/custom_kv_cache.py @@ -184,6 +184,26 @@ def _update_and_return_float_values(self, input_pos, k_val, v_val, indices=None) self.cache_fp_type, ) + # k_out/v_out are always self.cache_fp_type (float32), but k_val/v_val + # carry whatever dtype the live model activations use (e.g. bfloat16 + # or float16 for on-device inference). Neither update path below + # tolerates that mismatch: + # - torch.ops.llama.update_cache (use_custom_update_cache_op=True) is + # a raw memcpy-based op keyed on byte size, not dtype: a + # same-byte-size mismatch (e.g. bfloat16 into a float16-sized slot) + # silently reinterprets bits and corrupts the cache, while a + # different-byte-size mismatch (e.g. bfloat16/float16 into this + # float32 cache) hits a fatal assertion and aborts the process. + # - the plain-assignment fallback below (`k_out[:, input_pos] = k_val`) + # is advanced/fancy indexing (input_pos is a tensor), which lowers + # to index_put_ and raises `RuntimeError: Index put requires the + # source and destination dtypes match` instead of promoting, unlike + # plain slice assignment. + # Cast once, up front, so both paths get a value already in the + # cache's float dtype. + k_val = k_val.to(k_out.dtype) + v_val = v_val.to(v_out.dtype) + # When returning float values we just use the last value # instead of dequantized value. start_pos = input_pos[0].item() diff --git a/examples/models/llama/source_transformation/test_quantized_kv_cache.py b/examples/models/llama/source_transformation/test_quantized_kv_cache.py index 07c8e1bf9a0..8b83c61db5c 100644 --- a/examples/models/llama/source_transformation/test_quantized_kv_cache.py +++ b/examples/models/llama/source_transformation/test_quantized_kv_cache.py @@ -124,3 +124,80 @@ def test_simple_update_fetch_dynamic_shape_use_custom_op(self): self._test_simple_update_fetch( is_dynamic_shape=True, use_custom_update_cache_op=True ) + + def _test_update_with_non_float32_activation_dtype( + self, dtype, use_custom_update_cache_op + ): + """QuantizedKVCache stores int8 data internally but always dequantizes + back to a hardcoded `cache_fp_type` (float32), regardless of the dtype + of the live k/v activations it's updated with (`cache_fp_type` is not + parameterized by, and does not track, the model's actual compute + dtype). When those activations are bfloat16/float16 (the common case + for on-device LLM inference, which is the whole point of + ExecuTorch), `QuantizedKVCache.update()` must not crash and must + return values consistent with the un-quantized reference cache, + exactly as it already does for float32 activations in + `_test_simple_update_fetch` above. + """ + max_batch_size, max_context_len, n_kv_heads, head_dim = 1, 5, 8, 17 + kv_cache = KVCache( + max_batch_size, + max_context_len, + n_kv_heads, + head_dim, + False, + dtype=dtype, + ) + quantized_kv_cache = QuantizedKVCache.from_float( + kv_cache, + QuantizedCacheType.AffineAsymmetric, + use_custom_update_cache_op, + ) + + input_pos = torch.tensor([0, 1, 2]) + shape = (max_batch_size, n_kv_heads, input_pos.size(0), head_dim) + k = torch.rand(shape, dtype=dtype) + v = torch.rand(shape, dtype=dtype) + + updated_dequantized_k_cache, updated_dequantized_v_cache = ( + quantized_kv_cache.update(input_pos, k, v) + ) + + # Reference: un-quantized cache update, same dtype. + updated_k_cache, updated_v_cache = kv_cache.update(input_pos, k, v) + + def index(t, positions): + return t[:, :, positions, :] + + torch.testing.assert_close( + index(updated_k_cache, input_pos).float(), + index(updated_dequantized_k_cache, input_pos).float(), + rtol=1e-02, + atol=1e-02, + ) + torch.testing.assert_close( + index(updated_v_cache, input_pos).float(), + index(updated_dequantized_v_cache, input_pos).float(), + rtol=1e-02, + atol=1e-02, + ) + + def test_update_bfloat16_activation(self): + self._test_update_with_non_float32_activation_dtype( + torch.bfloat16, use_custom_update_cache_op=False + ) + + def test_update_bfloat16_activation_use_custom_op(self): + self._test_update_with_non_float32_activation_dtype( + torch.bfloat16, use_custom_update_cache_op=True + ) + + def test_update_float16_activation(self): + self._test_update_with_non_float32_activation_dtype( + torch.float16, use_custom_update_cache_op=False + ) + + def test_update_float16_activation_use_custom_op(self): + self._test_update_with_non_float32_activation_dtype( + torch.float16, use_custom_update_cache_op=True + )