Bug
KVPressTextGenerationPipeline's removeanswer_from_cache() crashes when used with QuantizedCache(backend="hqq", ...). It runs after every generated answer to trim the question+answer tokens back out of the cache, and unconditionally does:
cache.layers[layer_idx]._quantized_keys = cache.layers[layer_idx]._quantized_keys[:, :, :sequence_length]
This assumes ._quantized_keys is a single sliceable tensor, which holds for the quanto backend (WeightQBitsTensor, still shaped [batch, kv_heads, seq_len, head_dim]), but not for hqq: HQQQuantizedLayer._quantize() (transformers/cache_utils.py) returns a (qtensor, meta) tuple, and slicing a plain Python tuple with a 3-axis index raises:
TypeError: tuple indices must be integers or slices, not tuple
qtensor itself has also been reshaped/bit-packed into a flat [bytes_per_group, num_groups] layout (e.g. [32, 176] for a 22-token, 8-head, 128-dim key tensor at 4 bits) where seq_len is no longer a distinct axis, so even a tuple-aware slice couldn't correspond to "the first N tokens" - the fix needs to dequantize() back to the real shape first, slice there, then quantize() again.
To Reproduce
import torch
from transformers import QuantizedCache, pipeline
import kvpress
pipe = pipeline("kv-press-text-generation", model="meta-llama/Llama-3.2-3B-Instruct", device="cuda", dtype=torch.float16)
cache = QuantizedCache(backend="hqq", config=pipe.model.config, nbits=4)
pipe("The Eiffel Tower is located in Paris, France.", question="\nWhere is the Eiffel Tower located?", cache=cache, max_new_tokens=20)
Repository version
7331c23
Bug
KVPressTextGenerationPipeline's removeanswer_from_cache() crashes when used with QuantizedCache(backend="hqq", ...). It runs after every generated answer to trim the question+answer tokens back out of the cache, and unconditionally does:
This assumes ._quantized_keys is a single sliceable tensor, which holds for the quanto backend (WeightQBitsTensor, still shaped [batch, kv_heads, seq_len, head_dim]), but not for hqq: HQQQuantizedLayer._quantize() (transformers/cache_utils.py) returns a (qtensor, meta) tuple, and slicing a plain Python tuple with a 3-axis index raises:
qtensor itself has also been reshaped/bit-packed into a flat [bytes_per_group, num_groups] layout (e.g. [32, 176] for a 22-token, 8-head, 128-dim key tensor at 4 bits) where seq_len is no longer a distinct axis, so even a tuple-aware slice couldn't correspond to "the first N tokens" - the fix needs to dequantize() back to the real shape first, slice there, then quantize() again.
To Reproduce
Repository version
7331c23