Skip to content

vLLM 0.19 V1: prefill seems to bypass impl.forward monkey-patch (decode-only); notes from a fork attempt #6

Description

@Tonoken3

Hi — first off, thanks for releasing kvtc as OSS. We've been trying it out at Lna-Lab on a Blackwell box and ran into something that might be worth flagging upstream. Hopefully the trace/notes below are useful; happy to be told this is wrong and we missed something obvious.

What we observed

When we ran kvtc against vLLM 0.19.1rc1 on CUDA, vllm_backend.hook_model(...) only seems to intercept decode batches — not prefill. So state.capture(...) never sees the prefill K/V, and decode_request then has nothing in sequence.window to attend over. The shapes are all correct end-to-end, but the generated text diverges from the baseline starting at the second decode token.

Worth saying upfront: the design of hook_model is sound — this looks like a vLLM-side dispatch change, not a kvtc bug. We're flagging it here because anyone else trying kvtc on vLLM 0.19+ will probably hit the same thing.

Environment

  • vLLM 0.19.1rc1.dev310+g8ad6ff003
  • PyTorch 2.10
  • Qwen2.5-3B-Instruct (also tested calibration on Qwen3.6-27B hybrid)
  • 1×RTX PRO 6000 Blackwell (SM120), enforce_eager=True, enable_prefix_caching=False, async_scheduling=False

Reproduction

VLLM_ALLOW_INSECURE_SERIALIZATION=1 CUDA_VISIBLE_DEVICES=0 \
  python proof.py --mode worker --backend kvtc --auto-activate \
  --model Qwen/Qwen2.5-3B-Instruct \
  --max-tokens 64 --gpu-memory-utilization 0.55 --max-model-len 4096 \
  --enforce-eager --calibration-path kvtc_qwen25_calib.pt
  • auto_activate=False (passthrough, capture-only): correct output, ~91 tok/s
  • auto_activate=True (KVTC decode active): the very first generated token is correct (it's sampled from prefill logits), but every subsequent token attends over an empty sequence.window and the output diverges.

Where the trace points

We instrumented PatchedForward.__call__ and a torch register_forward_hook on the Attention nn.Module itself (not on impl). The Attention.forward hook does fire for prefill — but PatchedForward doesn't:

[Attn.forward HOOK] inputs=[(2, 2048),  (2, 256),  (2, 256)]    # warmup prefill
[Attn.forward HOOK] inputs=[(61, 2048), (61, 256), (61, 256)]   # real prefill ← never reaches PatchedForward
[PatchedForward]    q=(1, 16, 128) k=(1, 2, 128) ← only decodes show up here

The dispatch goes:

  • Attention.forward() → torch.ops.vllm.unified_attention_with_output(...) → attn_layer.impl.forward(...)

…and on CUDA (opaque_attention_op() == True ⇒ use_direct_call = False), it looks like the prefill batch goes through a fast-path that captures the original impl.forward reference rather than re-resolving the attribute, so a runtime monkey-patch isn't honoured. We didn't manage to pin down the exact line in vLLM where prefill and decode dispatch differently — would welcome any pointers if you've looked at this area before.

Capture-side workaround we tried

To get past the bypass for testing, we wrapped Attention.forward directly (instead of impl.forward) for the capture side. The existing PatchedForward decode path is left in place; a state flag prevents double-counting decode tokens (which DO reach PatchedForward):

def _wrap_attention_forward_for_capture(layer, state, handle):
    from vllm.forward_context import get_forward_context
    original_forward = layer.forward
    state._skip_patched_capture = True

    def patched_attention_forward(query, key, value, output_shape=None):
        try:
            attn_metadata = get_forward_context().attn_metadata
            if isinstance(attn_metadata, dict):
                attn_metadata = attn_metadata.get(layer.layer_name, None)
        except Exception:
            attn_metadata = None
        num_tokens = query.shape[0]
        k3d = key.view(-1, state.num_kv_heads, state.head_dim) if (key is not None and key.dim() == 2) else key
        v3d = value.view(-1, state.num_kv_heads, state.head_dim) if (value is not None and value.dim() == 2) else value
        if k3d is not None and v3d is not None:
            spans = extract_request_spans(attn_metadata, num_tokens, device=k3d.device)
            state.capture(k3d, v3d, spans)
        return original_forward(query, key, value, output_shape=output_shape)

    layer.forward = patched_attention_forward

After this, sequence.window correctly grows by prefill_len after the first decode call (window=62 after a 61-token prefill + 1 decode token, in our trace), and auto_activate triggers as expected.

This is just one approach — there's probably a cleaner integration that avoids the second wrapper layer (e.g. capturing at the unified_attention_with_output op-registration layer, or moving the capture into a register_forward_hook). Would defer to whatever shape you'd prefer.

Caveat — there's a residual decode-quality issue we haven't sorted yet

Even with the wrapper, decode_request reads from a populated sequence.window (verified via instrumentation: sinks=0 window=62 middle=0 on the first decode), but the reconstructed attention output still diverges from baseline. We've ruled out:

  • PCA / quantization (same divergence with window_tokens >= seq_len, i.e. raw-only path)
  • Triton (same divergence with --no-triton, i.e. pure-torch dense_attention_state)
  • RoPE base mismatch (byte-identical output with rope_theta=10000 vs rope_theta=1000000)

Most likely a K/V layout / contiguity / dtype thing on our end — possibly we're capturing the K in the wrong shape relative to what dense_attention_state expects under vLLM 0.19's output= / output_scale= contract. Still digging; will follow up here once we know more (or open a separate issue if it turns out to be its own thing rather than user error).

A couple of small unrelated changes we made downstream while debugging

These are independent of the dispatch question above — happen to be small enough to mention. Not asking you to take them; just flagging in case any are useful, and so the link to our fork doesn't look mysterious:

  1. relative imports in src/ — gpu_ops.py, fused_ops.py, adaptive_budget.py switched from from common import … to from .common import …. Lets us import the package as kvtc.* without sys.path hacks (might just be us; happy to drop if it's not the layout you intend).
  2. calibrate.py — handle hybrid attention models — modern past_key_values.layers can have layers without KV (layer.keys is None, e.g. linear/SSM/Mamba mixed in). Skipping those + remapping to a contiguous index over the full-attention layers got Qwen3.6-27B (16 of 64 layers carry KV) through calibration cleanly.
  3. vLLM 0.19+ child-process integration — hook_engine / free_kv_cache_engine route through llm.llm_engine.apply_model(func) (with VLLM_ALLOW_INSECURE_SERIALIZATION=1) so the hook lands inside the EngineCore worker process, plus the output_scale / output_block_scale kwarg pop in the decode signature. This is the part most coupled to the bypass issue above.

If any of these would be welcome as a PR (one combined or three small), happy to send whichever shape works best for you — and equally happy to leave them as fork-only if the design doesn't match where you'd like the project to go.

Fork (master): https://github.com/Shinka-Man/kvtc

  • 71932cf — relative imports
  • a0e9e7b — hybrid attention calibration
  • a04003a — vLLM 0.19+ integration + V1 prefill capture wrapper

Thanks again for the paper + the clean OSS release. Let us know if there's anything we can test on Blackwell / Qwen-hybrid as you keep iterating.

— Lna-Lab (Shinka-Man / TonoKen3)

Activity

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

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions