Repository navigation
Fix _remove_answer_from_cache crash with QuantizedCache(backend='hqq') - #290
Conversation
Signed-off-by: bonginn <caco.sc11@nycu.edu.tw>
|
Hi @bonginn , thanks a lot for spotting the bug, and opening a PR. I reviewed the PR; the current implementation doesn't reset the hqq cache properly. Transformers library also faces this issue (reusable quantized cache), and proposes to copy over the whole cache before generation, then restoring it. We opt to copy HF's pattern to support quantized cache, althtough it temporarily copies the kv-cache. A proper fix would require an implementation within the quantized cache class (something like We are happy if you are interested in adapting your PR, below you find a detailed assesment by an agent that may help. Assessment: the PR removes the
The new test asks a single short question, so it covers none of these cases. Suggested fix: create a copy of the layer state that is taken before each answer and restored afterwards. This follows the prefix caching pattern in the transformers docs, which deep-copies the prefilled cache before each generation. A shallow copy is enough here, because Sketch of implementation, needs to be implemented: # in _forward, before generate_answer
layer_states = [vars(layer).copy() for layer in cache.layers] if isinstance(cache, QuantizedCache) else None
...
self._remove_answer_from_cache(cache, cache_seq_lengths, layer_states)
# in _remove_answer_from_cache, replacing the QuantizedCache branch
if isinstance(cache, QuantizedCache):
for layer, state in zip(cache.layers, layer_states):
vars(layer).update(state)
return
|
A QuantizedLayer flushes its full-precision residual into the quantized storage once it reaches residual_length, so the question and answer can't be sliced off. Slicing also broke both backends with more than one question: quanto's sliced WeightQBitsTensor is no longer quantized and dequantizes to float32, and hqq's cumulative_length was never reset. Save each layer's state before an answer and restore it afterwards, following the prefix caching pattern in the transformers docs. Add a two-question test with residual_length=4 for quanto and hqq, and add the missing config argument to the README example. Signed-off-by: bonginn <caco.sc11@nycu.edu.tw>
|
Thanks a lot for the detailed review @maxjeblick! I've updated the PR following your suggestion: the layer states are saved before each answer and restored afterwards instead of slicing the quantized cache. I added the two-question test with |
maxjeblick
left a comment
There was a problem hiding this comment.
LGTM, thanks a lot for this fix!
|
/ok to test aef39ea |
PR description
Fixes #289. Following @maxjeblick's review,
_remove_answer_from_cacheno longer slices aQuantizedCache: it saves each layer's state before an answer and restores it afterwards, as in the prefix caching pattern of the transformers docs.Slicing can't restore a quantized cache:
QuantizedLayerflushes its full-precision residual into the quantized storage once it reachesresidual_length, so the question and answer end up inside the quantized tensors. With more than one question, slicing failed on both backends:hqq:_quantized_keysis a(qtensor, meta)tuple (the originalTypeError), andcumulative_lengthwas never reset, so the next question hit an attention mask size mismatch.quanto: slicing aWeightQBitsTensorreturns a plain tensor that dequantizes to float32 (Expected query, key, and value to have the same dtype).A shallow copy of
vars(layer)is enough becauseQuantizedLayer.updateassigns new tensors instead of writing in place. The quantized branch returns before theDynamicCacheslicing, which is unchanged.Tests: added
test_pipeline_with_quantized_cache_multiple_questions(quanto and hqq,residual_length=4so the residual flushes while answering). It checks that the cache is back to the compressed context length and that each answer matches a fresh single-question call. It fails onmainfor both backends and passes with this change. Also fixed the README example, which was missing the requiredconfigargument.Checklist
Before submitting a PR, please make sure:
make test)make style, on errors try fix withmake format)git commit -s