Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
20 changes: 20 additions & 0 deletions examples/models/llama/source_transformation/custom_kv_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Loading