Skip to content

Code refactor 🤖🤖🤖 - #301

Open
maxjeblick wants to merge 57 commits into
mainfrom
max/agent_refactor
Open

maxjeblick wants to merge 57 commits into
mainfrom
max/agent_refactor

Conversation

@maxjeblick

Copy link
Copy Markdown
Collaborator

PR description

This is a larger PR designed to fix several bugs.
Bugs are described in the attached HTML.
The code was written with the help of opus 5.5. and went thorugh an initial review to ensure the scope of this PR is coherent. Each fix is a separate commit. When reviewing, it can make sense to reference the corresponding commit hash.

kvpress-bugfixes.html

Checklist

Before submitting a PR, please make sure:

  • Tests are working (make test)

  • Code is formatted correctly (make style, on errors try fix with make format)

  • Copyright header is included

  • All commits are signed-off using git commit -s

  • (new press) mypress_press.py is in the presses directory

  • (new press) MyPress is in __init__.py

  • (new press) README.md is updated with a 1 liner about the new press in the Available presses section

  • (new press) New press is in the default_presses list in tests/default_presses.py

  • (new press) A docstring is provided that follows the same structure as the existing ones

…g 🤖🤖🤖

The MaxJeblick/InfiniteBench answers arrive as numpy arrays of strings and
multiple-choice answers only keep the letter, so with perfect predictions
code_run scored 0, code_debug raised an IndexError, math_find a TypeError,
math_calc an AssertionError, and passkey/number_string relied on a
deprecated ndarray-to-float conversion.

calculate_metrics now converts each label to a list, casts code_run,
math_find and math_calc labels to numbers, and rebuilds the upstream
[option_text, letter] labels of code_debug and longbook_choice_eng from the
options at the end of the question (or of the context with query_aware).
The upstream scoring functions are unchanged.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Tokenizers without a BOS token (e.g. Qwen) return None for bos_token, so models without a chat template failed with a TypeError when building the context.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Fire passes "--trust_remote_code false" (likewise --query_aware and --fp8) as the truthy string "false", which
silently enabled remote code. Coerce "true"/"false" strings (case-insensitive) for every bool field of
EvaluationConfig and raise a ValueError for any other non-bool value.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
…ing 🤖🤖🤖

DMSPress(decoding=True) keeps masks from the first answer that point past the trimmed cache, so the next question indexes out of bounds. PrefillDecodingPress without a decoding press is safe for multiple questions.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
generation_config.eos_token_id may be None or miss the tokenizer EOS token, in which case the pipeline always generated max_new_tokens tokens.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
_setup_model_pipeline overwrote model_kwargs["attn_implementation"] with flash_attention_2 whenever flash_attn
was importable, even when the user set it (evaluate_config.yaml advertises model_kwargs). Only select flash
attention when attn_implementation is unset.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
model_kwargs in _setup_model_pipeline aliased config.model_kwargs, so with
--fp8 a FineGrainedFP8Config object (and the automatic attn_implementation)
ended up in config.yaml, which yaml.dump wrote as a python/object tag that
yaml.safe_load cannot read. Copy the dict before adding entries, and save
the config with yaml.safe_dump (which also writes tuples, e.g. from
--needle_depth 10,50, as plain lists).

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
With add_v_norm=True the token scores (context length) were multiplied
by the norms of all cached values, which also cover the re-encoded copy
of the context, so the press always crashed with a size mismatch.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The haystack dataset has a single row, so fraction <= 0.5 sampled it away
and insert_needle_in_haystack then failed with KeyError: 0 (it read the
first row by label). Skip fraction sampling for this dataset with a
warning, read the first row by position, and raise a ValueError when
max_context_length leaves no room for the haystack once the needle and the
150 reserved prompt tokens are subtracted (a negative budget used to keep
almost the whole essay).

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
…copy 🤖🤖🤖

The patched forward re-encoded the context after every call, decoding
steps included, so under model.generate the answer changed even at
compression_ratio=0 and the cache ended up holding the last token only.
The re-encoded copy of the context also stayed in the cache: structured
mode cropped it when compressing, but in unstructured mode a 64-token
context left 128 entries and the question positions collided with the
copy.

Only the first prefill (empty cache) is now scored, with the forward
hooks registered for its two passes only, and the copy is cropped right
after the scoring pass. Later calls run unchanged, and the replay no
longer reuses the positions and masks of the prefill call. The scoring
uses the cache returned by the model, so calling the model without a
cache also works.

Structured results are unchanged. Unstructured results change: the
kvcompose_unstructured registry entry needs to be re-evaluated.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
… 🤖🤖🤖

kvpress' attention patch resets masked_key_indices on prefill, but
KVComposePress prefills with eager attention, which is not patched. The
masks left by a previous head-wise press (e.g. KVzipPress) then applied
to the structured cache during decoding.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
… 🤖🤖🤖

The newest tokens of a transformers QuantizedLayer are stored in full precision in layer.keys and layer.values. Decoding compression (DecodingPress, CAMPress) ignored them and silently dropped the newest tokens. Extraction now appends the residual, and all write-backs go through set_keys_and_values, which re-quantizes every token and empties the residual.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The compression ran in the context manager's finally block, so a block without a forward pass raised
AttributeError and an exception raised in the block was replaced by the one raised while compressing.
Compress only when the block completed and a context was registered.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Structured compression rebuilt the cache from the first batch element
only, silently dropping the others, and unstructured compression raised
only after both prefill passes. Raise a ValueError when the context is
registered instead.

test_kvcompose_structured_batch_size_2_does_not_collapse_on_short_contexts
relied on batch size 2 being accepted and is removed.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Unstructured compression masks the evicted keys through kvpress'
attention patch, which does not wrap eager attention, so with an eager
model the masks were silently ignored. Raise a ValueError when entering
the context manager.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
In transformers 5, eager attention is not part of ALL_ATTENTION_FUNCTIONS, so it is never patched and masked_key_indices are silently ignored. AdaKVPress, CriticalAdaKVPress, LUKVPress and DMSPress now share check_masked_key_indices_support, which raises a ValueError for eager. BasePress also initializes masked_key_indices on hooked layers (without overwriting existing masks), so DMSPress no longer fails with an AttributeError.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
KVzipPress stored kwargs["past_key_values"], which is None when the
model is called without a cache (e.g. model(inputs) as in the README).
The replay then scored a throwaway cache while the masks applied to the
cache returned by the model, and KVgradPress crashed. The cache is now
taken from the forward output, input_ids may be passed positionally,
and a clear ValueError is raised without input_ids or without a cache.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
SnapKVPress, ExpectedAttentionPress and CAMPress divided by sqrt(head_dim), which differs from the model attention when query_pre_attn_scalar != head_dim (e.g. Gemma3-27B).

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The patched forward stored the input_ids and cache of every call, so
after model.generate the replay reconstructed the last decoded token
only and the cache was truncated to 1 token (KVgradPress scored that
token only). Only the first prefill (empty cache) is now captured and
later calls run unchanged. Tokens decoded after the prefill are cropped
before the replay, so the scores match the ones of a plain prefill.
This applies to KVgradPress and RestoreKVPress too.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
…on 🤖🤖🤖

Future positions started at hidden_states.shape[1], which is only the next position during a plain prefill. Under DecodingPress (buffered hidden states), ChunkPress or BlockPress they were placed too early. The keys are now rotated back by the difference to the next position (from position_ids, falling back to cache_position), which is equivalent for RoPE and leaves the get_query_statistics/apply_avg_rope signatures unchanged for ExpectedAttentionStatsPress. Prefill scores are unchanged.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The replayed chunks land in the residual buffer of the quantized
layers, so KVzip scored the wrong keys (up to 0.93 off with an identity
quantizer), and KVComposePress cannot crop or rebuild quantized layers.
Raise a ValueError on the prefill instead, for KVzipPress and its
subclasses (KVgradPress, RestoreKVPress) and for KVComposePress, and
drop the now unreachable QuantizedCache branch of KVzipPress's hook.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Padded keys were masked with -1e-9 instead of a large negative value, and the zero-padded queries of the last chunk spread a uniform softmax over all keys of that chunk. Padded keys are now masked with the float32 minimum and padded query rows are zeroed after the softmax, so the column sums total the number of real queries (#295).

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Inner presses that load data from the model (QFilterPress, KVzapPress, DuoAttentionPress, ...) failed because
post_init_from_model was not delegated to them.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Sinks were padded with the maximum score, so they tied with the top token and could be pruned.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Qwen2, Qwen2.5 and Phi3 configs have no head_dim attribute, so CriticalKVPress and CriticalAdaKVPress failed on these models (#213).

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Inner presses that load data from the model were never initialized.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The published gates store their weights in bfloat16 and are only moved
to the model device, so float16 and float32 models crashed with "mat1
and mat2 must have the same dtype". The test mock builds its gates in
the model dtype, which hid the issue. bfloat16 models are unaffected.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Gemma3Model has no layers attribute (they live in its language_model),
and the skipped sliding window layers left None scores that could not
be stacked, so FastKVzipPress failed on Gemma3 although it is listed as
supported. Look the layers up as BasePress.__call__ does and only stack
and compress the scored layers. A block without forward pass no longer
raises either.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Each layer scores its KV pairs on its own device, so torch.stack failed
when the layers were spread over several devices (device_map).

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Qwen2 and Qwen2.5 configs have no head_dim attribute, so both the stats
lookup and the statistics script raised an AttributeError. Fall back to
hidden_size // num_attention_heads.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
…e 🤖🤖🤖

The query statistics are loaded on model.device, so scoring failed for
the layers placed on other devices (device_map). Move the statistics of
a layer to the device of its hidden states when they are used.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The rotary embedding is shared by all layers, and with a device_map it
can return cos and sin on another device than the layer being scored,
so building the averaged RoPE matrix failed.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The averaged RoPE rotation used the positions following
hidden_states.shape[1], the number of tokens of the forward pass, rather
than the position following the last token (cache_position[-1] + 1).
Both match for a plain prefill, but not when the hidden states follow
cached tokens, e.g. with chunked prefilling or decoding presses.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Outside eager attention, DropKVPress recomputes the window attention
with 1/sqrt(head_dim) instead of the attention scaling of the model.
Both match for all supported models except Gemma3-27B, whose scaling is
query_pre_attn_scalar**-0.5.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The published Gemma3 gates only cover the full attention layers (8 gates
for the 48 layers of gemma-3-12b-it), in order, as in the original
implementation (layer_id_to_static_id). Indexing them by layer index
gave layer 5 the gate of the sixth full attention layer and raised an
IndexError from layer 11. Gates are now looked up by the position of the
layer among the scored layers, which is the layer index for the other
models.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
…ress 🤖🤖🤖

Window-based scorers received the buffered hidden states with the RoPE embeddings of the current step only, so every buffered query was rotated at the current position. The embeddings are now buffered with the hidden states (also in CAMPress) and passed in a kwargs copy. Attention weights are no longer passed to the base press as they only cover the current query.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The buffer and step counts of a layer survived a new prefill in the same press context (e.g. two generate calls), so the next compression mixed hidden states of the previous sequence. They are now reset per layer on prefill, as in CAMPress.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
CAMPress calls base_press.score, which AdaKVPress does not implement, so the cam_adakv_snapkv registry entry could not run. It is removed from the registry.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The two independent topk calls could select a token as both kept and merged, or neither, when scores are tied. Without ties the selected tokens are unchanged.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
In bfloat16 the running sum of small attention weights stalls (2000 steps of 2e-3 sum to 1.0 instead of 4.0). Merge contributions are cast back to the cache dtype.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The target size was derived from position_ids, which models fill from cache_position when they are not passed. cache_position counts the compressed cache, so the target shrank geometrically (67, 34, 19, ... instead of ~44) and the documented NotImplementedError was unreachable. The press now counts the prefill and decoding tokens per layer (reset on prefill), starting from the position of the first decoding step when the prefill ran outside of the press context. The tests of the removed position_ids logic are replaced by tests/presses/test_compression_ratio_decoding_press.py.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Only the inner forward hooks were called, never the inner __call__, so presses relying on it failed (e.g.
FinchPress: window_size must be provided). Both inner contexts are now entered with an ExitStack, which is
enough since BasePress hooks only compress during prefilling and DecodingPress hooks only during decoding.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The state was only reset when layer 0 prefilled, which never happens on Gemma3 where layer 0 is a skipped sliding window layer.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
keys[:, :, -n_recent + n_last:] becomes keys[:, :, 0:] when n_recent == n_last (the config of tests/default_presses.py), keeping the whole cache plus duplicated initial tokens. The recent window is now sliced from k_len - max(0, n_recent - n_last), and the reported ratio is computed from the kept length instead of assuming n_last == 1.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The float32 (batch, heads, n_evict, n_kept) similarity matrix needed about 8.6 GB per layer at 32k context. It
is now computed for chunk_size=1024 evicted tokens at a time, which gives the same nearest survivors
(identical for up to chunk_size evicted tokens, up to float32 rounding of the matrix product otherwise).

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
Chunks covered the whole sequence, so a remainder chunk made only of window (question) tokens evicted some of them. Chunks now only cover the context and the window is always kept.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
A sample without delimiter silently reused the window size of the previous sample. It now fails with 'window_size must be provided'.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The attention weights multiplied by the number of non-zero weights overflow float16 for long contexts (0.9 * 80000 = inf). The formula is unchanged.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
…dings 🤖🤖🤖

resize_token_embeddings(len(tokenizer)) shrank the embeddings of models whose vocabulary is padded (Qwen2.5: 151,936 -> 151,666 rows) on the shared model, so they are only resized when the delimiter id does not fit. model.model.embed_tokens does not exist on Gemma3ForConditionalGeneration, model.get_input_embeddings() is used instead.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
DecodingPress and CAMPress remove tokens from the cache, which shifts the positions of keys masked by a head-wise press applied during prefilling (e.g. PrefillDecodingPress(prefilling_press=AdaKVPress(...))). This raised an IndexError in attention_patch or masked the wrong keys; it now raises a clear ValueError.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
The randomly initialized model could stop generating early, changing the number of decoding steps.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
transformers 5 returns a BatchEncoding from apply_chat_template, so the script crashed on the first question.
Request return_dict=True, call model.generate(**inputs) and slice the output with the input_ids length.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
With an empty mask (e.g. a head-wise press at compression_ratio=0), search_hyperplane ran on every decoding step for nothing and could fail to find a hyperplane on some models.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
In transformers 5.2, load_adapter keeps the adapter weights on the CPU unless a device map is given, so RestoreKVPress failed with a device mismatch on single-device CUDA models.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
make style fails on main because black 24 formats these three files differently. No code changes.

Signed-off-by: Maximilian Jeblick <maximilianjeblick@gmail.com>
@copy-pr-bot

copy-pr-bot Bot commented Oct 6, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@maxjeblick

Copy link
Copy Markdown
Collaborator Author

/ok to test ef33440

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant