From a443fdabecb4e8f2aa122c3ace3b3349501cf964 Mon Sep 17 00:00:00 2001 From: Beichen-Ma Date: Sat, 27 Jun 2026 23:03:34 +0000 Subject: [PATCH 1/3] feat: add out-of-tree Gemma4 KV connector --- .../data_generation/gemma4_kv_connector.py | 351 ++++++++++++++++++ tests/e2e/run_gemma4_kv_extraction.py | 126 +++++++ tests/e2e/smoke/test_gemma4_kv_connector.py | 34 ++ tests/e2e/utils.py | 129 +++++++ 4 files changed, 640 insertions(+) create mode 100644 src/speculators/data_generation/gemma4_kv_connector.py create mode 100644 tests/e2e/run_gemma4_kv_extraction.py create mode 100644 tests/e2e/smoke/test_gemma4_kv_connector.py diff --git a/src/speculators/data_generation/gemma4_kv_connector.py b/src/speculators/data_generation/gemma4_kv_connector.py new file mode 100644 index 000000000..b61ea520b --- /dev/null +++ b/src/speculators/data_generation/gemma4_kv_connector.py @@ -0,0 +1,351 @@ +"""Out-of-tree KV connector for Gemma4 MTP training-data extraction. + +Load it out-of-tree via --kv-transfer-config: + + { + "kv_connector": "Gemma4KVConnector", + "kv_connector_module_path": + "speculators.data_generation.gemma4_kv_connector", + "kv_role": "kv_producer", + "kv_connector_extra_config": {"shared_storage_path": ""} + } +""" + +import fcntl +import os +from collections import defaultdict +from dataclasses import dataclass +from functools import partial +from typing import TYPE_CHECKING, Any + +import torch +from vllm.distributed.communication_op import tensor_model_parallel_gather +from vllm.distributed.kv_transfer.kv_connector.v1.example_hidden_states_connector import ( # noqa: E501 + ExampleHiddenStatesConnector, + PendingSave, + extract_from_kv_cache, +) +from vllm.logger import init_logger + +if TYPE_CHECKING: + from vllm.config import VllmConfig + from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorRole + from vllm.v1.kv_cache_interface import KVCacheConfig + from vllm.v1.request import Request + +logger = init_logger(__name__) + +_KV_CACHE_NDIM = 5 +_KV_CACHE_KV_AXIS_SIZE = 2 + +GEMMA4_KV_KEYS = ( + "kv_last_local_k", + "kv_last_local_v", + "kv_last_global_k", + "kv_last_global_v", +) + + +def extract_real_kv_from_cache( + kv_cache: torch.Tensor, + slot_mapping: torch.Tensor, + num_tokens: int, +) -> tuple[torch.Tensor, torch.Tensor]: + """Extract per-token K and V from a real attention layer's paged cache. + + Args: + kv_cache: The layer's paged KV buffer, 5-D as described above. + slot_mapping: Per-token physical cache slots (1-D, length >= num_tokens). + num_tokens: Number of leading tokens to extract. + + Returns: + A (K, V) tuple, each of shape + (num_tokens, num_kv_heads, head_size). + """ + _bad_layout_msg = ( + "Gemma4KVConnector expects the 5-D KV layout " + "(num_blocks, 2, block_size, num_kv_heads, head_size) used by the " + f"FlashAttention/Triton backends; got shape {tuple(kv_cache.shape)}." + ) + assert kv_cache.dim() == _KV_CACHE_NDIM, _bad_layout_msg + assert kv_cache.shape[1] == _KV_CACHE_KV_AXIS_SIZE, _bad_layout_msg + block_size = kv_cache.shape[2] + block_idx = slot_mapping // block_size + block_off = slot_mapping % block_size + k = kv_cache[:, 0][block_idx, block_off][:num_tokens] + v = kv_cache[:, 1][block_idx, block_off][:num_tokens] + return k, v + + +@dataclass +class Gemma4PendingSave(PendingSave): + """PendingSave plus the per-group block_ids of the two verifier KV + layers, carried scheduler -> worker so their KV can be extracted.""" + + local_block_ids: list[int] | None = None + global_block_ids: list[int] | None = None + + +class Gemma4KVConnector(ExampleHiddenStatesConnector): + """Hidden-states + verifier-KV extraction connector for Gemma4 MTP training.""" + + def __init__( + self, + vllm_config: "VllmConfig", + role: "KVConnectorRole", + kv_cache_config: "KVCacheConfig", + ): + super().__init__(vllm_config, role, kv_cache_config) + + self._tp_size = vllm_config.parallel_config.tensor_parallel_size + pp = vllm_config.parallel_config.pipeline_parallel_size + assert pp == 1, ( + f"Gemma4KVConnector does not support pipeline_parallel_size>1 " + f"(got pp={pp})." + ) + + tcfg = vllm_config.model_config.hf_config.get_text_config() + self._local_total_kv_heads = tcfg.num_key_value_heads + + global_kv = getattr(tcfg, "num_global_key_value_heads", None) + self._global_total_kv_heads = global_kv or tcfg.num_key_value_heads + + cache_dtype = vllm_config.cache_config.cache_dtype + assert cache_dtype == "auto", ( + "Gemma4KVConnector requires an unquantized KV cache " + f"(cache_dtype='auto'); got {cache_dtype!r}. Quantized-cache " + "handling (per-token head scales) is not yet supported." + ) + + self.verifier_local_layer, self._local_group_idx = self._resolve_layer( + "sliding_attention" + ) + self.verifier_global_layer, self._global_group_idx = self._resolve_layer( + "full_attention" + ) + + self._local_kv_cache: torch.Tensor | None = None + self._global_kv_cache: torch.Tensor | None = None + + logger.info( + "Gemma4KVConnector: local=%s (group %d) global=%s (group %d)", + self.verifier_local_layer, + self._local_group_idx, + self.verifier_global_layer, + self._global_group_idx, + ) + + def _resolve_layer(self, attn_type: str) -> tuple[str, int]: + """Resolve the last non-KV-shared verifier layer of attn_type. + + Args: + attn_type: "sliding_attention" or "full_attention". + + Returns: + A (layer_name, kv_cache_group_idx) tuple. + """ + tcfg = self._vllm_config.model_config.hf_config.get_text_config() + layer_types = list(getattr(tcfg, "layer_types", [])) + if not layer_types: + raise ValueError( + "Gemma4KVConnector requires the verifier config to expose " + "'layer_types'; none found." + ) + n_shared = getattr(tcfg, "num_kv_shared_layers", 0) + num_non_shared = len(layer_types) - n_shared + + type_to_indices: dict[str, list[int]] = defaultdict(list) + for idx, lt in enumerate(layer_types[:num_non_shared]): + type_to_indices[lt].append(idx) + indices = type_to_indices.get(attn_type) + if not indices: + counts = {k: len(v) for k, v in type_to_indices.items()} + raise ValueError( + f"Gemma4KVConnector found no non-KV-shared '{attn_type}' " + f"layer; layer-type counts: {counts}" + ) + + prefix = self._resolve_attn_layer_prefix() + layer_name = f"{prefix}.{indices[-1]}.self_attn.attn" + return layer_name, self._find_group_idx(layer_name) + + def register_kv_caches(self, kv_caches: dict[str, torch.Tensor]): + super().register_kv_caches(kv_caches) + for name in (self.verifier_local_layer, self.verifier_global_layer): + if name not in kv_caches: + raise ValueError( + f"Resolved verifier KV layer {name!r} is not among the " + f"registered kv_caches keys: {sorted(kv_caches)[:8]}..." + ) + self._local_kv_cache = kv_caches[self.verifier_local_layer] + self._global_kv_cache = kv_caches[self.verifier_global_layer] + + def _resolve_attn_layer_prefix(self) -> str: + assert self._kv_cache_config is not None + for group in self._kv_cache_config.kv_cache_groups: + for name in group.layer_names: + if ".layers." in name and name.endswith(".self_attn.attn"): + return name.split(".layers.")[0] + ".layers" + raise ValueError( + "Gemma4KVConnector could not resolve the verifier attention-layer " + "name prefix (no '*.layers.N.self_attn.attn' layer found)." + ) + + def _find_group_idx(self, layer_name: str) -> int: + assert self._kv_cache_config is not None + for i, group in enumerate(self._kv_cache_config.kv_cache_groups): + if layer_name in group.layer_names: + return i + raise ValueError( + f"Gemma4KVConnector: layer {layer_name!r} not found in any kv_cache_group." + ) + + def request_finished_all_groups( + self, + request: "Request", + block_ids: tuple[list[int], ...], + ) -> tuple[bool, dict[str, Any] | None]: + result = super().request_finished_all_groups(request, block_ids) + + req_id = request.request_id + base = self._pending_saves.get(req_id) + if base is not None: + self._pending_saves[req_id] = Gemma4PendingSave( + req_id=base.req_id, + filename=base.filename, + token_ids=base.token_ids, + block_ids=base.block_ids, + local_block_ids=list(block_ids[self._local_group_idx]), + global_block_ids=list(block_ids[self._global_group_idx]), + ) + return result + + @staticmethod + def _slot_mapping_from_blocks( + block_ids: list[int], block_size: int + ) -> torch.Tensor: + block_ids_t = torch.tensor(block_ids, dtype=torch.long) + num_blocks = block_ids_t.shape[0] + offsets = torch.arange(0, block_size, dtype=torch.long) + slot_mapping = ( + offsets.reshape((1, block_size)) + + block_ids_t.reshape((num_blocks, 1)) * block_size + ) + return slot_mapping.flatten() + + def _gather_kv_heads( + self, x: torch.Tensor, total_kv_heads: int + ) -> torch.Tensor | None: + """Gather a per-rank KV-head shard onto rank 0 of the TP group.. + + Args: + x: This rank's shard, shape (num_tokens, num_kv_heads_local, head_size). + total_kv_heads: The unsharded KV-head count for this layer, used to + dedup under GQA replication. + + Returns: + (num_tokens, total_kv_heads, head_size) on rank 0; None on other ranks. + """ + if self._tp_size == 1: + return x + + gathered = tensor_model_parallel_gather(x, dst=0, dim=1) + if gathered is None: + return None + + stride = gathered.shape[1] // total_kv_heads + if stride > 1: + gathered = gathered[:, ::stride, :] + + return gathered.contiguous() + + def _submit_async_write(self, pending: PendingSave) -> None: + num_tokens = pending.token_ids.shape[0] + + gathered_kv: dict[str, torch.Tensor] = {} + if isinstance(pending, Gemma4PendingSave) and pending.local_block_ids: + for role, cache, blocks, total_heads in ( + ( + "local", + self._local_kv_cache, + pending.local_block_ids, + self._local_total_kv_heads, + ), + ( + "global", + self._global_kv_cache, + pending.global_block_ids, + self._global_total_kv_heads, + ), + ): + assert cache is not None + assert blocks is not None + slots = self._slot_mapping_from_blocks(blocks, cache.shape[2]).to( + device=cache.device, non_blocking=True + ) + k, v = extract_real_kv_from_cache(cache, slots, num_tokens) + # Returns None on non-rank-0; only rank 0 keeps the result. + gk = self._gather_kv_heads(k, total_heads) + gv = self._gather_kv_heads(v, total_heads) + if gk is not None and gv is not None: + gathered_kv[f"kv_last_{role}_k"] = gk + gathered_kv[f"kv_last_{role}_v"] = gv + + # Only rank 0 writes to disk. + if not self._is_tp_rank_zero: + return + assert self._kv_cache is not None + if not gathered_kv: + logger.warning( + "Gemma4KVConnector: req %s has no verifier block_ids; " + "writing hidden states only.", + pending.req_id, + ) + + copy_stream = self._get_copy_stream() + ready_event = torch.cuda.Event() + ready_event.record() + copy_stream.wait_event(ready_event) + + tensors: dict[str, torch.Tensor] = {} + with torch.cuda.stream(copy_stream): + hs_slots = self._slot_mapping_from_blocks( + pending.block_ids, self._kv_cache.shape[1] + ).to(device=self._kv_cache.device, non_blocking=True) + tensors["hidden_states"] = self._pin( + extract_from_kv_cache(self._kv_cache, hs_slots, num_tokens) + ) + # Pin the already-gathered (full-head) verifier KV tensors. + for key, full in gathered_kv.items(): + tensors[key] = self._pin(full) + + copy_done = torch.cuda.Event() + copy_done.record(copy_stream) + + assert not pending.token_ids.is_cuda + tensors["token_ids"] = pending.token_ids.clone() + + prior = self._req_futures.get(pending.req_id) + assert prior is None, "Found another KV transfer request with same req_id!" + + os.makedirs(os.path.dirname(pending.filename), exist_ok=True) + lock_fd = self._lock_fds.pop(pending.req_id, None) + if lock_fd is None and self.use_lock: + lock_path = pending.filename + ".lock" + lock_fd = os.open(lock_path, os.O_CREAT | os.O_WRONLY | os.O_TRUNC, 0o644) + fcntl.flock(lock_fd, fcntl.LOCK_EX) + + future = self._executor.submit( + self._write_tensors, tensors, copy_done, pending.filename, lock_fd + ) + self._req_copy_events[pending.req_id] = copy_done + self._req_futures[pending.req_id] = future + future.add_done_callback(partial(self._on_write_done, pending.req_id)) + + @staticmethod + def _pin(gpu_tensor: torch.Tensor) -> torch.Tensor: + """Async DtoH copy into pinned host memory. Call inside the copy + stream so the copy is ordered with the other per-request copies.""" + pinned = torch.empty_like(gpu_tensor, device="cpu", pin_memory=True) + pinned.copy_(gpu_tensor, non_blocking=True) + return pinned diff --git a/tests/e2e/run_gemma4_kv_extraction.py b/tests/e2e/run_gemma4_kv_extraction.py new file mode 100644 index 000000000..beac8624a --- /dev/null +++ b/tests/e2e/run_gemma4_kv_extraction.py @@ -0,0 +1,126 @@ +"""Subprocess driver for the Gemma4KVConnector extraction test. + +Generic, assertion-free counterpart to run_vllm.py: it is parametrized by +JSON (--llm-args, --prompts) and dumps raw per-request facts (tensor +shapes, token alignment, finiteness) to --results-file. All assertions live +parent-side in tests/e2e/utils.run_gemma4_kv_extraction. +""" + +import argparse +import json +import sys +import traceback +from pathlib import Path +from typing import Any + +_REPO_ROOT = Path(__file__).resolve().parents[2] +if str(_REPO_ROOT) not in sys.path: + sys.path.insert(0, str(_REPO_ROOT)) + +# isort: off +from tests.e2e.run_vllm import LLM, SamplingParams # noqa: E402 + +import torch # noqa: E402 +from vllm.distributed.kv_transfer.kv_connector.v1 import ( # noqa: E402 + example_hidden_states_connector, +) + +# isort: on + +from speculators.data_generation.gemma4_kv_connector import ( # noqa: E402 + GEMMA4_KV_KEYS as KV_KEYS, +) + + +def run_extraction(llm_args: dict, prompts: list[str]) -> dict: + """Run the extraction forward pass and collect raw per-request facts. + + No assertions: this only reports JSON-serializable facts (shapes, token + alignment, finiteness) for the parent process to assert on. + """ + llm = LLM(**llm_args) + + model_config = llm.llm_engine.model_config + hidden_size = model_config.get_hidden_size() + text_config = model_config.hf_config.get_text_config() + local_total_heads = text_config.num_key_value_heads + global_total_heads = ( + getattr(text_config, "num_global_key_value_heads", None) or local_total_heads + ) + + outputs = llm.generate(prompts, SamplingParams(max_tokens=1, temperature=0.0)) + + per_request: list[dict[str, Any]] = [] + for output in outputs: + kv_params = output.kv_transfer_params + path = kv_params.get("hidden_states_path") if kv_params else None + if path is None: + per_request.append({"hidden_states_path": None}) + continue + + obj = example_hidden_states_connector.load_hidden_states(path) + n = len(output.prompt_token_ids) + per_request.append( + { + "hidden_states_path": path, + "num_tokens": n, + "token_ids_aligned": bool( + torch.equal(obj["token_ids"], torch.tensor(output.prompt_token_ids)) + ), + "hidden_states_shape": list(obj["hidden_states"].shape), + "kv": { + key: { + "present": key in obj, + "shape": list(obj[key].shape) if key in obj else None, + "finite": bool(torch.isfinite(obj[key]).all()) + if key in obj + else None, + } + for key in KV_KEYS + }, + } + ) + + return { + "num_outputs": len(outputs), + "hidden_size": hidden_size, + "local_total_heads": local_total_heads, + "global_total_heads": global_total_heads, + "per_request": per_request, + } + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument( + "--llm-args", + type=str, + required=True, + help="JSON-serialized kwargs for LLM instantiation", + ) + parser.add_argument( + "--prompts", type=str, required=True, help="JSON-serialized list of prompts" + ) + parser.add_argument( + "--results-file", + type=Path, + required=True, + help="File to write the JSON-serialized extraction facts", + ) + return parser.parse_args() + + +if __name__ == "__main__": + args = parse_args() + try: + results = run_extraction(json.loads(args.llm_args), json.loads(args.prompts)) + exit_code = 0 + except Exception as exc: # noqa: BLE001 + traceback.print_exc() + results = {"ok": False, "error": f"{type(exc).__name__}: {exc}"} + exit_code = 1 + + with args.results_file.open("w", encoding="utf-8") as f: + json.dump(results, f) + + sys.exit(exit_code) diff --git a/tests/e2e/smoke/test_gemma4_kv_connector.py b/tests/e2e/smoke/test_gemma4_kv_connector.py new file mode 100644 index 000000000..d54e30b38 --- /dev/null +++ b/tests/e2e/smoke/test_gemma4_kv_connector.py @@ -0,0 +1,34 @@ +"""E2E test for the out-of-tree Gemma4KVConnector. + +Gemma4KVConnector extends vLLM's ExampleHiddenStatesConnector to also export the +verifier's last sliding-window (local) and last full-attention (global) layer KV +caches alongside the hidden states, into one .safetensors per request. This is +the training-data producer for finetuning Gemma4 MTP draft models, whose +query-only attention borrows exactly those two verifier KV caches. +""" + +from pathlib import Path + +import pytest + +from tests.conftest import requires_cuda, requires_multi_gpu +from tests.e2e.utils import run_gemma4_kv_extraction + +MODEL = "google/gemma-4-E4B-it" + + +@pytest.mark.e2e +@pytest.mark.slow +@requires_cuda +def test_gemma4_kv_connector_smoke(tmp_path: Path): + """Single-GPU (tp=1) extraction: shapes, alignment, both verifier KV caches.""" + run_gemma4_kv_extraction(MODEL, tmp_path, tensor_parallel_size=1) + + +@pytest.mark.e2e +@pytest.mark.slow +@requires_multi_gpu +def test_gemma4_kv_connector_tp2(tmp_path: Path): + """tp=2: verifier KV heads are sharded per rank; the connector all-gathers + (and dedups GQA replicas) so saved tensors carry the full head count.""" + run_gemma4_kv_extraction(MODEL, tmp_path, tensor_parallel_size=2) diff --git a/tests/e2e/utils.py b/tests/e2e/utils.py index 0891818e6..04462e490 100644 --- a/tests/e2e/utils.py +++ b/tests/e2e/utils.py @@ -23,6 +23,7 @@ "launch_vllm_server_context", "purge_newfiles", "run_data_generation_offline", + "run_gemma4_kv_extraction", "run_prepare_data", "run_stitch_mtp", "run_training", @@ -31,6 +32,13 @@ "wait_for_server", ] +GEMMA4_KV_KEYS = ( + "kv_last_local_k", + "kv_last_local_v", + "kv_last_global_k", + "kv_last_global_v", +) + def purge_newfiles(fn: Callable[..., Path]): """Decorator that turns a Path-returning function into a context manager. @@ -452,3 +460,124 @@ def run_vllm_engine( assert acci >= thresholdi, ( f"Acceptance {acci} at token {i} is less than threshold {thresholdi}" ) + + +def run_gemma4_kv_extraction( + model: str, + tmp_path: Path, + tensor_parallel_size: int = 1, + aux_layer_ids: Iterable[int] = (2, 21, 39, 42), + max_model_len: int = 2048, + gpu_memory_utilization: float = 0.4, + timeout: float | None = None, +): + """Drive Gemma4KVConnector extraction via run_gemma4_kv_extraction.py.""" + logger.info("vLLM Python executable: {}", VLLM_PYTHON) + + driver_file = str(Path(__file__).with_name("run_gemma4_kv_extraction.py")) + shared_storage = tmp_path / "kv_out" + shared_storage.mkdir(exist_ok=True) + results_file = tmp_path / "results.json" + aux_layer_ids = list(aux_layer_ids) + + llm_args_dict = { + "model": model, + "tensor_parallel_size": tensor_parallel_size, + "speculative_config": { + "method": "extract_hidden_states", + "num_speculative_tokens": 1, + "draft_model_config": { + "hf_config": {"eagle_aux_hidden_state_layer_ids": aux_layer_ids} + }, + }, + "kv_transfer_config": { + "kv_connector": "Gemma4KVConnector", + "kv_connector_module_path": ( + "speculators.data_generation.gemma4_kv_connector" + ), + "kv_role": "kv_producer", + "kv_connector_extra_config": {"shared_storage_path": str(shared_storage)}, + }, + "max_model_len": max_model_len, + "enforce_eager": True, + "enable_chunked_prefill": False, + "gpu_memory_utilization": gpu_memory_utilization, + "load_format": "dummy", + } + + # Short, and a long prompt past the 512 sliding window, in one batch. + prompts = [ + "Hello world", + "Test prompt with several tokens", + " ".join(["word"] * 600), + ] + + command = [ + VLLM_PYTHON, + driver_file, + "--llm-args", + json.dumps(llm_args_dict), + "--prompts", + json.dumps(prompts), + "--results-file", + str(results_file), + ] + logger.info("run_gemma4_kv_extraction.py command:\n {}", command) + + env = os.environ.copy() + # Avoid a flashinfer JIT-build failure. + env["VLLM_USE_FLASHINFER_SAMPLER"] = "0" + + result = subprocess.run( # noqa: S603 + command, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + check=False, + env=env, + timeout=timeout, + ) + logger.info( + "run_gemma4_kv_extraction.py output:\n{}", indent(result.stdout, " ") + ) + assert result.returncode == 0, ( + f"run_gemma4_kv_extraction.py exited with non-zero return code: " + f"{result.returncode}" + ) + + with results_file.open(encoding="utf-8") as f: + results = json.load(f) + + assert results["num_outputs"] == len(prompts) + assert len(results["per_request"]) == len(prompts) + expected_heads = { + "kv_last_local_k": results["local_total_heads"], + "kv_last_local_v": results["local_total_heads"], + "kv_last_global_k": results["global_total_heads"], + "kv_last_global_v": results["global_total_heads"], + } + hidden_size = results["hidden_size"] + for req in results["per_request"]: + path = req.get("hidden_states_path") + assert path is not None, "request produced no hidden_states_path" + n = req["num_tokens"] + assert req["token_ids_aligned"], "saved token_ids do not match prompt tokens" + assert req["hidden_states_shape"] == [n, len(aux_layer_ids), hidden_size] + for key in GEMMA4_KV_KEYS: + kv = req["kv"][key] + assert kv["present"], f"missing {key}" + shape = kv["shape"] + assert len(shape) == 3, f"{key} expected 3-D, got {shape}" + assert shape[0] == n, f"{key} rows {shape[0]} != n_tokens {n}" + # full (unsharded) head count even under TP + assert shape[1] == expected_heads[key], ( + f"{key} head count {shape[1]} != unsharded {expected_heads[key]} " + f"(tp={tensor_parallel_size})" + ) + assert kv["finite"], f"{key} has non-finite values" + # local (sliding) and global (full) caches use different head dims. + local_shape = req["kv"]["kv_last_local_k"]["shape"] + global_shape = req["kv"]["kv_last_global_k"]["shape"] + assert local_shape[1:] != global_shape[1:] + + return results From d5fd02d5f243bb852a9718eeb3149edd6db77404 Mon Sep 17 00:00:00 2001 From: Beichen-Ma Date: Sun, 28 Jun 2026 04:16:52 +0000 Subject: [PATCH 2/3] fix --- .../data_generation/gemma4_kv_connector.py | 33 ++++++++++--------- 1 file changed, 17 insertions(+), 16 deletions(-) diff --git a/src/speculators/data_generation/gemma4_kv_connector.py b/src/speculators/data_generation/gemma4_kv_connector.py index b61ea520b..5fe8dc72e 100644 --- a/src/speculators/data_generation/gemma4_kv_connector.py +++ b/src/speculators/data_generation/gemma4_kv_connector.py @@ -62,13 +62,12 @@ def extract_real_kv_from_cache( A (K, V) tuple, each of shape (num_tokens, num_kv_heads, head_size). """ - _bad_layout_msg = ( - "Gemma4KVConnector expects the 5-D KV layout " - "(num_blocks, 2, block_size, num_kv_heads, head_size) used by the " - f"FlashAttention/Triton backends; got shape {tuple(kv_cache.shape)}." - ) - assert kv_cache.dim() == _KV_CACHE_NDIM, _bad_layout_msg - assert kv_cache.shape[1] == _KV_CACHE_KV_AXIS_SIZE, _bad_layout_msg + if kv_cache.dim() != _KV_CACHE_NDIM or kv_cache.shape[1] != _KV_CACHE_KV_AXIS_SIZE: + raise ValueError( + "Gemma4KVConnector expects the 5-D KV layout " + "(num_blocks, 2, block_size, num_kv_heads, head_size) used by the " + f"FlashAttention/Triton backends; got shape {tuple(kv_cache.shape)}." + ) block_size = kv_cache.shape[2] block_idx = slot_mapping // block_size block_off = slot_mapping % block_size @@ -99,10 +98,11 @@ def __init__( self._tp_size = vllm_config.parallel_config.tensor_parallel_size pp = vllm_config.parallel_config.pipeline_parallel_size - assert pp == 1, ( - f"Gemma4KVConnector does not support pipeline_parallel_size>1 " - f"(got pp={pp})." - ) + if pp != 1: + raise ValueError( + f"Gemma4KVConnector does not support pipeline_parallel_size>1 " + f"(got pp={pp})." + ) tcfg = vllm_config.model_config.hf_config.get_text_config() self._local_total_kv_heads = tcfg.num_key_value_heads @@ -111,11 +111,12 @@ def __init__( self._global_total_kv_heads = global_kv or tcfg.num_key_value_heads cache_dtype = vllm_config.cache_config.cache_dtype - assert cache_dtype == "auto", ( - "Gemma4KVConnector requires an unquantized KV cache " - f"(cache_dtype='auto'); got {cache_dtype!r}. Quantized-cache " - "handling (per-token head scales) is not yet supported." - ) + if cache_dtype != "auto": + raise ValueError( + "Gemma4KVConnector requires an unquantized KV cache " + f"(cache_dtype='auto'); got {cache_dtype!r}. Quantized-cache " + "handling (per-token head scales) is not yet supported." + ) self.verifier_local_layer, self._local_group_idx = self._resolve_layer( "sliding_attention" From 332d36be1569715eb6310a8f9d2e6a988df4284d Mon Sep 17 00:00:00 2001 From: Beichen-Ma Date: Sun, 28 Jun 2026 09:13:58 +0000 Subject: [PATCH 3/3] fix: trim slot_mapping before indexing --- src/speculators/data_generation/gemma4_kv_connector.py | 6 ++++-- 1 file changed, 4 insertions(+), 2 deletions(-) diff --git a/src/speculators/data_generation/gemma4_kv_connector.py b/src/speculators/data_generation/gemma4_kv_connector.py index 5fe8dc72e..ebaefcfe3 100644 --- a/src/speculators/data_generation/gemma4_kv_connector.py +++ b/src/speculators/data_generation/gemma4_kv_connector.py @@ -69,10 +69,12 @@ def extract_real_kv_from_cache( f"FlashAttention/Triton backends; got shape {tuple(kv_cache.shape)}." ) block_size = kv_cache.shape[2] + + slot_mapping = slot_mapping[:num_tokens] block_idx = slot_mapping // block_size block_off = slot_mapping % block_size - k = kv_cache[:, 0][block_idx, block_off][:num_tokens] - v = kv_cache[:, 1][block_idx, block_off][:num_tokens] + k = kv_cache[:, 0][block_idx, block_off] + v = kv_cache[:, 1][block_idx, block_off] return k, v