From a18231ca158d2ad747807498ec9231fac80f62cf Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 17:39:35 +0800 Subject: [PATCH 01/11] feat: implement per-head h2o prefill retention --- docs/en/features/sparse-methods.md | 25 +- docs/zh/features/sparse-methods.md | 21 +- src/sparsevllm/configs/groups.py | 1 + src/sparsevllm/configs/sparse.py | 15 +- src/sparsevllm/engine/cache_manager/h2o.py | 288 +++++++----------- .../engine/cache_manager/h2o_retention.py | 120 ++++++++ .../cache_manager/storage/explicit_kv.py | 32 ++ src/sparsevllm/engine/chain_cache.py | 1 + src/sparsevllm/engine/llm_engine.py | 1 + src/sparsevllm/engine/sparse_controller.py | 9 + src/sparsevllm/engine/sparse_methods/base.py | 19 ++ src/sparsevllm/engine/sparse_methods/h2o.py | 163 ++++++++-- .../engine/sparse_methods/h2o_selection.py | 42 +++ .../kernels/triton/prefill_score.py | 84 +++-- src/sparsevllm/layers/attention.py | 3 +- src/sparsevllm/layers/mla_attention.py | 3 +- src/sparsevllm/method_registry.py | 46 +-- tests/test_chain_prefix_cache.py | 2 + tests/test_h2o_cache_manager.py | 212 +++++-------- tests/test_h2o_per_head_prefill.py | 195 ++++++++++++ tests/test_h2o_prefill_probability_kernel.py | 83 +++++ tests/test_prefill_attention_provider.py | 22 +- 22 files changed, 943 insertions(+), 444 deletions(-) create mode 100644 src/sparsevllm/engine/cache_manager/h2o_retention.py create mode 100644 src/sparsevllm/engine/sparse_methods/h2o_selection.py create mode 100644 tests/test_h2o_per_head_prefill.py create mode 100644 tests/test_h2o_prefill_probability_kernel.py diff --git a/docs/en/features/sparse-methods.md b/docs/en/features/sparse-methods.md index d203de4f..4dff645b 100644 --- a/docs/en/features/sparse-methods.md +++ b/docs/en/features/sparse-methods.md @@ -49,16 +49,21 @@ compatibility rule. SnapKV defaults `sparse_prefill_score_mode` to `logits`; `probability` remains an explicit reproducibility option because its additional normalized QK sweep is substantially more expensive in measured long-context prefill. PyramidKV -and H2O continue to default to `probability`. For the shared H2O prompt-scoring -state this is the canonical path: every KV layer independently sums its -normalized softmax attention probabilities over the full current query chunk, -then accumulates that attention mass across prefill chunks. Decode score -collection and eviction are intentionally disabled. Sparse-vLLM -reuses FA3's softmax LSE and performs one additional QK sweep because FlashAttention -does not materialize its probability matrix. `h2o_prefill_score_window=0` selects -the full current chunk and is the canonical default. A nonzero window in `[1, 128]` -or explicit `logits` mode is a non-canonical approximation; neither changes the -requirement that every H2O KV layer computes and retains its own prefill score. +and H2O use `probability`. H2O accumulates FP32 probability sums per query +head across prefill chunks. At eviction, `h2o_head_reduction=max` (default) or +`mean` combines the cumulative scores within each GQA KV group, or across the +whole layer for MLA. MHA heads select independently. Each selection retains +heavy hitters plus recent tokens within its token budget; native GQA KV sharing +and MLA latent storage are preserved. MLA H2O prefill currently requires TP1. + +`h2o_prefill_score_window=0` scores all queries in each chunk. Windows in +`[1, 128]` are explicit approximations. H2O rejects `logits` mode because a +reduced logit vector cannot represent per-head cumulative probabilities. +When attention provides its softmax LSE, scoring reuses it; otherwise it +recomputes normalization from the same visible keys. Intermediate chunk eviction +changes subsequent attention, so results can depend on chunk size and budgets. +Decode scoring and eviction remain disabled; `h2o_decode_budget` determines the +final prefill retention budget, and the cache grows during generation. ## Prefill Scheduling Policies diff --git a/docs/zh/features/sparse-methods.md b/docs/zh/features/sparse-methods.md index f06fef69..0ea1fb87 100644 --- a/docs/zh/features/sparse-methods.md +++ b/docs/zh/features/sparse-methods.md @@ -38,15 +38,18 @@ prefill attention 计算。它们是同一条轴上的备选项,可以分别 SnapKV 的 `sparse_prefill_score_mode` 默认值改为 `logits`;`probability` 仍可显式启用以复现实验,但它需要额外执行归一化 QK sweep,在已测长上下文 -prefill 中开销明显更高。PyramidKV 和 H2O 继续默认使用 `probability`。对两个阶段 -共享的 H2O prompt scoring state 而言,这是 canonical 路径:每个 KV layer 都独立地对 -完整当前 query chunk 的归一化 softmax attention probability 求和,并在 -prefill chunk 之间累计 attention mass。当前明确关闭 decode score 收集与 -淘汰。Sparse-vLLM 复用 FA3 的 -softmax LSE;由于 FlashAttention 不物化 probability matrix,还需额外执行 -一遍 QK。`h2o_prefill_score_window=0` 表示完整当前 chunk,是 canonical -默认设置;`[1, 128]` 的非零 window 或显式 `logits` 模式均属于非 canonical -近似,但都不会改变每个 H2O KV layer 必须独立计算并保存 prefill score 的要求。 +prefill 中开销明显更高。PyramidKV 和 H2O 使用 `probability`。H2O 逐 query head 跨 prefill chunks +累计 FP32 概率和;驱逐时通过 `h2o_head_reduction=max`(默认)或 `mean`, +在 GQA 的每个 KV 组内或 MLA 的整个 layer 内归约累计分数。MHA 各 head 独立选择。 +每套选择在预算内保留 heavy hitters 和 recent tokens,保持 GQA 的原生 KV 共享 +以及 MLA 的原生 latent 存储。MLA H2O prefill 当前要求 TP1。 + +`h2o_prefill_score_window=0` 观察完整当前 chunk;`[1, 128]` 的窗口属于显式近似。 +H2O 拒绝 `logits` 模式,因为归约后的 logits 无法表示逐 head 累计概率。 +Attention 提供 softmax LSE 时复用该结果,否则使用同一可见 KV 集合重新计算归一化。 +中间 chunk 的实际驱逐会改变后续 attention,结果因此可能随 chunk size 和预算变化。 +Decode 评分和驱逐仍保持关闭;`h2o_decode_budget` 用于最后一个 prefill chunk 的保留预算, +之后缓存随生成增长。 ## Prefill Scheduling Policy diff --git a/src/sparsevllm/configs/groups.py b/src/sparsevllm/configs/groups.py index 4dcf9e23..6c1846ac 100644 --- a/src/sparsevllm/configs/groups.py +++ b/src/sparsevllm/configs/groups.py @@ -60,6 +60,7 @@ class SparseMethodConfig: h2o_prefill_budget: int = 8192 h2o_recent_ratio: float = 0.5 h2o_prefill_score_window: int = 0 + h2o_head_reduction: str = "max" rkv_compression_interval: int = 128 rkv_observation_tokens: int = 8 diff --git a/src/sparsevllm/configs/sparse.py b/src/sparsevllm/configs/sparse.py index d4990f95..b86959ee 100644 --- a/src/sparsevllm/configs/sparse.py +++ b/src/sparsevllm/configs/sparse.py @@ -148,6 +148,12 @@ def _normalize_snapkv(config) -> None: def _normalize_h2o(config) -> None: + reduction = str(getattr(config, "h2o_head_reduction", "max")).strip().lower() + if reduction not in {"max", "mean"}: + raise ValueError("h2o_head_reduction must be 'max' or 'mean'.") + config.h2o_head_reduction = reduction + if getattr(config, "sparse_prefill_score_mode", "probability") != "probability": + raise ValueError("H2O per-head accumulation requires sparse_prefill_score_mode='probability'.") _normalize_positive_int(config, "h2o_decode_budget", fallback=0) _normalize_positive_int(config, "h2o_decode_eviction_interval", fallback=0) _normalize_int_attr(config, "h2o_prefill_budget", fallback=0) @@ -162,14 +168,7 @@ def _normalize_h2o(config) -> None: f"h2o_recent_ratio must be in (0, 1), got {config.h2o_recent_ratio}." ) _normalize_int_attr(config, "h2o_prefill_score_window", fallback=0) - score_mode = getattr(config, "sparse_prefill_score_mode", "probability") - if score_mode == "logits": - if config.h2o_prefill_score_window < 0: - raise ValueError( - "h2o_prefill_score_window must be non-negative in logits " - f"mode (0 means the full chunk), got {config.h2o_prefill_score_window}." - ) - elif not 0 <= config.h2o_prefill_score_window <= 128: + if not 0 <= config.h2o_prefill_score_window <= 128: raise ValueError( "h2o_prefill_score_window must be in [0, 128] in probability mode " "(0 means the full current chunk), got " diff --git a/src/sparsevllm/engine/cache_manager/h2o.py b/src/sparsevllm/engine/cache_manager/h2o.py index 29e32e9d..77df5b3a 100644 --- a/src/sparsevllm/engine/cache_manager/h2o.py +++ b/src/sparsevllm/engine/cache_manager/h2o.py @@ -17,6 +17,7 @@ from .base import ExplicitKVPayload, PrefillComputeView from .snapkv import SnapKVCacheManager +from .h2o_retention import H2OPrefillRetentionMixin, H2O_PREFILL_QUERY_TILE from .storage import ExplicitKVStorage @@ -26,16 +27,8 @@ class _H2ORowRef(NamedTuple): seq_id: int -class H2OCacheManager(SnapKVCacheManager): - """H2O physical KV eviction with one score vector per layer and sequence. - - Sparse-vLLM owns one physical token row shared by all KV heads, so this v1 - implementation maintains one cumulative normalized token-importance vector - aligned with that row. The probability prefill path accumulates normalized - attention mass. The logits path max-reduces raw QK over the observation - queries and query heads, normalizes that token vector, and then accumulates - it on the same per-query mass scale. - """ +class H2OCacheManager(H2OPrefillRetentionMixin, SnapKVCacheManager): + """Native KV storage with cumulative per-query-head prefill probabilities.""" def __init__( self, @@ -50,6 +43,7 @@ def __init__( allocation_budget_bytes=allocation_budget_bytes, ) self._h2o_scores: dict[tuple[int, int], torch.Tensor] = {} + self._h2o_positions: dict[tuple[int, int], torch.Tensor] = {} # Decode rows remain reclaimable while temporarily absent from a # scheduled batch. Keep only ids here: caching full Sequence objects # would retain their logical token histories on every worker. @@ -97,6 +91,52 @@ def h2o_decode_enabled(self) -> bool: getattr(self.config, "sparse_method", None) ) == "h2o" + def _get_available_slots_info(self) -> tuple[int, int]: + from .storage import MlaLatentStorage + + available, payload_bytes = super()._get_available_slots_info() + mla = isinstance(self.attention_cache_storage, MlaLatentStorage) + if mla and self.tp_size != 1: + raise ValueError("H2O MLA prefill currently requires TP1 for layer-wide head reduction.") + heads = int(self.hf_config.num_attention_heads) // self.tp_size + groups = self.h2o_selection_groups + if heads <= 0 or groups <= 0 or heads % groups: + raise ValueError("H2O requires complete local query-to-KV head groups.") + + # Account for score/position retention copies and selection indices in + # addition to native payload. Decode may grow the cache but adds no scores. + metadata_per_slot = 20 * heads + 64 * groups + 64 + batch = min(int(self.max_buffer_rows), int(self.config.max_num_seqs_in_batch)) + chunk = min(int(self.config.engine_prefill_chunk_size), int(self.max_model_len)) + width = min(int(self.max_model_len), self.h2o_prefill_budget + chunk) + if bool(getattr(self.config, "enable_prefix_caching", False)): + width = int(self.max_model_len) + window = int(self.config.h2o_prefill_score_window) + queries = min(H2O_PREFILL_QUERY_TILE, chunk, window or chunk) + # Probability QK stats use at most two padded head rows per real head + # and at most one bounded query tile at a time. + padded_queries = max(16, 1 << (queries - 1).bit_length()) + stats = 8 * batch * (2 * heads) * padded_queries * ((width + 63) // 64 + 1) + scores = 8 * batch * heads * width + # Include replacement allocation while the previous workspace is live. + workspace_bytes = 2 * (stats + scores) + 2 * self.h2o_prefill_budget * payload_bytes + if workspace_bytes >= available: + raise RuntimeError( + "Not enough memory for H2O prefill score/retention workspaces: " + f"required={workspace_bytes} available={available}." + ) + self._h2o_reserved_workspace_bytes = workspace_bytes + self._h2o_metadata_bytes_per_slot = metadata_per_slot + return available - workspace_bytes, payload_bytes + metadata_per_slot + + def _iter_accounting_tensors(self): + yield from super()._iter_accounting_tensors() + storage = getattr(self, "attention_cache_storage", None) + workspace = getattr(storage, "_head_copy_workspace", None) + if workspace is not None: + for index, tensor in enumerate(workspace): + yield f"h2o_compaction_workspace.{index}", tensor + def _prefill_append_peak( self, resident_len: int, @@ -584,15 +624,12 @@ def _require_score_length( expected_len: int, ) -> torch.Tensor: score = self.h2o_score(layer_idx, seq.seq_id) - if score is None: - raise RuntimeError( - f"H2O score vector is missing: layer={layer_idx} seq_id={seq.seq_id}." - ) - if int(score.numel()) != int(expected_len): + if score is None or int(score.shape[-1]) != int(expected_len): raise RuntimeError( - "H2O score vector is not aligned with the physical KV row: " - f"layer={layer_idx} seq_id={seq.seq_id} scores={int(score.numel())} " - f"physical_len={int(expected_len)}." + "H2O scores are not aligned with the physical KV row: " + f"layer={layer_idx} seq_id={seq.seq_id} " + f"shape={None if score is None else tuple(score.shape)} " + f"physical_len={expected_len}." ) return score @@ -603,17 +640,13 @@ def _expand_score( *, device: torch.device, ) -> torch.Tensor: - new_len = int(new_len) - if new_len < 0: - raise ValueError(f"H2O score length must be non-negative, got {new_len}.") - old_len = 0 if score is None else int(score.numel()) - if old_len > new_len: - raise RuntimeError( - f"H2O score vector cannot shrink without keep_indices: old={old_len} new={new_len}." - ) - expanded = torch.zeros((new_len,), dtype=torch.float32, device=device) - if score is not None and old_len > 0: - expanded[:old_len].copy_(score.to(device=device, dtype=torch.float32)) + old_len = 0 if score is None else int(score.shape[-1]) + if new_len < old_len or new_len < 0: + raise RuntimeError("H2O scores cannot shrink without retention indices.") + shape = () if score is None else tuple(score.shape[:-1]) + expanded = torch.zeros((*shape, new_len), dtype=torch.float32, device=device) + if score is not None: + expanded[..., :old_len].copy_(score) return expanded @classmethod @@ -625,13 +658,14 @@ def _accumulate_score( new_len: int, weight: float, ) -> torch.Tensor: - if step_score.dim() != 1 or int(step_score.numel()) < int(new_len): - raise ValueError( - "H2O step score must be a 1D vector covering new_len: " - f"shape={tuple(step_score.shape)} new_len={int(new_len)}." - ) + if step_score.ndim not in (1, 2) or step_score.shape[-1] < new_len: + raise ValueError("H2O step scores must cover the physical cache length.") + if previous is not None and previous.shape[:-1] != step_score.shape[:-1]: + raise ValueError("H2O cannot change score heads during accumulation.") + if previous is None: + previous = step_score.new_zeros((*step_score.shape[:-1], 0)) cumulative = cls._expand_score(previous, new_len, device=step_score.device) - cumulative.add_(step_score[:new_len].float(), alpha=float(weight)) + cumulative.add_(step_score[..., :new_len].float(), alpha=float(weight)) return cumulative @staticmethod @@ -705,6 +739,7 @@ def _prepare_prefill(self, seqs: list[Sequence]): if logical_start == 0: for layer_idx in score_layer_ids: self._h2o_scores.pop(self._score_key(layer_idx, seq.seq_id), None) + self._h2o_positions.pop(self._score_key(layer_idx, seq.seq_id), None) for layer_idx in layer_ids: row_idx = self._get_free_row(layer_idx, int(seq.seq_id)) @@ -716,7 +751,21 @@ def _prepare_prefill(self, seqs: list[Sequence]): ) if logical_start > 0 and layer_idx in score_layer_ids: self._require_score_length(layer_idx, seq, physical_start) + key = self._score_key(layer_idx, seq.seq_id) + groups = self.h2o_selection_groups + old_positions = self._h2o_positions.get(key) + if physical_start and ( + old_positions is None or tuple(old_positions.shape) != (groups, physical_start) + ): + raise RuntimeError("H2O positions are not aligned before prefill append.") + new_positions = torch.arange( + logical_start, logical_end, dtype=torch.int64, device=self.device, + ).expand(groups, -1) + positions = new_positions.clone() if old_positions is None else torch.cat( + (old_positions, new_positions), dim=-1, + ) self._allocate(layer_idx, int(seq.seq_id), chunk_size) + self._h2o_positions[key] = positions physical_end = physical_start + chunk_size layers_slot_mapping[ layer_idx, token_offset : token_offset + chunk_size @@ -790,143 +839,10 @@ def prefill_score_ranges( ranges.append((batch_idx, seq, prompt_cache_len, score_start, score_end)) return ranges - @torch.no_grad() - def collect_prefill_attention_score( - self, - layer_idx: int, - q: torch.Tensor, - view: PrefillComputeView, - *, - b_start_loc: torch.Tensor, - chunk_lens: torch.Tensor, - attention_lse: torch.Tensor | None = None, - ): - ctx = get_context() - if not ctx.is_prefill: - raise RuntimeError("H2O prefill score collection was called outside prefill.") - seqs = getattr(ctx, "seqs", None) - if seqs is None: - raise RuntimeError("H2O prefill score collection requires current seqs in context.") - if int(chunk_lens.ndim) != 1 or int(chunk_lens.shape[0]) != len(seqs): - raise RuntimeError( - "H2O prefill scoring chunk-length batch mismatch: " - f"shape={tuple(chunk_lens.shape)} seqs={len(seqs)}." - ) - ranges = self.prefill_score_ranges(layer_idx, seqs) - if not ranges: - return None - if not isinstance(view.payload, ExplicitKVPayload): - raise TypeError( - "H2O prefill scoring requires ExplicitKVPayload, got " - f"{type(view.payload).__name__}." - ) - meta = view.meta - payload = view.payload - - context_lens = tuple(int(item[4]) for item in ranges) - prepared_context_lens = getattr( - self, - "_prefill_context_lens_cpu_by_layer", - {}, - ).get(int(layer_idx)) - if prepared_context_lens is None and meta.context_lens.device.type == "cpu": - prepared_context_lens = tuple( - int(value) for value in meta.context_lens.tolist() - ) - if prepared_context_lens is None: - raise RuntimeError( - "H2O prefill scoring requires CPU context lengths prepared " - f"for layer={layer_idx}." - ) - if tuple(prepared_context_lens) != context_lens: - raise RuntimeError( - "H2O prefill score view is not in compressed physical coordinates: " - f"layer={layer_idx} view={tuple(prepared_context_lens)} " - f"physical={context_lens}." - ) - prompt_cache_lens_cpu = tuple(int(item[2]) for item in ranges) - score_starts_cpu = tuple(int(item[3]) for item in ranges) - score_ends_cpu = tuple(int(item[4]) for item in ranges) - ( - prompt_cache_lens, - _batch_indices, - score_starts, - score_ends, - ) = self._cached_prefill_score_metadata_tensors( - device=q.device, - context_lens=context_lens, - prompt_cache_lens=prompt_cache_lens_cpu, - batch_indices=tuple(range(len(ranges))), - score_starts=score_starts_cpu, - score_ends=score_ends_cpu, - ) - max_context_len = max(context_lens) - if meta.attn_score is None: - step_score = self._prefill_step_score_buffer( - batch_size=len(seqs), - max_context_len=max_context_len, - device=q.device, - ) - max_score_len = max(item[4] - item[3] for item in ranges) - if attention_lse is None: - self._run_prefill_score( - q, - payload.k_cache, - step_score, - meta, - b_start_loc, - prompt_cache_lens, - max_score_len, - score_starts, - score_ends, - candidate_start=0, - recent_keep_tokens=0, - ) - else: - if self.config.sparse_prefill_score_mode != "probability": - raise RuntimeError( - "FA3 softmax LSE is only valid for probability H2O scoring." - ) - from sparsevllm.kernels.triton.prefill_score import ( - prefill_score_from_lse_fwd, - ) - - prefill_score_from_lse_fwd( - q, - payload.k_cache, - attention_lse, - step_score, - meta.req_indices, - b_start_loc, - meta.context_lens, - prompt_cache_lens, - max_score_len, - meta.active_slots, - score_starts, - score_ends, - workspace=getattr(self, "_prefill_score_workspace", None), - ) - else: - if ( - self.config.sparse_prefill_score_mode != "logits" - or int(self.config.h2o_prefill_score_window) != 0 - ): - raise RuntimeError( - "H2O main-attention prefill scores require logits mode with " - "h2o_prefill_score_window=0." - ) - step_score = meta.attn_score - if ( - step_score.ndim != 2 - or int(step_score.shape[0]) != len(seqs) - or int(step_score.shape[1]) < max_context_len - ): - raise ValueError( - "H2O fused prefill scores must have shape [batch, context], " - f"got {tuple(step_score.shape)} for batch={len(seqs)} " - f"max_context_len={max_context_len}." - ) - for batch_idx, seq, prompt_cache_len, score_start, score_end in ranges: + def accumulate_prefill_scores( + self, layer_idx: int, seqs: list[Sequence], step_score: torch.Tensor, + ) -> None: + for batch_idx, seq, prompt_cache_len, score_start, score_end in self.prefill_score_ranges(layer_idx, seqs): key = self._score_key(layer_idx, seq.seq_id) previous = self._h2o_scores.get(key) if prompt_cache_len > 0 and previous is None: @@ -935,24 +851,18 @@ def collect_prefill_attention_score( f"layer={layer_idx} seq_id={seq.seq_id} " f"prompt_cache_len={prompt_cache_len}." ) - if previous is not None and int(previous.numel()) != prompt_cache_len: + if previous is not None and int(previous.shape[-1]) != prompt_cache_len: raise RuntimeError( "H2O prefill score vector lost physical-row alignment before append: " - f"layer={layer_idx} seq_id={seq.seq_id} scores={int(previous.numel())} " + f"layer={layer_idx} seq_id={seq.seq_id} scores={int(previous.shape[-1])} " f"prompt_cache_len={prompt_cache_len}." ) - effective_queries = score_end - score_start score_row = step_score[batch_idx] - if self.config.sparse_prefill_score_mode == "logits": - score_row = self._normalize_logit_prefill_score( - score_row, - new_len=score_end, - ) cumulative = self._accumulate_score( previous, score_row, new_len=score_end, - weight=float(effective_queries), + weight=1.0, ) self._h2o_scores[key] = cumulative return None @@ -1814,6 +1724,9 @@ def free_seq(self, seq_id: int): for key in list(self._h2o_scores): if key[1] == seq_id: self._h2o_scores.pop(key, None) + for key in list(self._h2o_positions): + if key[1] == seq_id: + self._h2o_positions.pop(key, None) super().free_seq(seq_id) def on_chain_turn_finished( @@ -1841,10 +1754,21 @@ def on_chain_turn_finished( physical_len, device=score.device, ) + positions = self._h2o_positions.get(key) + if positions is not None: + appended = physical_len - positions.shape[-1] + if appended < 0: + raise RuntimeError("H2O chain positions exceed the physical cache length.") + suffix = torch.arange( + processed_token_count - appended, processed_token_count, + dtype=torch.int64, device=positions.device, + ).expand(positions.shape[0], -1) + self._h2o_positions[key] = torch.cat((positions, suffix), dim=-1) super().on_chain_turn_finished(seq_id, processed_token_count) def reset_after_warmup(self) -> None: self._h2o_scores.clear() + self._h2o_positions.clear() self._h2o_active_decode_seq_ids.clear() self._h2o_decode_score_workspace = None self._h2o_decode_score_signature = None @@ -1858,7 +1782,7 @@ def debug_state_summary(self) -> dict[str, object]: summary["h2o"] = { "counters": dict(self._h2o_counters), "score_lengths": { - f"{layer_idx}:{seq_id}": int(score.numel()) + f"{layer_idx}:{seq_id}": int(score.shape[-1]) for (layer_idx, seq_id), score in sorted(self._h2o_scores.items()) }, "final_prefill_workspace": ( diff --git a/src/sparsevllm/engine/cache_manager/h2o_retention.py b/src/sparsevllm/engine/cache_manager/h2o_retention.py new file mode 100644 index 00000000..5a3e4887 --- /dev/null +++ b/src/sparsevllm/engine/cache_manager/h2o_retention.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +from dataclasses import dataclass + +import torch + + +H2O_PREFILL_QUERY_TILE = 128 + + +@dataclass(frozen=True) +class H2ORetention: + """One layer/request selection in the current packed cache coordinates.""" + + layer_idx: int + seq_id: int + source_length: int + keep: torch.Tensor # [native KV heads, budget], or [1, budget] for MLA + final_prefill: bool + + +class H2OPrefillRetentionMixin: + """Cache-owned per-head state and overlap-safe physical retention.""" + + @property + def h2o_selection_groups(self) -> int: + from .storage import MlaLatentStorage + + return ( + 1 if isinstance(self.attention_cache_storage, MlaLatentStorage) + else self.num_kv_heads + ) + + def commit_h2o_retention(self, requests: list[H2ORetention]) -> None: + from .storage import ExplicitKVStorage + + prepared = [] + seen = set() + seen_rows = set() + release_counts = {} + storage = self.attention_cache_storage + for request in requests: + layer, seq_id, length = request.layer_idx, request.seq_id, request.source_length + key = (layer, seq_id) + if key in seen: + raise ValueError("Duplicate H2O retention request.") + seen.add(key) + row = self.seq_id_to_row[layer][seq_id] + if (layer, row) in seen_rows: + raise ValueError("H2O retention requests share a physical row.") + seen_rows.add((layer, row)) + if int(self.row_seq_lens[layer][row]) != length: + raise ValueError("H2O retention refers to a stale physical row.") + keep = request.keep + groups = self.h2o_selection_groups + if ( + keep.ndim != 2 or keep.shape[0] != groups + or keep.dtype != torch.long or keep.device != self.device + or not 0 < keep.shape[-1] < length + ): + raise ValueError("H2O retention requires [selection_groups, budget] int64 indices.") + budget = int(keep.shape[-1]) + self._assert_final_prefill_tensor( + ((keep >= 0) & (keep < length)).all() + & (keep[:, 1:] > keep[:, :-1]).all(), + "H2O retention indices must be in bounds and strictly increasing.", + ) + score = self._h2o_scores[key] + positions = self._h2o_positions[key] + if ( + score.ndim != 2 or score.shape[-1] != length + or score.shape[0] % groups or tuple(positions.shape) != (groups, length) + ): + raise ValueError("H2O retention score/position metadata is not aligned.") + slots = self.buffer_req_to_token_slots[layer][row, :length].long().clone() + ordered_slots = slots.sort().values + self._assert_final_prefill_tensor( + ((slots >= 0) & (slots < storage.slot_capacity())).all() + & (ordered_slots[1:] > ordered_slots[:-1]).all(), + "H2O retention physical slots must be valid and unique.", + ) + release_counts[layer] = release_counts.get(layer, 0) + length - budget + end = int(self._num_free_slots[layer]) + release_counts[layer] + if end > self.free_slots_stack[layer].numel(): + raise RuntimeError("H2O retention would overflow the free-slot stack.") + # Every query head keeps its own history, including heads that did + # not supply the group's maximum on this step. + score_keep = keep.repeat_interleave(score.shape[0] // groups, dim=0) + kept_score = score.gather(1, score_keep).contiguous() + kept_positions = positions.gather(1, keep).contiguous() + prepared.append((request, row, slots, ordered_slots, kept_score, kept_positions)) + + # Validate the complete submission before publishing any row mutation. + for request, row, slots, ordered_slots, score, positions in prepared: + layer, seq_id, length = request.layer_idx, request.seq_id, request.source_length + keep = request.keep + budget = int(keep.shape[-1]) + destination = ordered_slots[:budget] + released = ordered_slots[budget:] + kv_layer = self.kv_layer_index(layer) + selected = slots[keep] + if isinstance(storage, ExplicitKVStorage): + storage.copy_head_slots(kv_layer, selected, destination) + else: + storage.copy_slots(kv_layer, selected[0], destination) + ptr = int(self._num_free_slots[layer]) + self.free_slots_stack[layer][ptr:ptr + released.numel()] = released + self._num_free_slots[layer] = ptr + released.numel() + self.buffer_req_to_token_slots[layer][row, :budget] = destination + self.buffer_req_to_token_slots[layer][row, budget:length] = 0 + self.row_seq_lens[layer][row] = budget + self._h2o_scores[(layer, seq_id)] = score + self._h2o_positions[(layer, seq_id)] = positions + counter = "final_prefill_evictions" if request.final_prefill else "intermediate_prefill_evictions" + self._h2o_counters[counter] += 1 + self._h2o_counters["dropped_tokens"] += length - budget + if prepared: + self._uniform_decode_metadata = False + self._decode_static_state_binding_key = None + self._invalidate_h2o_decode_score_workspace() diff --git a/src/sparsevllm/engine/cache_manager/storage/explicit_kv.py b/src/sparsevllm/engine/cache_manager/storage/explicit_kv.py index eecade9b..83afffd2 100644 --- a/src/sparsevllm/engine/cache_manager/storage/explicit_kv.py +++ b/src/sparsevllm/engine/cache_manager/storage/explicit_kv.py @@ -33,6 +33,37 @@ def __init__( f"num_kv_heads={self.num_kv_heads} head_dim={self.head_dim}." ) self.kv_cache: torch.Tensor | None = None + self._head_copy_workspace: tuple[torch.Tensor, torch.Tensor] | None = None + + def copy_head_slots( + self, layer_idx: int, source_slots: torch.Tensor, destination_slots: torch.Tensor, + ) -> None: + """Pack [Hkv, B] source slices into B native-width slots, with overlap.""" + payload = self.layer_payload(layer_idx) + if ( + source_slots.ndim != 2 or source_slots.shape[0] != self.num_kv_heads + or destination_slots.ndim != 1 + or source_slots.shape[1] != destination_slots.numel() + ): + raise ValueError("Per-head KV copy requires [Hkv, B] sources and [B] destinations.") + heads = torch.arange(self.num_kv_heads, device=source_slots.device)[:, None] + source = (source_slots * self.num_kv_heads + heads).reshape(-1) + destination = (destination_slots[None, :] * self.num_kv_heads + heads).reshape(-1) + count = source.numel() + workspace = self._head_copy_workspace + if workspace is None or workspace[0].shape[0] < count: + workspace = ( + payload.k_cache.new_empty((count, self.head_dim)), + payload.v_cache.new_empty((count, self.head_dim)), + ) + self._head_copy_workspace = workspace + k = payload.k_cache.view(-1, self.head_dim) + v = payload.v_cache.view(-1, self.head_dim) + scratch_k, scratch_v = workspace[0][:count], workspace[1][:count] + torch.index_select(k, 0, source, out=scratch_k) + torch.index_select(v, 0, source, out=scratch_v) + k.index_copy_(0, destination, scratch_k) + v.index_copy_(0, destination, scratch_v) def allocate( self, @@ -57,6 +88,7 @@ def allocate( dtype=self.dtype, device=device, ) + self._head_copy_workspace = None def _require_cache(self) -> torch.Tensor: if self.kv_cache is None: diff --git a/src/sparsevllm/engine/chain_cache.py b/src/sparsevllm/engine/chain_cache.py index c0569447..fac6306a 100644 --- a/src/sparsevllm/engine/chain_cache.py +++ b/src/sparsevllm/engine/chain_cache.py @@ -225,6 +225,7 @@ def build_chain_cache_fingerprint(config: Any) -> bytes: "h2o_prefill_budget", "h2o_recent_ratio", "h2o_prefill_score_window", + "h2o_head_reduction", "sparse_prefill_score_mode", "sparse_attn_score_dtype", ), diff --git a/src/sparsevllm/engine/llm_engine.py b/src/sparsevllm/engine/llm_engine.py index e0489dac..8ebc1922 100644 --- a/src/sparsevllm/engine/llm_engine.py +++ b/src/sparsevllm/engine/llm_engine.py @@ -1294,6 +1294,7 @@ def worker_info( "h2o_prefill_budget", "h2o_recent_ratio", "h2o_prefill_score_window", + "h2o_head_reduction", "pool_kernel_size", "sparse_attn_score_dtype", "pyramid_layer_ratios", diff --git a/src/sparsevllm/engine/sparse_controller.py b/src/sparsevllm/engine/sparse_controller.py index e6e79f05..a7babf5d 100644 --- a/src/sparsevllm/engine/sparse_controller.py +++ b/src/sparsevllm/engine/sparse_controller.py @@ -16,11 +16,20 @@ create_sparse_method_runtime, ) from sparsevllm.utils.context import get_context +from sparsevllm.engine.sparse_methods.base import PrefillScoreEvent class SparseController: """Method-agnostic sparse lifecycle facade used by the inference engine.""" + def collect_prefill_attention_score( + self, layer_idx, q, view, *, b_start_loc, chunk_lens, + softmax_scale: float, attention_lse=None, + ) -> None: + self.runtime.collect_prefill_attention_score(PrefillScoreEvent( + layer_idx, q, view, b_start_loc, chunk_lens, softmax_scale, attention_lse, + )) + def __init__(self, config: Config, cache_manager: CacheManager): self.config = config self.cache_manager = cache_manager diff --git a/src/sparsevllm/engine/sparse_methods/base.py b/src/sparsevllm/engine/sparse_methods/base.py index a0642daa..770e07b8 100644 --- a/src/sparsevllm/engine/sparse_methods/base.py +++ b/src/sparsevllm/engine/sparse_methods/base.py @@ -10,6 +10,7 @@ from sparsevllm.config import Config from sparsevllm.engine.cache_manager import CacheManager, SparseSelection from sparsevllm.engine.cache_manager.base import ( + PrefillComputeView, _debug_tensor_summary, _debug_value_summary, ) @@ -61,6 +62,17 @@ class PrefillSelectionRequest: forward_context: Any +@dataclass(frozen=True) +class PrefillScoreEvent: + layer_idx: int + query: torch.Tensor + view: PrefillComputeView + b_start_loc: torch.Tensor + chunk_lens: torch.Tensor + softmax_scale: float + attention_lse: torch.Tensor | None = None + + @dataclass(frozen=True) class DecodeSelectionRequest: layer_idx: int @@ -84,6 +96,13 @@ class LayerEndEvent: class SparseMethodRuntime(ABC): """Method-owned logical sparse runtime behind the controller facade.""" + def collect_prefill_attention_score(self, event: PrefillScoreEvent) -> None: + self.cache_manager.collect_prefill_attention_score( + event.layer_idx, event.query, event.view, + b_start_loc=event.b_start_loc, chunk_lens=event.chunk_lens, + attention_lse=event.attention_lse, + ) + def __init__(self, config: Config, cache_manager: CacheManager): self.config = config self.cache_manager = cache_manager diff --git a/src/sparsevllm/engine/sparse_methods/h2o.py b/src/sparsevllm/engine/sparse_methods/h2o.py index 4981388f..61dcb5b2 100644 --- a/src/sparsevllm/engine/sparse_methods/h2o.py +++ b/src/sparsevllm/engine/sparse_methods/h2o.py @@ -4,19 +4,26 @@ from sparsevllm.engine.sequence import Sequence from sparsevllm.method_registry import ( - h2o_uses_fused_prefill_score, normalize_sparse_method, resolve_prefill_sparse_method, ) +from sparsevllm.engine.cache_manager.h2o_retention import H2ORetention, H2O_PREFILL_QUERY_TILE from sparsevllm.utils.profiler import profiler +from sparsevllm.utils.context import get_context +from sparsevllm.engine.cache_manager.base import ExplicitKVPayload +from sparsevllm.kernels.triton.prefill_score import PrefillScoreWorkspace -from .base import SparseStepContext +from .base import SparseStepContext, PrefillScoreEvent from .passthrough import PassThroughRuntime +from .h2o_selection import select_h2o_heads class H2ORuntime(PassThroughRuntime): def __init__(self, config, cache_manager): super().__init__(config, cache_manager) + self._prefill_score_workspace = PrefillScoreWorkspace() + self._prefill_head_score_buffer: torch.Tensor | None = None + self._prefill_head_score_total: torch.Tensor | None = None self._h2o_decode_attn_score_buffers: dict[ tuple[int, ...], torch.Tensor, @@ -38,12 +45,8 @@ def needs_attention_score( layer_idx: int, step: SparseStepContext, ) -> bool: - del layer_idx - return ( - h2o_uses_fused_prefill_score(self.config) - if step.is_prefill - else False - ) + del layer_idx, step + return False def prefill_score_shape( self, @@ -51,25 +54,143 @@ def prefill_score_shape( num_heads: int, max_len: int, ) -> tuple[int, ...]: - del num_heads - return batch_size, max_len + return batch_size, num_heads, max_len def prefill_score_fill_value(self) -> float: - return -torch.inf + return 0.0 + + @torch.no_grad() + def collect_prefill_attention_score(self, event: PrefillScoreEvent) -> None: + manager = self.cache_manager + layer_idx, q, view = event.layer_idx, event.query, event.view + b_start_loc, chunk_lens = event.b_start_loc, event.chunk_lens + attention_lse = event.attention_lse + ctx = get_context() + if not ctx.is_prefill: + raise RuntimeError("H2O prefill score collection was called outside prefill.") + seqs = getattr(ctx, "seqs", None) + if seqs is None: + raise RuntimeError("H2O prefill score collection requires current seqs in context.") + if int(chunk_lens.ndim) != 1 or int(chunk_lens.shape[0]) != len(seqs): + raise RuntimeError( + "H2O prefill scoring chunk-length batch mismatch: " + f"shape={tuple(chunk_lens.shape)} seqs={len(seqs)}." + ) + ranges = manager.prefill_score_ranges(layer_idx, seqs) + if not ranges: + return None + if not isinstance(view.payload, ExplicitKVPayload): + raise TypeError( + "H2O prefill scoring requires ExplicitKVPayload, got " + f"{type(view.payload).__name__}." + ) + meta = view.meta + payload = view.payload + + context_lens = tuple(int(item[4]) for item in ranges) + prepared_context_lens = getattr( + manager, + "_prefill_context_lens_cpu_by_layer", + {}, + ).get(int(layer_idx)) + if prepared_context_lens is None and meta.context_lens.device.type == "cpu": + prepared_context_lens = tuple( + int(value) for value in meta.context_lens.tolist() + ) + if prepared_context_lens is None: + raise RuntimeError( + "H2O prefill scoring requires CPU context lengths prepared " + f"for layer={layer_idx}." + ) + if tuple(prepared_context_lens) != context_lens: + raise RuntimeError( + "H2O prefill score view is not in compressed physical coordinates: " + f"layer={layer_idx} view={tuple(prepared_context_lens)} " + f"physical={context_lens}." + ) + prompt_cache_lens_cpu = tuple(int(item[2]) for item in ranges) + score_starts_cpu = tuple(int(item[3]) for item in ranges) + score_ends_cpu = tuple(int(item[4]) for item in ranges) + ( + prompt_cache_lens, + _batch_indices, + score_starts, + score_ends, + ) = manager._cached_prefill_score_metadata_tensors( + device=q.device, + context_lens=context_lens, + prompt_cache_lens=prompt_cache_lens_cpu, + batch_indices=tuple(range(len(ranges))), + score_starts=score_starts_cpu, + score_ends=score_ends_cpu, + ) + max_context_len = max(context_lens) + if meta.attn_score is not None: + raise ValueError("H2O requires per-head probability scores after attention.") + from sparsevllm.kernels.triton.prefill_score import ( + prefill_score_fwd, prefill_score_from_lse_fwd, + ) + + shape = (len(seqs), int(q.shape[1]), max_context_len) + buffer = self._prefill_head_score_buffer + if buffer is None or any(old < new for old, new in zip(buffer.shape, shape)): + buffer = torch.empty(shape, dtype=torch.float32, device=q.device) + self._prefill_head_score_buffer = buffer + self._prefill_head_score_total = torch.empty_like(buffer) + step_score = buffer[:shape[0], :shape[1], :shape[2]] + total = self._prefill_head_score_total[:shape[0], :shape[1], :shape[2]] + total.zero_() + score_kwargs = dict( + workspace=self._prefill_score_workspace, per_head=True, + softmax_scale=event.softmax_scale, + ) + max_queries = max(item[4] - item[3] for item in ranges) + # Bound QK statistics workspace independently of prompt chunk size. + # Tiles contribute probability sums; no head reduction occurs here. + for offset in range(0, max_queries, H2O_PREFILL_QUERY_TILE): + tile_start = torch.minimum(score_starts + offset, score_ends) + tile_end = torch.minimum(tile_start + H2O_PREFILL_QUERY_TILE, score_ends) + score_args = ( + step_score, meta.req_indices, b_start_loc, meta.context_lens, + prompt_cache_lens, min(H2O_PREFILL_QUERY_TILE, max_queries - offset), + meta.active_slots, tile_start, tile_end, + ) + if attention_lse is None: + prefill_score_fwd(q, payload.k_cache, *score_args, **score_kwargs) + else: + prefill_score_from_lse_fwd( + q, payload.k_cache, attention_lse, *score_args, **score_kwargs, + ) + total.add_(step_score) + manager.accumulate_prefill_scores(layer_idx, seqs, total) def finish_step(self, step: SparseStepContext) -> None: if not step.is_prefill: return - prefill_method = resolve_prefill_sparse_method( - getattr(self.config, "prefill_sparse_method", None), - sparse_method=getattr(self.config, "sparse_method", None), - ) - if prefill_method == "h2o_prefill": - self.cache_manager.evict_after_intermediate_prefill(step.seqs) - if normalize_sparse_method( - getattr(self.config, "sparse_method", None) - ) == "h2o": - self.cache_manager.compact_final_prefill_for_decode(step.seqs) + manager = self.cache_manager + requests = [] + for layer_idx in manager.kv_transformer_layer_indices(): + for seq in step.seqs: + final = bool(seq.is_last_chunk_prefill) + if final: + if not manager.h2o_decode_enabled: + continue + elif not manager.h2o_prefill_enabled: + continue + budget = manager.h2o_decode_budget if final else manager.h2o_prefill_budget + length = manager._physical_row_len(layer_idx, seq) + scores = manager._require_score_length(layer_idx, seq, length) + if length <= budget: + continue + keep = select_h2o_heads( + scores, + selection_groups=manager.h2o_selection_groups, + budget=budget, + recent_ratio=float(self.config.h2o_recent_ratio), + reduction=self.config.h2o_head_reduction, + ) + requests.append(H2ORetention(layer_idx, int(seq.seq_id), length, keep, final)) + manager.commit_h2o_retention(requests) def _h2o_kv_layer_indices(self) -> list[int]: return [ diff --git a/src/sparsevllm/engine/sparse_methods/h2o_selection.py b/src/sparsevllm/engine/sparse_methods/h2o_selection.py new file mode 100644 index 00000000..d6b78072 --- /dev/null +++ b/src/sparsevllm/engine/sparse_methods/h2o_selection.py @@ -0,0 +1,42 @@ +from __future__ import annotations + +import torch + + +def select_h2o_heads( + cumulative: torch.Tensor, + *, + selection_groups: int, + budget: int, + recent_ratio: float, + reduction: str, +) -> torch.Tensor: + """Rank cumulative [Hq, L] probabilities, returning [groups, min(B,L)]. + + A group is one native KV head for explicit KV, or the whole layer for + shared MLA latent storage. Reduction happens only after time accumulation. + Equal scores prefer older positions, independently of the device top-k. + """ + if cumulative.ndim != 2 or cumulative.shape[0] == 0: + raise ValueError("H2O cumulative scores must have shape [query_heads, length].") + heads, length = cumulative.shape + if selection_groups <= 0 or heads % selection_groups: + raise ValueError("H2O query heads must divide into complete selection groups.") + if budget <= 0 or not 0 < recent_ratio < 1: + raise ValueError("H2O requires a positive budget and recent_ratio in (0, 1).") + if reduction not in {"max", "mean"}: + raise ValueError("H2O head reduction must be 'max' or 'mean'.") + if length <= budget: + return torch.arange(length, device=cumulative.device).expand(selection_groups, -1) + + grouped = cumulative.reshape(selection_groups, heads // selection_groups, length) + ranks = grouped.amax(dim=1) if reduction == "max" else grouped.mean(dim=1) + recent_count = min(budget, max(1, int(budget * recent_ratio))) + recent_start = length - recent_count + heavy = torch.argsort( + ranks[:, :recent_start], dim=-1, descending=True, stable=True, + )[:, :budget - recent_count] + recent = torch.arange(recent_start, length, device=cumulative.device).expand( + selection_groups, -1, + ) + return torch.cat((heavy, recent), dim=-1).sort(dim=-1).values diff --git a/src/sparsevllm/kernels/triton/prefill_score.py b/src/sparsevllm/kernels/triton/prefill_score.py index c93967c0..fb9ec0e9 100644 --- a/src/sparsevllm/kernels/triton/prefill_score.py +++ b/src/sparsevllm/kernels/triton/prefill_score.py @@ -118,6 +118,7 @@ def _prefill_score_partial_stats_kernel( NUM_BLOCKS: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_ROWS: tl.constexpr, + HEAD_DIM: tl.constexpr, BLOCK_DMODEL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, @@ -165,7 +166,7 @@ def _prefill_score_partial_stats_kernel( + q_head[:, None] * stride_qh + offs_d[None, :] * stride_qd ) - q = tl.load(Q + off_q, mask=q_row_valid[:, None], other=0.0) + q = tl.load(Q + off_q, mask=q_row_valid[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0) start_n = cur_n_block * BLOCK_N kv_pos = start_n + offs_n @@ -177,7 +178,7 @@ def _prefill_score_partial_stats_kernel( other=0, ) off_k = kv_loc[None, :] * stride_ks + cur_kv_head * stride_kh + offs_d[:, None] * stride_kd - k = tl.load(K + off_k, mask=kv_in_candidate[None, :], other=0.0) + k = tl.load(K + off_k, mask=kv_in_candidate[None, :] & (offs_d[:, None] < HEAD_DIM), other=0.0) qk = tl.dot(q, k) * sm_scale causal_mask = q_abs_pos[:, None] >= kv_pos[None, :] @@ -261,6 +262,7 @@ def _prefill_score_final_kernel( HEAD_BLOCKS: tl.constexpr, QUERY_BLOCKS: tl.constexpr, WRITE_PER_HEAD: tl.constexpr, + SUM_QUERIES: tl.constexpr, USE_BATCH_INDICES: tl.constexpr, candidate_start: tl.constexpr, recent_keep_tokens: tl.constexpr, @@ -268,6 +270,7 @@ def _prefill_score_final_kernel( NUM_BLOCKS: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_ROWS: tl.constexpr, + HEAD_DIM: tl.constexpr, BLOCK_DMODEL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, @@ -317,7 +320,7 @@ def _prefill_score_final_kernel( + q_head[:, None] * stride_qh + offs_d[None, :] * stride_qd ) - q = tl.load(Q + off_q, mask=q_row_valid[:, None], other=0.0) + q = tl.load(Q + off_q, mask=q_row_valid[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0) start_n = cur_n_block * BLOCK_N kv_pos = start_n + offs_n @@ -329,7 +332,7 @@ def _prefill_score_final_kernel( other=0, ) off_k = kv_loc[None, :] * stride_ks + cur_kv_head * stride_kh + offs_d[:, None] * stride_kd - k = tl.load(K + off_k, mask=kv_in_candidate[None, :], other=0.0) + k = tl.load(K + off_k, mask=kv_in_candidate[None, :] & (offs_d[:, None] < HEAD_DIM), other=0.0) qk = tl.dot(q, k) * sm_scale valid = ( @@ -351,7 +354,7 @@ def _prefill_score_final_kernel( head_rows = row_head_in_block == head_idx head_score = tl.sum( tl.where(head_rows[:, None], probs, 0.0), axis=0 - ) / (score_q_len * 1.0) + ) / (1.0 if SUM_QUERIES else score_q_len * 1.0) if WRITE_PER_HEAD: output_local_head = cur_head_block * BLOCK_H + head_idx output_head = cur_kv_head * H_PER_KV + output_local_head @@ -407,10 +410,12 @@ def _prefill_probability_from_lse_kernel( HEAD_BLOCKS: tl.constexpr, QUERY_BLOCKS: tl.constexpr, WRITE_PER_HEAD: tl.constexpr, + SUM_QUERIES: tl.constexpr, SCORE_WIDTH: tl.constexpr, sm_scale: tl.constexpr, BLOCK_H: tl.constexpr, BLOCK_ROWS: tl.constexpr, + HEAD_DIM: tl.constexpr, BLOCK_DMODEL: tl.constexpr, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, @@ -450,7 +455,7 @@ def _prefill_probability_from_lse_kernel( + kv_slot[None, :] * stride_ks + cur_kv_head * stride_kh + offs_d[:, None] * stride_kd, - mask=kv_valid[None, :], + mask=kv_valid[None, :] & (offs_d[:, None] < HEAD_DIM), other=0.0, ) score_q_len = tl.maximum(score_q_end - score_q_start, 1) @@ -471,7 +476,7 @@ def _prefill_probability_from_lse_kernel( + q_token[:, None] * stride_qt + q_head[:, None] * stride_qh + offs_d[None, :] * stride_qd, - mask=q_valid[:, None], + mask=q_valid[:, None] & (offs_d[None, :] < HEAD_DIM), other=0.0, ) row_lse = tl.load( @@ -494,7 +499,7 @@ def _prefill_probability_from_lse_kernel( head_rows = row_head_in_block == head_idx head_score = tl.sum( tl.where(head_rows[:, None], probabilities, 0.0), axis=0 - ) / (score_q_len * 1.0) + ) / (1.0 if SUM_QUERIES else score_q_len * 1.0) if WRITE_PER_HEAD: output_local_head = cur_head_block * BLOCK_H + head_idx output_head = cur_kv_head * H_PER_KV + output_local_head @@ -679,13 +684,20 @@ def prefill_score_fwd( score_mode: str = "probability", workspace: PrefillScoreWorkspace | None = None, batch_indices: torch.Tensor | None = None, + per_head: bool = False, + softmax_scale: float | None = None, ): head_dim = q.shape[-1] assert k.shape[-1] == head_dim assert q.dtype == k.dtype assert q.stride(-1) == 1 and k.stride(-1) == 1 - assert attn_score.dim() == 2 - assert head_dim in {16, 32, 64, 128, 256} + assert attn_score.dim() == (3 if per_head else 2) + if per_head and (attn_score.shape[1] != q.shape[1] or attn_score.dtype != torch.float32): + raise ValueError("Per-head scores require FP32 [batch, query_heads, length].") + if per_head and score_mode != "probability": + raise ValueError("Per-head scores require probability mode.") + assert head_dim in {16, 32, 64, 128, 256} or (per_head and 16 <= head_dim <= 256) + sm_scale = float(head_dim) ** -0.5 if softmax_scale is None else float(softmax_scale) batch, head = score_q_start.shape[0], q.shape[1] if score_q_end.shape != score_q_start.shape: raise ValueError( @@ -732,7 +744,7 @@ def prefill_score_fwd( block_m = min(32, max(16, triton.next_power_of_2(max_score_len))) query_blocks = triton.cdiv(max_score_len, block_m) - max_candidate_end = int(attn_score.shape[1]) + max_candidate_end = int(attn_score.shape[-1]) if max_candidate_end <= 0: return @@ -801,8 +813,12 @@ def prefill_score_fwd( reduce_rows //= 2 workspace = PrefillScoreWorkspace() if workspace is None else workspace - write_per_head = query_blocks > 1 - if write_per_head: + write_per_head = per_head or query_blocks > 1 + if per_head: + head_score = attn_score + head_score.zero_() + head_score_strides = head_score.stride() + elif write_per_head: head_score = workspace.probability_head_score_buffer( batch_size=batch, query_heads=head, @@ -852,11 +868,12 @@ def prefill_score_fwd( USE_BATCH_INDICES=batch_indices is not None, candidate_start=int(candidate_start), recent_keep_tokens=int(recent_keep_tokens), - sm_scale=float(head_dim) ** -0.5, + sm_scale=sm_scale, NUM_BLOCKS=candidate_blocks, BLOCK_H=block_h, BLOCK_ROWS=block_rows, - BLOCK_DMODEL=head_dim, + HEAD_DIM=head_dim, + BLOCK_DMODEL=triton.next_power_of_2(head_dim), BLOCK_M=block_m, BLOCK_N=block_n, num_warps=dot_warps, @@ -906,20 +923,22 @@ def prefill_score_fwd( HEAD_BLOCKS=head_blocks, QUERY_BLOCKS=query_blocks, WRITE_PER_HEAD=write_per_head, + SUM_QUERIES=per_head, USE_BATCH_INDICES=batch_indices is not None, candidate_start=int(candidate_start), recent_keep_tokens=int(recent_keep_tokens), - sm_scale=float(head_dim) ** -0.5, + sm_scale=sm_scale, NUM_BLOCKS=candidate_blocks, BLOCK_H=block_h, BLOCK_ROWS=block_rows, - BLOCK_DMODEL=head_dim, + HEAD_DIM=head_dim, + BLOCK_DMODEL=triton.next_power_of_2(head_dim), BLOCK_M=block_m, BLOCK_N=block_n, num_warps=dot_warps, num_stages=3, ) - if write_per_head: + if write_per_head and not per_head: reduce_heads = triton.next_power_of_2(head) _prefill_probability_head_reduce_kernel[(batch, candidate_blocks)]( head_score, @@ -951,6 +970,8 @@ def prefill_score_from_lse_fwd( score_q_end: torch.Tensor, *, workspace: PrefillScoreWorkspace | None = None, + per_head: bool = False, + softmax_scale: float | None = None, ) -> None: """Reduce exact FA3 probabilities into one token vector per layer row.""" @@ -983,11 +1004,13 @@ def prefill_score_from_lse_fwd( or b_prompt_cache_len.numel() != batch ): raise ValueError("FA3 probability score metadata must have one entry per batch row.") - if attn_score.ndim != 2 or int(attn_score.shape[0]) != batch: + if attn_score.ndim != (3 if per_head else 2) or int(attn_score.shape[0]) != batch: raise ValueError( "FA3 probability score output must be [batch, context], got " f"{tuple(attn_score.shape)}." ) + if per_head and (attn_score.shape[1] != q.shape[1] or attn_score.dtype != torch.float32): + raise ValueError("Per-head scores require FP32 [batch, query_heads, length].") max_score_len = int(max_query_len) if max_score_len <= 0 or batch == 0: return @@ -1007,7 +1030,7 @@ def prefill_score_from_lse_fwd( block_m = min(32, max(16, triton.next_power_of_2(max_score_len))) query_blocks = triton.cdiv(max_score_len, block_m) block_n = 64 if head_dim >= 128 else 128 - candidate_blocks = triton.cdiv(int(attn_score.shape[1]), block_n) + candidate_blocks = triton.cdiv(int(attn_score.shape[-1]), block_n) max_rows = 256 block_h = min( triton.next_power_of_2(heads_per_kv), @@ -1017,12 +1040,15 @@ def prefill_score_from_lse_fwd( block_rows = block_h * block_m group_count = batch * kv_heads * head_blocks * query_blocks workspace = PrefillScoreWorkspace() if workspace is None else workspace - write_per_head = query_blocks > 1 - if write_per_head: + write_per_head = per_head or query_blocks > 1 + if per_head: + head_score = attn_score + head_score_strides = head_score.stride() + elif write_per_head: head_score = workspace.probability_head_score_buffer( batch_size=batch, query_heads=query_heads, - score_width=int(attn_score.shape[1]), + score_width=int(attn_score.shape[-1]), device=q.device, ) head_score.zero_() @@ -1066,17 +1092,19 @@ def prefill_score_from_lse_fwd( HEAD_BLOCKS=head_blocks, QUERY_BLOCKS=query_blocks, WRITE_PER_HEAD=write_per_head, - SCORE_WIDTH=int(attn_score.shape[1]), - sm_scale=float(head_dim) ** -0.5, + SUM_QUERIES=per_head, + SCORE_WIDTH=int(attn_score.shape[-1]), + sm_scale=float(head_dim) ** -0.5 if softmax_scale is None else float(softmax_scale), BLOCK_H=block_h, BLOCK_ROWS=block_rows, - BLOCK_DMODEL=head_dim, + HEAD_DIM=head_dim, + BLOCK_DMODEL=triton.next_power_of_2(head_dim), BLOCK_M=block_m, BLOCK_N=block_n, num_warps=8 if block_rows >= 128 else 4, num_stages=3, ) - if write_per_head: + if write_per_head and not per_head: reduce_heads = triton.next_power_of_2(query_heads) _prefill_probability_head_reduce_kernel[(batch, candidate_blocks)]( head_score, @@ -1085,7 +1113,7 @@ def prefill_score_from_lse_fwd( *attn_score.stride(), QUERY_HEADS=query_heads, REDUCE_HEADS=reduce_heads, - SCORE_WIDTH=int(attn_score.shape[1]), + SCORE_WIDTH=int(attn_score.shape[-1]), BLOCK_N=block_n, num_warps=8 if reduce_heads >= 64 else 4, num_stages=3, diff --git a/src/sparsevllm/layers/attention.py b/src/sparsevllm/layers/attention.py index 782da7fb..a9da3e12 100644 --- a/src/sparsevllm/layers/attention.py +++ b/src/sparsevllm/layers/attention.py @@ -121,13 +121,14 @@ def forward( attention_lse = prefill_result.softmax_lse else: o = prefill_result - cache_manager.collect_prefill_attention_score( + sparse_controller.collect_prefill_attention_score( layer_idx, q, prefill_view, b_start_loc=b_start_loc, chunk_lens=chunk_lens, attention_lse=attention_lse, + softmax_scale=self.scale, ) cache_manager.record_prefill_query( layer_idx, diff --git a/src/sparsevllm/layers/mla_attention.py b/src/sparsevllm/layers/mla_attention.py index f91c4ef5..0418864b 100644 --- a/src/sparsevllm/layers/mla_attention.py +++ b/src/sparsevllm/layers/mla_attention.py @@ -989,12 +989,13 @@ def run_cached_attention( chunk_lens=chunk_lens, ) explicit_view = self.build_prefill_explicit_view(workset) - cache_manager.collect_prefill_attention_score( + sparse_controller.collect_prefill_attention_score( layer_idx, q, explicit_view, b_start_loc=b_start_loc, chunk_lens=chunk_lens, + softmax_scale=self.spec.softmax_scale, ) cache_manager.record_prefill_query( layer_idx, diff --git a/src/sparsevllm/method_registry.py b/src/sparsevllm/method_registry.py index b4964723..d1ef39b3 100644 --- a/src/sparsevllm/method_registry.py +++ b/src/sparsevllm/method_registry.py @@ -167,6 +167,7 @@ class ModelRuntimeCompatibility: class PrefillScoreCollectionKind(Enum): NONE = auto() METHOD_OWNED_POSTHOC_REDUCED = auto() + METHOD_OWNED_POSTHOC_PER_HEAD = auto() MAIN_ATTENTION_REDUCED = auto() @@ -238,26 +239,14 @@ def sparse_prefill_attention_contract( prefill_sparse_method, sparse_method=normalized, ) - cache_method = resolve_cache_sparse_method( - normalized, - prefill_sparse_method=resolved_prefill_method, - ) - h2o_score_collection = cache_method == "h2o" + cache_method = resolve_cache_sparse_method(normalized, prefill_sparse_method=resolved_prefill_method) layer_varying_page_table = _PREFILL_LAYER_VARYING_PAGE_TABLE[cache_method] - fused_h2o_score = ( - h2o_score_collection - and resolved_prefill_method != "flashprefill_v2" - and resolve_sparse_prefill_score_mode( - normalized, - sparse_prefill_score_mode, - ) - == "logits" - and int(h2o_prefill_score_window) == 0 - ) - if fused_h2o_score: + if cache_method == "h2o": + if resolve_sparse_prefill_score_mode(normalized, sparse_prefill_score_mode) != "probability": + raise ValueError("H2O requires per-head probability prefill scoring.") return SparsePrefillAttentionContract( - main_score_kind=AttentionScoreKind.RAW_QK_REDUCED, - score_collection=PrefillScoreCollectionKind.MAIN_ATTENTION_REDUCED, + main_score_kind=AttentionScoreKind.NONE, + score_collection=PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_PER_HEAD, layer_varying_page_table=layer_varying_page_table, ) collection = ( @@ -272,27 +261,6 @@ def sparse_prefill_attention_contract( ) -def h2o_uses_fused_prefill_score(config) -> bool: - return ( - resolve_cache_sparse_method( - getattr(config, "sparse_method", None), - prefill_sparse_method=getattr(config, "prefill_sparse_method", None), - ) - == "h2o" - and resolve_prefill_sparse_method( - getattr(config, "prefill_sparse_method", None), - sparse_method=getattr(config, "sparse_method", None), - ) - != "flashprefill_v2" - and resolve_sparse_prefill_score_mode( - "h2o", - getattr(config, "sparse_prefill_score_mode", None), - ) - == "logits" - and int(getattr(config, "h2o_prefill_score_window", 0)) == 0 - ) - - def sparse_decode_attention_requires_scores(method: str | None) -> bool: """Return whether a prepared decode implementation must support scores.""" diff --git a/tests/test_chain_prefix_cache.py b/tests/test_chain_prefix_cache.py index 5b273d6d..62d8a992 100644 --- a/tests/test_chain_prefix_cache.py +++ b/tests/test_chain_prefix_cache.py @@ -1000,6 +1000,7 @@ def _h2o_fingerprint_config(**overrides): "h2o_prefill_budget": 8, "h2o_recent_ratio": 0.5, "h2o_prefill_score_window": 4, + "h2o_head_reduction": "max", "sparse_attn_score_dtype": "float32", } values.update(overrides) @@ -1014,6 +1015,7 @@ def _h2o_fingerprint_config(**overrides): ("h2o_prefill_budget", 9), ("h2o_recent_ratio", 0.25), ("h2o_prefill_score_window", 8), + ("h2o_head_reduction", "mean"), ("sparse_attn_score_dtype", "float16"), ], ) diff --git a/tests/test_h2o_cache_manager.py b/tests/test_h2o_cache_manager.py index e13a226a..cdc7b78e 100644 --- a/tests/test_h2o_cache_manager.py +++ b/tests/test_h2o_cache_manager.py @@ -28,6 +28,8 @@ from sparsevllm.engine.sparse_controller import SparseController from sparsevllm.engine.sparse_methods import SparseStepContext from sparsevllm.engine.sparse_methods.h2o import H2ORuntime +from sparsevllm.engine.sparse_methods.base import PrefillScoreEvent +from sparsevllm.kernels.triton.prefill_score import PrefillScoreWorkspace from sparsevllm.method_registry import ( PREFILL_POLICY_ALL_CHUNKED, ) @@ -93,6 +95,7 @@ def _manager_with_layer_rows( manager.head_dim = 2 manager.hf_config = SimpleNamespace(dtype=torch.float32) manager._h2o_scores = {} + manager._h2o_positions = {} manager._h2o_active_decode_seq_ids = set() manager._h2o_counters = { "intermediate_prefill_evictions": 0, @@ -283,7 +286,7 @@ def test_h2o_decode_does_not_request_scores_or_run_eviction(): ) assert runtime.needs_attention_score(0, decode) is False - assert runtime.needs_attention_score(0, prefill) is True + assert runtime.needs_attention_score(0, prefill) is False runtime.finish_step(decode) runtime.cache_manager.evict_after_decode.assert_not_called() @@ -501,142 +504,45 @@ def test_h2o_logit_prefill_score_rejects_nan_or_all_inf(): ) -def test_h2o_prefill_score_collection_accumulates_in_physical_coordinates(): - manager = _manager_with_rows([6]) - seq = _seq(0, 20, prefilled=8, chunk=2) - manager._h2o_scores[(0, 0)] = torch.tensor([1.0, 2.0, 3.0, 4.0]) - view = PrefillComputeView( - meta=AttentionViewMeta( - active_slots=manager.buffer_req_to_token_slots[0], - req_indices=torch.tensor([0], dtype=torch.int32), - context_lens=torch.tensor([6], dtype=torch.int32), - max_context_len=6, - ), - payload=ExplicitKVPayload( - k_cache=torch.empty((16, 1, 1)), - v_cache=torch.empty((16, 1, 1)), - ), - ) - set_context(is_prefill=True, cache_manager=manager, seqs=[seq]) - - def fake_run_prefill_score( - q, - k_cache, - attn_score, - meta, - b_start_loc, - prompt_cache_lens, - max_query_len, - score_starts, - score_ends, - **kwargs, - ): - del q, k_cache, meta, b_start_loc, max_query_len - assert prompt_cache_lens.tolist() == [4] - assert score_starts.tolist() == [4] - assert score_ends.tolist() == [6] - assert kwargs == {"candidate_start": 0, "recent_keep_tokens": 0} - attn_score[0, :6] = torch.tensor([0.1, 0.2, 0.3, 0.4, 0.5, 0.6]) - - with patch.object( - manager, - "_run_prefill_score", - side_effect=fake_run_prefill_score, - ): - manager.collect_prefill_attention_score( - 0, - torch.empty((2, 1, 1)), - view, - b_start_loc=torch.tensor([0], dtype=torch.int32), - chunk_lens=torch.tensor([2], dtype=torch.int32), - ) - - assert manager._h2o_scores[(0, 0)].tolist() == pytest.approx( - [1.2, 2.4, 3.6, 4.8, 1.0, 1.2] - ) - - -def test_h2o_logit_prefill_score_collection_normalizes_logits(): - manager = _manager_with_rows([6]) - manager.config.sparse_prefill_score_mode = "logits" - manager.config.h2o_prefill_score_window = 0 - seq = _seq(0, 20, prefilled=8, chunk=2) - manager._h2o_scores[(0, 0)] = torch.tensor([1.0, 2.0, 3.0, 4.0]) - view = PrefillComputeView( - meta=AttentionViewMeta( - active_slots=manager.buffer_req_to_token_slots[0], - req_indices=torch.tensor([0], dtype=torch.int32), - context_lens=torch.tensor([6], dtype=torch.int32), - max_context_len=6, - ), - payload=ExplicitKVPayload( - k_cache=torch.empty((16, 1, 1)), - v_cache=torch.empty((16, 1, 1)), - ), - ) - set_context(is_prefill=True, cache_manager=manager, seqs=[seq]) - logits = torch.arange(6, dtype=torch.float32) - - def fake_run_prefill_score(*args, **kwargs): - del kwargs - args[2][0, :6].copy_(logits) - - with patch.object( - manager, - "_run_prefill_score", - side_effect=fake_run_prefill_score, - ): - manager.collect_prefill_attention_score( - 0, - torch.empty((2, 1, 1)), - view, - b_start_loc=torch.tensor([0], dtype=torch.int32), - chunk_lens=torch.tensor([2], dtype=torch.int32), - ) - - expected = torch.tensor([1.0, 2.0, 3.0, 4.0, 0.0, 0.0]) - expected.add_(torch.softmax(logits, dim=0), alpha=2.0) - assert torch.equal(manager._h2o_scores[(0, 0)], expected) +def _score_runtime(manager): + runtime = object.__new__(H2ORuntime) + runtime.cache_manager = manager + runtime.config = manager.config + runtime._prefill_score_workspace = PrefillScoreWorkspace() + runtime._prefill_head_score_buffer = None + return runtime -def test_h2o_logit_prefill_score_collection_consumes_fused_main_score(): +def test_h2o_prefill_score_collection_accumulates_in_physical_coordinates(): manager = _manager_with_rows([6]) - manager.config.sparse_prefill_score_mode = "logits" - manager.config.h2o_prefill_score_window = 0 seq = _seq(0, 20, prefilled=8, chunk=2) - manager._h2o_scores[(0, 0)] = torch.tensor([1.0, 2.0, 3.0, 4.0]) - logits = torch.arange(6, dtype=torch.float32).unsqueeze(0) + previous = torch.tensor([[1., 2., 3., 4.], [4., 3., 2., 1.]]) + manager._h2o_scores[(0, 0)] = previous.clone() view = PrefillComputeView( meta=AttentionViewMeta( active_slots=manager.buffer_req_to_token_slots[0], req_indices=torch.tensor([0], dtype=torch.int32), context_lens=torch.tensor([6], dtype=torch.int32), max_context_len=6, - attn_score=logits, - ), - payload=ExplicitKVPayload( - k_cache=torch.empty((16, 1, 1)), - v_cache=torch.empty((16, 1, 1)), ), + payload=ExplicitKVPayload(k_cache=torch.empty(16, 2, 16), v_cache=torch.empty(16, 2, 16)), ) set_context(is_prefill=True, cache_manager=manager, seqs=[seq]) + increment = torch.tensor([[.1, .2, .3, .4, .5, .5], [.5, .5, .4, .3, .2, .1]]) - with patch.object( - manager, - "_run_prefill_score", - side_effect=AssertionError("launched posthoc scorer"), - ): - manager.collect_prefill_attention_score( - 0, - torch.empty((2, 1, 1)), - view, - b_start_loc=torch.tensor([0], dtype=torch.int32), - chunk_lens=torch.tensor([2], dtype=torch.int32), - ) + def produce_scores(q, k, output, reqs, starts, contexts, prefixes, max_queries, slots, score_start, score_end, **kwargs): + assert prefixes.tolist() == [4] + assert score_start.tolist() == [4] + assert score_end.tolist() == [6] + assert kwargs['per_head'] is True + assert kwargs['softmax_scale'] == .25 + output[0].copy_(increment) - expected = torch.tensor([1.0, 2.0, 3.0, 4.0, 0.0, 0.0]) - expected.add_(torch.softmax(logits[0], dim=0), alpha=2.0) - assert torch.equal(manager._h2o_scores[(0, 0)], expected) + event = PrefillScoreEvent(0, torch.empty(2, 2, 16), view, torch.tensor([0], dtype=torch.int32), torch.tensor([2], dtype=torch.int32), .25) + with patch('sparsevllm.kernels.triton.prefill_score.prefill_score_fwd', side_effect=produce_scores): + _score_runtime(manager).collect_prefill_attention_score(event) + expected = torch.nn.functional.pad(previous, (0, 2)) + increment + torch.testing.assert_close(manager._h2o_scores[(0, 0)], expected) def test_h2o_prefill_score_collection_rejects_misaligned_physical_view(): @@ -658,19 +564,17 @@ def test_h2o_prefill_score_collection_rejects_misaligned_physical_view(): set_context(is_prefill=True, cache_manager=manager, seqs=[seq]) with pytest.raises(RuntimeError, match="compressed physical coordinates"): - manager.collect_prefill_attention_score( - 0, - torch.empty((2, 1, 1)), - view, - b_start_loc=torch.tensor([0], dtype=torch.int32), - chunk_lens=torch.tensor([2], dtype=torch.int32), - ) + _score_runtime(manager).collect_prefill_attention_score(PrefillScoreEvent( + 0, torch.empty((2, 1, 1)), view, + torch.tensor([0], dtype=torch.int32), + torch.tensor([2], dtype=torch.int32), 1., + )) def test_h2o_missing_score_for_existing_physical_prefix_fails_fast(): manager = _manager_with_rows([8]) seq = _seq(0, 100, prefilled=64, chunk=3) - with pytest.raises(RuntimeError, match="score vector is missing"): + with pytest.raises(RuntimeError, match="not aligned"): manager._require_score_length(0, seq, 8) @@ -2290,3 +2194,51 @@ def test_h2o_debug_summary_exposes_auditable_eviction_counters(): assert summary["h2o"]["counters"]["dropped_tokens"] == 0 assert summary["h2o"]["counters"]["decode_eviction_bursts"] == 0 assert summary["h2o"]["score_lengths"] == {"0:0": 2} + + +def test_h2o_runtime_tiles_full_chunk_without_reducing_or_reweighting_heads(): + manager = _manager_with_rows([0]) + manager.config.h2o_prefill_score_window = 0 + # No physical payload is read by this producer mock; the oracle describes + # normalized probability mass, and tests runtime tiling/accumulation only. + manager.row_seq_lens[0][0] = 264 + manager._h2o_scores[(0, 0)] = torch.ones(2, 5) + seq = _seq(0, 300, prefilled=5, chunk=259) + view = PrefillComputeView( + meta=AttentionViewMeta( + active_slots=torch.zeros(1, 264, dtype=torch.int32), + req_indices=torch.tensor([0], dtype=torch.int32), + context_lens=torch.tensor([264], dtype=torch.int32), max_context_len=264, + ), + payload=ExplicitKVPayload(k_cache=torch.empty(264, 2, 16), v_cache=torch.empty(264, 2, 16)), + ) + set_context(is_prefill=True, cache_manager=manager, seqs=[seq]) + queries = torch.arange(5, 264)[:, None] + keys = torch.arange(264)[None, :] + logits = torch.stack([keys.expand(259, -1).float() * .01, -keys.expand(259, -1).float() * .01]) + probabilities = logits.masked_fill((keys > queries)[None], -torch.inf).softmax(-1) + ranges = [] + + def produce(q, k, output, reqs, starts, contexts, prefixes, max_queries, slots, begin, end, **kwargs): + first, last = int(begin[0]), int(end[0]) + ranges.append((first, last)) + output.zero_() + output[0].copy_(probabilities[:, first - 5:last - 5].sum(1)) + + event = PrefillScoreEvent(0, torch.empty(259, 2, 16), view, torch.tensor([0], dtype=torch.int32), torch.tensor([259], dtype=torch.int32), .25) + with patch('sparsevllm.kernels.triton.prefill_score.prefill_score_fwd', side_effect=produce): + _score_runtime(manager).collect_prefill_attention_score(event) + expected = probabilities.sum(1) + expected[:, :5] += 1 + torch.testing.assert_close(manager._h2o_scores[(0, 0)], expected) + assert ranges[0][0] == 5 and ranges[-1][1] == 264 + assert all(left[1] == right[0] for left, right in zip(ranges, ranges[1:])) + + +def test_h2o_failed_prefill_free_releases_positions_before_scores_exist(): + manager = _manager_with_rows([2]) + manager._h2o_positions[(0, 0)] = torch.arange(2)[None].expand(2, -1) + with patch.object(SnapKVCacheManager, 'free_seq', autospec=True): + manager.free_seq(0) + assert manager._h2o_positions == {} + assert manager._h2o_scores == {} diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py new file mode 100644 index 00000000..8dc9216a --- /dev/null +++ b/tests/test_h2o_per_head_prefill.py @@ -0,0 +1,195 @@ +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from sparsevllm.engine.cache_manager.h2o import H2OCacheManager +from sparsevllm.engine.cache_manager.h2o_retention import H2ORetention +from sparsevllm.engine.cache_manager.storage import ExplicitKVStorage, MlaLatentStorage +from sparsevllm.engine.sparse_methods.h2o_selection import select_h2o_heads + + +def test_late_max_preserves_head_history_across_chunks(): + # Alternating heads favor token 0 under early max; cumulative max favors 1. + p = torch.tensor([ + [[.4, .3, .3], [0., .3, .7]], + [[0., .3, .7], [.4, .3, .3]], + ]) + cumulative = None + for chunk in p: + cumulative = H2OCacheManager._accumulate_score( + cumulative, chunk, new_len=3, weight=1, + ) + torch.testing.assert_close(cumulative, p.sum(0)) + keep = select_h2o_heads(cumulative, selection_groups=1, budget=2, recent_ratio=.5, reduction='max') + assert keep.tolist() == [[1, 2]] + assert p.amax(1).sum(0)[0] > p.amax(1).sum(0)[1] + + +def test_groups_select_independently_and_mean_changes_only_ranking(): + cumulative = torch.tensor([[9., 6., 0., 0.], [0., 6., 0., 0.], [0., 0., 8., 0.], [0., 0., 7., 0.]]) + args = dict(selection_groups=2, budget=2, recent_ratio=.5) + assert select_h2o_heads(cumulative, reduction='max', **args).tolist() == [[0, 3], [2, 3]] + assert select_h2o_heads(cumulative, reduction='mean', **args).tolist() == [[1, 3], [2, 3]] + mha = select_h2o_heads(cumulative, selection_groups=4, budget=2, recent_ratio=.5, reduction='max') + assert mha.tolist() == [[0, 3], [1, 3], [2, 3], [2, 3]] + + +def test_ties_and_recent_suffix_fill_budget_in_logical_order(): + scores = torch.zeros(3, 9) + keep = select_h2o_heads(scores, selection_groups=3, budget=4, recent_ratio=.5, reduction='max') + assert keep.tolist() == [[0, 1, 7, 8]] * 3 + all_tokens = select_h2o_heads(scores, selection_groups=3, budget=12, recent_ratio=.5, reduction='mean') + assert all_tokens.tolist() == [list(range(9))] * 3 + + +def manager_with_storage(mla=False): + manager = object.__new__(H2OCacheManager) + manager.device = torch.device('cpu') + manager.num_kv_heads = 2 + manager.runtime_layout = SimpleNamespace(kv_layer_index=lambda layer: layer) + if mla: + storage = MlaLatentStorage(kv_lora_rank=512, rope_dim=64, dtype=torch.bfloat16) + else: + storage = ExplicitKVStorage(num_kv_heads=2, head_dim=4, dtype=torch.float32) + storage.allocate(num_layers=1, num_slots=12, device=manager.device) + manager.attention_cache_storage = storage + manager.seq_id_to_row = [{7: 0}] + manager.row_seq_lens = [np.array([6], dtype=np.int32)] + manager.buffer_req_to_token_slots = [torch.tensor([[8, 2, 6, 4, 0, 9, 0, 0]], dtype=torch.int32)] + manager.free_slots_stack = [torch.zeros(12, dtype=torch.int32)] + manager._num_free_slots = [6] + manager._h2o_scores = {(0, 7): torch.arange(24).reshape(4, 6).float()} + manager._h2o_positions = {(0, 7): torch.tensor([[1, 3, 5, 7, 9, 11]]).expand(1 if mla else 2, -1).clone()} + manager._h2o_counters = dict(final_prefill_evictions=0, intermediate_prefill_evictions=0, dropped_tokens=0) + return manager + + +@pytest.mark.parametrize('final', [False, True]) +def test_explicit_retention_packs_different_heads_without_union(final): + manager = manager_with_storage() + storage = manager.attention_cache_storage + storage.cache.copy_(torch.arange(storage.cache.numel()).reshape_as(storage.cache)) + before = storage.cache.clone() + old_slots = manager.buffer_req_to_token_slots[0][0, :6].long().clone() + old_scores = manager._h2o_scores[(0, 7)].clone() + keep = torch.tensor([[0, 2, 5], [1, 4, 5]]) + manager.commit_h2o_retention([H2ORetention(0, 7, 6, keep, final)]) + destinations = manager.buffer_req_to_token_slots[0][0, :3].long() + # Independent scalar oracle: the packed slot can contain different tokens. + for head in range(2): + for packed in range(3): + source = old_slots[keep[head, packed]] + torch.testing.assert_close(storage.cache[:, 0, destinations[packed], head], before[:, 0, source, head]) + for query_head in range(2 * head, 2 * head + 2): + assert manager._h2o_scores[(0, 7)][query_head, packed] == old_scores[query_head, keep[head, packed]] + assert manager.row_seq_lens[0][0] == 3 + assert manager._num_free_slots == [9] + assert storage.cache.shape == before.shape + assert manager._h2o_positions[(0, 7)].tolist() == [[1, 5, 11], [3, 9, 11]] + released = set(manager.free_slots_stack[0][6:9].tolist()) + assert released.isdisjoint(destinations.tolist()) + assert released | set(destinations.tolist()) == set(old_slots.tolist()) + + +def test_mla_retention_keeps_latent_rope_and_all_score_heads_paired(): + manager = manager_with_storage(mla=True) + storage = manager.attention_cache_storage + for slot in range(12): + storage.latent_cache[:, slot].fill_(slot) + storage.rope_cache[:, slot].fill_(slot + 20) + score_before = manager._h2o_scores[(0, 7)].clone() + manager.commit_h2o_retention([H2ORetention(0, 7, 6, torch.tensor([[1, 3, 5]]), True)]) + for dest, source in zip(manager.buffer_req_to_token_slots[0][0, :3], [2, 4, 9]): + assert (storage.latent_cache[0, dest] == source).all() + assert (storage.rope_cache[0, dest] == source + 20).all() + torch.testing.assert_close(manager._h2o_scores[(0, 7)], score_before[:, [1, 3, 5]]) + assert manager._h2o_positions[(0, 7)].tolist() == [[3, 7, 11]] + + +@pytest.mark.parametrize('keep', [torch.tensor([[0, 0, 5], [1, 4, 5]]), torch.tensor([[0, 2, 6], [1, 4, 5]])]) +def test_invalid_retention_does_not_mutate_cache_or_allocator(keep): + manager = manager_with_storage() + before = manager.attention_cache_storage.cache.clone() + slots = manager.buffer_req_to_token_slots[0].clone() + with pytest.raises(RuntimeError, match='indices'): + manager.commit_h2o_retention([H2ORetention(0, 7, 6, keep, True)]) + torch.testing.assert_close(manager.attention_cache_storage.cache, before, equal_nan=True) + torch.testing.assert_close(manager.buffer_req_to_token_slots[0], slots) + assert manager.row_seq_lens[0][0] == 6 + assert manager._num_free_slots == [6] + + +def test_chunk_prefill_append_retention_and_final_handoff_follow_logical_positions(): + from sparsevllm.engine.cache_manager.base import LayerBatchStates + from sparsevllm.engine.sequence import Sequence + from sparsevllm.engine.sparse_methods.base import SparseStepContext + from sparsevllm.engine.sparse_methods.h2o import H2ORuntime + + manager = manager_with_storage() + manager.num_layers = manager.num_kv_layers = 1 + manager.runtime_layout.kv_idx_to_layer_idx = (0,) + manager.max_model_len = 16 + manager.layer_batch_states = [LayerBatchStates()] + manager.buffer_req_to_token_slots = [torch.zeros(1, 16, dtype=torch.int32)] + manager.free_slots_stack = [torch.arange(12, dtype=torch.int32)] + manager._num_free_slots = [12] + manager.row_seq_lens[0][0] = 0 + manager._h2o_scores.clear() + manager._h2o_positions.clear() + manager.config = SimpleNamespace(h2o_prefill_budget=4, h2o_decode_budget=3, + h2o_prefill_score_window=0, h2o_recent_ratio=.5, + h2o_head_reduction='max') + runtime = object.__new__(H2ORuntime) + runtime.config = manager.config + runtime.cache_manager = manager + seq = Sequence(list(range(9))) + seq.seq_id = 7 + histories = [[], []] + cumulative_reference = [dict() for _ in range(4)] + for logical_start in (0, 3, 6): + seq.num_prefilled_tokens = logical_start + seq.current_chunk_size = 3 + _, positions, _ = manager._prepare_prefill([seq]) + assert positions.tolist() == list(range(logical_start, logical_start + 3)) + slots = manager.layer_batch_states[0].slot_mapping.long() + for index, slot in enumerate(slots): + for head in range(2): + manager.attention_cache_storage.cache[:, 0, slot, head].fill_(logical_start + index + head * 100) + for history in histories: + history.extend(range(logical_start, logical_start + 3)) + length = len(histories[0]) + delta = torch.zeros(1, 4, length) + for head in range(4): + history = histories[head // 2] + # Independently generated normalized causal probabilities; this CPU + # test establishes state/coordinate correctness, not GPU numerics. + for query_pos in range(logical_start, logical_start + 3): + visible = [token for token in history if token <= query_pos] + weights = [float(1 + ((token + head * 3) % 7)) for token in visible] + denominator = sum(weights) + for token, weight in zip(visible, weights): + mass = weight / denominator + delta[0, head, history.index(token)] += mass + cumulative_reference[head][token] = cumulative_reference[head].get(token, 0) + mass + manager.accumulate_prefill_scores(0, [seq], delta) + runtime.finish_step(SparseStepContext([seq], True, None)) + budget = 3 if logical_start == 6 else 4 + for group in range(2): + history = histories[group] + if len(history) > budget: + recent_count = max(1, int(budget * .5)) + rank = lambda token: max(cumulative_reference[2 * group + h][token] for h in range(2)) + heavy = sorted(history[:-recent_count], key=lambda token: (-rank(token), token))[:budget - recent_count] + histories[group] = sorted(heavy + history[-recent_count:]) + for head in (2 * group, 2 * group + 1): + cumulative_reference[head] = {token: cumulative_reference[head][token] for token in histories[group]} + torch.testing.assert_close(manager._h2o_scores[(0, 7)][head], torch.tensor(list(cumulative_reference[head].values()))) + assert manager._h2o_positions[(0, 7)].tolist() == histories + resident = int(manager.row_seq_lens[0][0]) + assert resident == min(length, budget) + for packed, slot in enumerate(manager.buffer_req_to_token_slots[0][0, :resident]): + for head in range(2): + assert (manager.attention_cache_storage.cache[:, 0, slot, head] == histories[head][packed] + 100 * head).all() + assert manager._num_free_slots[0] + resident == 12 diff --git a/tests/test_h2o_prefill_probability_kernel.py b/tests/test_h2o_prefill_probability_kernel.py new file mode 100644 index 00000000..f98dd3d1 --- /dev/null +++ b/tests/test_h2o_prefill_probability_kernel.py @@ -0,0 +1,83 @@ +"""CUDA numerical checks; run on an idle device against independent attention.""" + +import pytest +import torch + +from sparsevllm.kernels.triton.prefill_score import ( + prefill_score_fwd, prefill_score_from_lse_fwd, +) + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason='requires CUDA') + + +@pytest.mark.parametrize('heads,kv_heads,dim', [(4, 4, 32), (8, 2, 128), (4, 4, 192)]) +@pytest.mark.parametrize('query_lengths', [(7, 3), (137, 19)]) +@pytest.mark.parametrize('use_lse', [False, True]) +def test_head_probability_sums_match_causal_attention(heads, kv_heads, dim, query_lengths, use_lse): + # MLA uses already projected explicit K for scoring; dimension 192 exercises + # its non-power-of-two projection, with no persistent KV expansion here. + torch.manual_seed(91) + device = torch.device('cuda') + prefixes = [11, 5] + contexts = [p + q for p, q in zip(prefixes, query_lengths)] + width = max(contexts) + 9 + q = torch.randn(sum(query_lengths), heads, dim, device=device, dtype=torch.bfloat16) + k = torch.randn(2 * width, kv_heads, dim, device=device, dtype=torch.bfloat16) + slots = torch.randperm(2 * width, device=device).reshape(2, width).int() + starts = torch.tensor([0, query_lengths[0]], device=device, dtype=torch.int32) + context_tensor = torch.tensor(contexts, device=device, dtype=torch.int32) + prefix_tensor = torch.tensor(prefixes, device=device, dtype=torch.int32) + reqs = torch.arange(2, device=device, dtype=torch.int32) + expected = torch.zeros(2, heads, width, device=device) + lse = torch.empty(heads, sum(query_lengths), device=device) + scale = .07 # Verifies the actual attention scale is carried through. + for batch, (prefix, query_len, context) in enumerate(zip(prefixes, query_lengths, contexts)): + q_start = int(starts[batch]) + for head in range(heads): + keys = k[slots[batch, :context].long(), head // (heads // kv_heads)].float() + logits = q[q_start:q_start + query_len, head].float() @ keys.T * scale + mask = torch.arange(context, device=device)[None, :] > (prefix + torch.arange(query_len, device=device))[:, None] + logits.masked_fill_(mask, -torch.inf) + expected[batch, head, :context] = logits.softmax(-1).sum(0) + lse[head, q_start:q_start + query_len] = logits.logsumexp(-1) + # A strided output checks the public wrapper, not just contiguous kernels. + backing = torch.full((2, heads, width * 2), float('nan'), device=device) + output = backing[..., ::2] + args = (output, reqs, starts, context_tensor, prefix_tensor, max(query_lengths), slots, prefix_tensor, context_tensor) + for _ in range(2): + if use_lse: + prefill_score_from_lse_fwd(q, k, lse, *args, per_head=True, softmax_scale=scale) + else: + prefill_score_fwd(q, k, *args, per_head=True, softmax_scale=scale) + torch.testing.assert_close(output, expected, rtol=5e-3, atol=5e-3) + torch.testing.assert_close(output.sum(-1), torch.tensor(query_lengths, device=device).float()[:, None].expand(-1, heads), rtol=5e-3, atol=5e-3) + assert torch.isnan(backing[..., 1::2]).all() + + +def test_mla_score_uses_projected_nonrope_and_rope_logits(): + torch.manual_seed(27) + device = torch.device('cuda') + length, heads, nope_dim, rope_dim, latent_dim = 17, 4, 128, 64, 512 + latent = torch.randn(length, latent_dim, device=device, dtype=torch.bfloat16) + rope = torch.randn(length, rope_dim, device=device, dtype=torch.bfloat16) + projection = torch.randn(latent_dim, heads, nope_dim, device=device, dtype=torch.bfloat16) * .04 + expanded_nope = (latent @ projection.flatten(1)).reshape(length, heads, nope_dim) + q_nope = torch.randn(length, heads, nope_dim, device=device, dtype=torch.bfloat16) + q_rope = torch.randn(length, heads, rope_dim, device=device, dtype=torch.bfloat16) + q = torch.cat((q_nope, q_rope), -1) + k = torch.cat((expanded_nope, rope[:, None, :].expand(-1, heads, -1)), -1) + expected = torch.empty(heads, length, device=device) + scale = (nope_dim + rope_dim) ** -.5 + causal_mask = torch.arange(length, device=device)[None, :] > torch.arange(length, device=device)[:, None] + for head in range(heads): + # Independent decomposition of the two MLA attention terms. + logits = (q_nope[:, head].float() @ expanded_nope[:, head].float().T + + q_rope[:, head].float() @ rope.float().T) * scale + expected[head] = logits.masked_fill(causal_mask, -torch.inf).softmax(-1).sum(0) + output = torch.empty(1, heads, length, device=device) + zero = torch.zeros(1, device=device, dtype=torch.int32) + end = torch.tensor([length], device=device, dtype=torch.int32) + slots = torch.arange(length, device=device, dtype=torch.int32)[None] + prefill_score_fwd(q, k, output, zero, zero, end, zero, length, slots, zero, end, + per_head=True, softmax_scale=scale) + torch.testing.assert_close(output[0], expected, rtol=5e-3, atol=5e-3) diff --git a/tests/test_prefill_attention_provider.py b/tests/test_prefill_attention_provider.py index 103e3498..2b0a0b00 100644 --- a/tests/test_prefill_attention_provider.py +++ b/tests/test_prefill_attention_provider.py @@ -168,7 +168,7 @@ def _mock_flashinfer_paged_prefill_contract(): ("", AttentionScoreKind.NONE, PrefillScoreCollectionKind.NONE), ("snapkv", AttentionScoreKind.NONE, PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_REDUCED), ("pyramidkv", AttentionScoreKind.NONE, PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_REDUCED), - ("h2o", AttentionScoreKind.NONE, PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_REDUCED), + ("h2o", AttentionScoreKind.NONE, PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_PER_HEAD), ("rkv", AttentionScoreKind.NONE, PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_REDUCED), ("omnikv", AttentionScoreKind.NONE, PrefillScoreCollectionKind.NONE), ("deltakv", AttentionScoreKind.NONE, PrefillScoreCollectionKind.NONE), @@ -228,18 +228,9 @@ def test_sm120_dense_prefill_sparse_method_selects_flashinfer_fa2(method): assert resolved.report.selection_basis == "upstream_default" -def test_h2o_full_query_logits_request_fused_reduced_prefill_score(): - contract = sparse_prefill_attention_contract( - "h2o", - sparse_prefill_score_mode="logits", - h2o_prefill_score_window=0, - ) - - assert contract.main_score_kind is AttentionScoreKind.RAW_QK_REDUCED - assert ( - contract.score_collection - is PrefillScoreCollectionKind.MAIN_ATTENTION_REDUCED - ) +def test_h2o_rejects_reduced_logits_before_provider_binding(): + with pytest.raises(ValueError, match="per-head probability"): + sparse_prefill_attention_contract("h2o", sparse_prefill_score_mode="logits") @pytest.mark.parametrize( @@ -300,14 +291,14 @@ def test_h2o_flashprefill_uses_method_owned_posthoc_prefill_scoring(): contract = sparse_prefill_attention_contract( "h2o", prefill_sparse_method="flashprefill_v2", - sparse_prefill_score_mode="logits", + sparse_prefill_score_mode="probability", h2o_prefill_score_window=0, ) assert contract.main_score_kind is AttentionScoreKind.NONE assert ( contract.score_collection - is PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_REDUCED + is PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_PER_HEAD ) @@ -919,6 +910,7 @@ def run_sgl(*args, **_kwargs): ) sparse_controller = SimpleNamespace( get_prefill_selection=Mock(return_value=object()), + collect_prefill_attention_score=cache_manager.collect_prefill_attention_score, on_layer_attention_end=Mock(), ) context = SimpleNamespace( From 517db1a45fbe0f0dc34e495fd4ff61e32969c40c Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 17:45:26 +0800 Subject: [PATCH 02/11] refactor: retire shared h2o prefill eviction paths --- src/sparsevllm/engine/cache_manager/h2o.py | 512 +----------------- .../engine/cache_manager/h2o_retention.py | 18 +- tests/test_glm_mla_prefix_cache.py | 25 +- tests/test_h2o_cache_manager.py | 344 +----------- tests/test_h2o_per_head_prefill.py | 91 ++++ tests/test_static_eviction_compaction.py | 36 -- 6 files changed, 122 insertions(+), 904 deletions(-) diff --git a/src/sparsevllm/engine/cache_manager/h2o.py b/src/sparsevllm/engine/cache_manager/h2o.py index 77df5b3a..087c3057 100644 --- a/src/sparsevllm/engine/cache_manager/h2o.py +++ b/src/sparsevllm/engine/cache_manager/h2o.py @@ -15,10 +15,8 @@ from sparsevllm.utils.context import get_context from sparsevllm.utils.profiler import profiler -from .base import ExplicitKVPayload, PrefillComputeView from .snapkv import SnapKVCacheManager from .h2o_retention import H2OPrefillRetentionMixin, H2O_PREFILL_QUERY_TILE -from .storage import ExplicitKVStorage class _H2ORowRef(NamedTuple): @@ -58,7 +56,6 @@ def __init__( "decode_evictions": 0, "dropped_tokens": 0, } - self._h2o_final_prefill_workspace: torch.Tensor | None = None self._h2o_decode_score_workspace: torch.Tensor | None = None self._h2o_decode_score_signature: tuple[tuple[int, ...], tuple[int, ...]] | None = None self._h2o_decode_score_length = 0 @@ -668,35 +665,6 @@ def _accumulate_score( cumulative.add_(step_score[..., :new_len].float(), alpha=float(weight)) return cumulative - @staticmethod - def _normalize_logit_prefill_score( - step_score: torch.Tensor, - *, - new_len: int, - ) -> torch.Tensor: - """Normalize max-reduced prefill logits to step token probabilities. - - Unscored positions retain -inf and map to zero probability under - softmax. - """ - if step_score.dim() != 1 or int(step_score.numel()) < int(new_len): - raise ValueError( - "H2O logit prefill score must be a 1D vector covering new_len: " - f"shape={tuple(step_score.shape)} new_len={int(new_len)}." - ) - logits = step_score[: int(new_len)].float() - has_invalid = (torch.isnan(logits) | (logits == torch.inf)).any() - has_finite = torch.isfinite(logits).any() - valid = (~has_invalid) & has_finite - if valid.is_cuda: - torch._assert_async(valid) - elif not bool(valid.item()): - raise RuntimeError( - "H2O logit prefill score contains invalid non-finite values (NaN or +inf) " - "or lacks any finite score." - ) - return torch.softmax(logits, dim=0) - def _physical_row_len(self, layer_idx: int, seq: Sequence) -> int: row_idx = self.seq_id_to_row[layer_idx].get(int(seq.seq_id)) if row_idx is None: @@ -1113,473 +1081,6 @@ def free_part_slots( ) self._h2o_scores[self._score_key(layer_idx, seq.seq_id)] = kept_score - def _get_final_prefill_workspace( - self, - *, - batch_size: int, - budget: int, - k_cache: torch.Tensor, - v_cache: torch.Tensor, - ) -> torch.Tensor: - if batch_size <= 0 or budget <= 0: - raise RuntimeError( - "H2O final-prefill workspace requires positive batch and budget: " - f"batch={batch_size} budget={budget}." - ) - if not isinstance(k_cache, torch.Tensor) or not isinstance(v_cache, torch.Tensor): - raise TypeError("H2O final-prefill dense compaction requires tensor K/V caches.") - if k_cache.dim() != 3 or v_cache.dim() != 3 or tuple(k_cache.shape) != tuple(v_cache.shape): - raise RuntimeError( - "H2O final-prefill dense compaction requires matching [slots, heads, dim] caches: " - f"k_shape={tuple(k_cache.shape)} v_shape={tuple(v_cache.shape)}." - ) - if k_cache.dtype != v_cache.dtype or k_cache.device != v_cache.device: - raise RuntimeError( - "H2O final-prefill dense compaction requires matching K/V dtype and device: " - f"k={k_cache.dtype}/{k_cache.device} v={v_cache.dtype}/{v_cache.device}." - ) - - required_shape = ( - 2, - int(batch_size), - int(budget), - int(k_cache.shape[1]), - int(k_cache.shape[2]), - ) - workspace = getattr(self, "_h2o_final_prefill_workspace", None) - needs_allocation = ( - workspace is None - or workspace.dtype != k_cache.dtype - or workspace.device != k_cache.device - or int(workspace.shape[1]) < int(batch_size) - or tuple(workspace.shape[2:]) != required_shape[2:] - ) - if needs_allocation: - workspace = torch.empty( - required_shape, - dtype=k_cache.dtype, - device=k_cache.device, - ) - self._h2o_final_prefill_workspace = workspace - return workspace[:, :batch_size] - - @staticmethod - def _assert_final_prefill_tensor(condition: torch.Tensor, message: str) -> None: - if condition.numel() != 1: - raise RuntimeError( - "H2O final-prefill validation must reduce to one boolean: " - f"shape={tuple(condition.shape)} message={message}." - ) - if condition.is_cuda: - torch._assert_async(condition) - elif not bool(condition.item()): - raise RuntimeError(message) - - def _preflight_final_prefill_dense_capacity( - self, - seqs: list[Sequence], - ) -> None: - """Validate every final-prefill free-stack update before moving any KV.""" - final_seqs = [seq for seq in seqs if bool(seq.is_last_chunk_prefill)] - if not final_seqs: - return - seq_ids = [int(seq.seq_id) for seq in final_seqs] - if len(seq_ids) != len(set(seq_ids)): - raise RuntimeError( - "H2O final-prefill capacity preflight received duplicate seq ids: " - f"{seq_ids}." - ) - - budget = self.h2o_decode_budget - for layer_idx in self.kv_transformer_layer_indices(): - row_indices = [] - release_count = 0 - for seq in final_seqs: - row_idx = self.seq_id_to_row[layer_idx].get(int(seq.seq_id)) - if row_idx is None: - raise RuntimeError( - "H2O final-prefill capacity preflight is missing a physical row: " - f"layer={layer_idx} seq_id={int(seq.seq_id)}." - ) - row_indices.append(int(row_idx)) - release_count += max( - 0, - int(self.row_seq_lens[layer_idx][row_idx]) - budget, - ) - if len(row_indices) != len(set(row_indices)): - raise RuntimeError( - "H2O final-prefill capacity preflight received duplicate physical rows: " - f"layer={layer_idx} rows={row_indices}." - ) - if release_count == 0: - continue - - free_stack = self.free_slots_stack[layer_idx] - if free_stack is None or free_stack.dim() != 1: - raise RuntimeError( - "H2O final-prefill dense compaction requires a one-dimensional " - f"free-slot stack: layer={layer_idx}." - ) - free_ptr = int(self._num_free_slots[layer_idx]) - if free_ptr < 0 or free_ptr + release_count > int(free_stack.numel()): - raise RuntimeError( - "H2O final-prefill dense compaction would overflow the free-slot " - f"stack: layer={layer_idx} ptr={free_ptr} release={release_count} " - f"capacity={int(free_stack.numel())}." - ) - - @torch.no_grad() - def _compact_final_prefill_dense_batch( - self, - layer_idx: int, - seqs: list[Sequence], - keep_indices: torch.Tensor, - ) -> None: - """Move final H2O selections into ascending physical destination slots.""" - if not seqs: - raise RuntimeError("H2O final-prefill dense compaction requires sequences.") - kv_idx = self.kv_layer_index(layer_idx) - budget = self.h2o_decode_budget - batch_size = len(seqs) - keep_indices = keep_indices.to( - device=self.device, - dtype=torch.long, - ).contiguous() - if keep_indices.dim() != 2 or tuple(keep_indices.shape) != (batch_size, budget): - raise RuntimeError( - "H2O final-prefill keep indices must have shape [batch, decode_budget]: " - f"expected={(batch_size, budget)} got={tuple(keep_indices.shape)}." - ) - - seq_ids = [int(seq.seq_id) for seq in seqs] - if len(seq_ids) != len(set(seq_ids)): - raise RuntimeError( - f"H2O final-prefill dense compaction received duplicate seq ids: {seq_ids}." - ) - row_indices = [] - cur_lens = [] - for seq_id in seq_ids: - row_idx = self.seq_id_to_row[layer_idx].get(seq_id) - if row_idx is None: - raise RuntimeError( - "H2O final-prefill dense compaction is missing a physical row: " - f"layer={layer_idx} seq_id={seq_id}." - ) - row_indices.append(int(row_idx)) - cur_lens.append(int(self.row_seq_lens[layer_idx][row_idx])) - if len(row_indices) != len(set(row_indices)): - raise RuntimeError( - "H2O final-prefill dense compaction received duplicate physical rows: " - f"layer={layer_idx} rows={row_indices}." - ) - kv_len = int(cur_lens[0]) - if any(int(length) != kv_len for length in cur_lens[1:]): - raise RuntimeError( - "H2O final-prefill dense batch requires uniform physical lengths; " - "nonuniform callers must use batch-one compaction: " - f"layer={layer_idx} lengths={cur_lens}." - ) - if kv_len <= budget: - raise RuntimeError( - "H2O final-prefill dense compaction requires an over-budget row: " - f"layer={layer_idx} kv_len={kv_len} budget={budget}." - ) - - free_count = (kv_len - budget) * batch_size - free_ptr = int(self._num_free_slots[layer_idx]) - free_stack = self.free_slots_stack[layer_idx] - if free_stack is None or free_stack.dim() != 1: - raise RuntimeError( - "H2O final-prefill dense compaction requires a one-dimensional free-slot stack: " - f"layer={layer_idx}." - ) - if free_ptr < 0 or free_ptr + free_count > int(free_stack.numel()): - raise RuntimeError( - "H2O final-prefill dense compaction would overflow the free-slot stack: " - f"layer={layer_idx} ptr={free_ptr} release={free_count} " - f"capacity={int(free_stack.numel())}." - ) - - storage = getattr(self, "attention_cache_storage", None) - uses_explicit_kv = storage is None or isinstance(storage, ExplicitKVStorage) - if uses_explicit_kv: - k_cache, v_cache = self.get_layer_kv_cache(layer_idx) - workspace = self._get_final_prefill_workspace( - batch_size=batch_size, - budget=budget, - k_cache=k_cache, - v_cache=v_cache, - ) - slot_capacity = int(k_cache.shape[0]) - else: - slot_capacity = storage.slot_capacity() - rows_gpu = torch.tensor(row_indices, dtype=torch.long, device=self.device) - old_slots = self.buffer_req_to_token_slots[layer_idx][ - rows_gpu, :kv_len - ].clone() - self._assert_final_prefill_tensor( - ((keep_indices >= 0) & (keep_indices < kv_len)).all(), - "H2O final-prefill keep indices are out of bounds: " - f"layer={layer_idx} kv_len={kv_len}.", - ) - if budget > 1: - self._assert_final_prefill_tensor( - (keep_indices[:, 1:] > keep_indices[:, :-1]).all(), - "H2O final-prefill keep indices must be strictly increasing in logical order: " - f"layer={layer_idx}.", - ) - self._assert_final_prefill_tensor( - ((old_slots >= 0) & (old_slots < slot_capacity)).all(), - "H2O final-prefill slot map contains an out-of-range physical slot: " - f"layer={layer_idx} num_slots={slot_capacity}.", - ) - - globally_sorted_slots = torch.sort(old_slots.reshape(-1)).values - if int(globally_sorted_slots.numel()) > 1: - self._assert_final_prefill_tensor( - (globally_sorted_slots[1:] != globally_sorted_slots[:-1]).all(), - "H2O final-prefill active physical slots must be unique across the batch: " - f"layer={layer_idx} rows={row_indices}.", - ) - - selected_slots = old_slots.gather(1, keep_indices).to(torch.long) - sorted_old_slots = torch.sort(old_slots, dim=1).values - destination_slots = sorted_old_slots[:, :budget].contiguous() - released_slots = sorted_old_slots[:, budget:].reshape(-1).contiguous() - - selected_flat = selected_slots.reshape(-1) - destination_flat = destination_slots.reshape(-1).to(torch.long) - if uses_explicit_kv: - workspace[0].copy_( - k_cache.index_select(0, selected_flat).view( - batch_size, - budget, - int(k_cache.shape[1]), - int(k_cache.shape[2]), - ) - ) - workspace[1].copy_( - v_cache.index_select(0, selected_flat).view( - batch_size, - budget, - int(v_cache.shape[1]), - int(v_cache.shape[2]), - ) - ) - k_cache.index_copy_( - 0, - destination_flat, - workspace[0].reshape( - batch_size * budget, - int(k_cache.shape[1]), - int(k_cache.shape[2]), - ), - ) - v_cache.index_copy_( - 0, - destination_flat, - workspace[1].reshape( - batch_size * budget, - int(v_cache.shape[1]), - int(v_cache.shape[2]), - ), - ) - else: - storage.copy_slots(kv_idx, selected_flat, destination_flat) - - free_stack[free_ptr : free_ptr + free_count] = released_slots.to( - dtype=free_stack.dtype, - device=free_stack.device, - ) - self._num_free_slots[layer_idx] = free_ptr + free_count - self.buffer_req_to_token_slots[layer_idx][ - rows_gpu, :budget - ] = destination_slots.to(torch.int32) - self.buffer_req_to_token_slots[layer_idx][rows_gpu, budget:kv_len] = 0 - self.row_seq_lens[layer_idx][row_indices] = budget - self._uniform_decode_metadata = False - - def _evict(self, seqs: list[Sequence], *, is_prefill: bool): - if is_prefill: - self._preflight_final_prefill_dense_capacity(seqs) - if self._try_batched_evict(seqs, is_prefill=is_prefill): - return - ratio = float(self.config.h2o_recent_ratio) - for layer_idx in self.kv_transformer_layer_indices(): - for seq in seqs: - is_final_prefill = bool(is_prefill and seq.is_last_chunk_prefill) - budget = ( - self.h2o_decode_budget - if not is_prefill or is_final_prefill - else self.h2o_prefill_budget - ) - kv_len = self._physical_row_len(layer_idx, seq) - if kv_len <= budget: - continue - score = self._require_score_length(layer_idx, seq, kv_len) - keep_indices = self.select_h2o_indices( - score, - budget=budget, - recent_ratio=ratio, - ) - dropped = kv_len - int(keep_indices.numel()) - if is_final_prefill: - kept_score = score.index_select(0, keep_indices).contiguous() - self._compact_final_prefill_dense_batch( - layer_idx, - [seq], - keep_indices.unsqueeze(0), - ) - self._h2o_scores[self._score_key(layer_idx, seq.seq_id)] = kept_score - else: - self.free_part_slots( - layer_idx, - seq, - keep_indices, - keep_indices_sorted=True, - ) - if is_prefill: - counter = ( - "final_prefill_evictions" - if is_final_prefill - else "intermediate_prefill_evictions" - ) - else: - counter = "decode_evictions" - self._h2o_counters[counter] += 1 - self._h2o_counters["dropped_tokens"] += int(dropped) - - def _try_batched_evict(self, seqs: list[Sequence], *, is_prefill: bool) -> bool: - """Compact each layer across uniform sequence rows with one batch op.""" - if not seqs: - return False - layer_indices = list(self.kv_transformer_layer_indices()) - if not layer_indices: - return False - - final_flags = [bool(seq.is_last_chunk_prefill) for seq in seqs] if is_prefill else [] - if is_prefill and any(flag != final_flags[0] for flag in final_flags[1:]): - return False - budget = ( - self.h2o_decode_budget - if not is_prefill or final_flags[0] - else self.h2o_prefill_budget - ) - - layer_rows: list[tuple[int, int, list[torch.Tensor]]] = [] - for layer_idx in layer_indices: - physical_lens = [ - self._physical_row_len(layer_idx, seq) for seq in seqs - ] - kv_len = int(physical_lens[0]) - if any(int(length) != kv_len for length in physical_lens[1:]): - return False - if kv_len <= budget: - continue - score_rows = [ - self._require_score_length(layer_idx, seq, kv_len) for seq in seqs - ] - layer_rows.append((int(layer_idx), kv_len, score_rows)) - - compact_layers = [] - compact_indices = [] - compact_scores = [] - compact_lengths = [] - if layer_rows: - self._invalidate_h2o_decode_score_workspace() - for layer_idx, kv_len, score_rows in layer_rows: - with profiler.record("h2o_evict_score_stack"): - scores = torch.stack(score_rows, dim=0) - with profiler.record("h2o_evict_select"): - keep_indices = self.select_h2o_indices_batch( - scores, - budget=budget, - recent_ratio=float(self.config.h2o_recent_ratio), - ) - kept_scores = scores.gather(1, keep_indices) - compact_layers.append(layer_idx) - compact_indices.append(keep_indices) - compact_scores.append(kept_scores) - compact_lengths.append(kv_len) - - if compact_layers: - # Each layer supplies its own keep indices. Page-table compaction - # avoids copying retained K/V rows into a dense destination range. - with profiler.record("h2o_evict_compact"): - SnapKVCacheManager.free_part_slots_batch_layers( - self, - compact_layers, - seqs, - torch.stack(compact_indices, dim=0), - keep_indices_sorted=True, - ) - for local_layer, layer_idx in enumerate(compact_layers): - for batch_idx, seq in enumerate(seqs): - self._h2o_scores[ - self._score_key(layer_idx, seq.seq_id) - ] = compact_scores[local_layer][batch_idx] - - num_evicted_rows = len(compact_layers) * len(seqs) - dropped_tokens = sum( - (kv_len - budget) * len(seqs) for kv_len in compact_lengths - ) - - if is_prefill: - counter = ( - "final_prefill_evictions" - if final_flags[0] - else "intermediate_prefill_evictions" - ) - else: - counter = "decode_evictions" - self._h2o_counters[counter] += int(num_evicted_rows) - self._h2o_counters["dropped_tokens"] += int(dropped_tokens) - return True - - def evict_after_intermediate_prefill(self, seqs: list[Sequence]) -> None: - """Apply the prefill method only between prompt chunks.""" - - intermediate = [seq for seq in seqs if not seq.is_last_chunk_prefill] - if intermediate: - self._evict(intermediate, is_prefill=True) - - def compact_final_prefill_for_decode(self, seqs: list[Sequence]) -> None: - """Apply the decode method at the final prompt boundary.""" - - final = [seq for seq in seqs if seq.is_last_chunk_prefill] - if not final: - return - self._evict(final, is_prefill=True) - for layer_idx in self.kv_transformer_layer_indices(): - for seq in final: - kv_len = self._physical_row_len(layer_idx, seq) - if kv_len > self.h2o_decode_budget: - raise RuntimeError( - "H2O final prefill did not compact to the decode budget: " - f"layer={layer_idx} seq_id={seq.seq_id} " - f"kv_len={kv_len} budget={self.h2o_decode_budget}." - ) - if self.num_free_slots <= 0: - self._evict_decode_rows([]) - - def evict_after_prefill(self, seqs: list[Sequence]) -> None: - """Compatibility wrapper for the legacy combined H2O lifecycle.""" - - self._evict(seqs, is_prefill=True) - for layer_idx in self.kv_transformer_layer_indices(): - for seq in seqs: - if not seq.is_last_chunk_prefill: - continue - kv_len = self._physical_row_len(layer_idx, seq) - if kv_len > self.h2o_decode_budget: - raise RuntimeError( - "H2O final prefill did not compact to the decode budget: " - f"layer={layer_idx} seq_id={seq.seq_id} " - f"kv_len={kv_len} budget={self.h2o_decode_budget}." - ) - if self.num_free_slots <= 0: - self._evict_decode_rows([]) - def _decode_eviction_groups( self, seqs: list[Sequence], @@ -1778,22 +1279,13 @@ def reset_after_warmup(self) -> None: def debug_state_summary(self) -> dict[str, object]: summary = super().debug_state_summary() - workspace = getattr(self, "_h2o_final_prefill_workspace", None) summary["h2o"] = { "counters": dict(self._h2o_counters), "score_lengths": { f"{layer_idx}:{seq_id}": int(score.shape[-1]) for (layer_idx, seq_id), score in sorted(self._h2o_scores.items()) }, - "final_prefill_workspace": ( - None - if workspace is None - else { - "shape": list(workspace.shape), - "dtype": str(workspace.dtype), - "device": str(workspace.device), - "nbytes": int(workspace.untyped_storage().nbytes()), - } - ), + "reserved_workspace_bytes": getattr(self, "_h2o_reserved_workspace_bytes", 0), + "metadata_bytes_per_slot": getattr(self, "_h2o_metadata_bytes_per_slot", 0), } return summary diff --git a/src/sparsevllm/engine/cache_manager/h2o_retention.py b/src/sparsevllm/engine/cache_manager/h2o_retention.py index 5a3e4887..19ad2134 100644 --- a/src/sparsevllm/engine/cache_manager/h2o_retention.py +++ b/src/sparsevllm/engine/cache_manager/h2o_retention.py @@ -22,6 +22,13 @@ class H2ORetention: class H2OPrefillRetentionMixin: """Cache-owned per-head state and overlap-safe physical retention.""" + @staticmethod + def _assert_retention_tensor(condition: torch.Tensor, message: str) -> None: + if condition.is_cuda: + torch._assert_async(condition) + elif not bool(condition.item()): + raise RuntimeError(message) + @property def h2o_selection_groups(self) -> int: from .storage import MlaLatentStorage @@ -60,7 +67,7 @@ def commit_h2o_retention(self, requests: list[H2ORetention]) -> None: ): raise ValueError("H2O retention requires [selection_groups, budget] int64 indices.") budget = int(keep.shape[-1]) - self._assert_final_prefill_tensor( + self._assert_retention_tensor( ((keep >= 0) & (keep < length)).all() & (keep[:, 1:] > keep[:, :-1]).all(), "H2O retention indices must be in bounds and strictly increasing.", @@ -74,15 +81,16 @@ def commit_h2o_retention(self, requests: list[H2ORetention]) -> None: raise ValueError("H2O retention score/position metadata is not aligned.") slots = self.buffer_req_to_token_slots[layer][row, :length].long().clone() ordered_slots = slots.sort().values - self._assert_final_prefill_tensor( + self._assert_retention_tensor( ((slots >= 0) & (slots < storage.slot_capacity())).all() & (ordered_slots[1:] > ordered_slots[:-1]).all(), "H2O retention physical slots must be valid and unique.", ) release_counts[layer] = release_counts.get(layer, 0) + length - budget - end = int(self._num_free_slots[layer]) + release_counts[layer] - if end > self.free_slots_stack[layer].numel(): - raise RuntimeError("H2O retention would overflow the free-slot stack.") + pointer = int(self._num_free_slots[layer]) + end = pointer + release_counts[layer] + if pointer < 0 or end > self.free_slots_stack[layer].numel(): + raise RuntimeError(f"H2O retention would overflow the free-slot stack: layer={layer}.") # Every query head keeps its own history, including heads that did # not supply the group's maximum on this step. score_keep = keep.repeat_interleave(score.shape[0] // groups, dim=0) diff --git a/tests/test_glm_mla_prefix_cache.py b/tests/test_glm_mla_prefix_cache.py index 92551639..4750c196 100644 --- a/tests/test_glm_mla_prefix_cache.py +++ b/tests/test_glm_mla_prefix_cache.py @@ -13,6 +13,8 @@ PrefillComputeView, ) from sparsevllm.engine.cache_manager.h2o import H2OCacheManager +from sparsevllm.engine.sparse_methods.h2o import H2ORuntime +from sparsevllm.engine.sparse_methods.base import SparseStepContext from sparsevllm.engine.cache_manager.rkv import RKVCacheManager from sparsevllm.engine.cache_manager.snapkv import SnapKVCacheManager from sparsevllm.engine.cache_manager.storage import MlaLatentStorage @@ -54,6 +56,7 @@ def _latent_chain_manager(manager_type, method: str): h2o_prefill_budget=8, h2o_recent_ratio=0.5, h2o_prefill_score_window=2, + h2o_head_reduction="max", rkv_compression_interval=2, rkv_observation_tokens=2, rkv_alpha=0.5, @@ -105,6 +108,7 @@ def _latent_chain_manager(manager_type, method: str): manager._prefill_attn_score_accumulators = {} manager._uniform_decode_metadata = False manager._h2o_scores = {} + manager._h2o_positions = {} manager._h2o_active_decode_seq_ids = set() manager._h2o_counters = { "intermediate_prefill_evictions": 0, @@ -113,7 +117,6 @@ def _latent_chain_manager(manager_type, method: str): "decode_evictions": 0, "dropped_tokens": 0, } - manager._h2o_final_prefill_workspace = None manager._rkv_query_cache_enabled = True manager._rkv_observation_tokens = 2 manager._rkv_vectorized_prefill_query_cache = True @@ -259,6 +262,9 @@ def test_snapkv_chain_resume_preserves_latent_payload_and_resets_prefill_scores( def test_h2o_chain_resume_preserves_aligned_scores_and_cleans_side_state(): manager, storage, config = _latent_chain_manager(H2OCacheManager, "h2o") coordinator = ChainCacheCoordinator(config, manager) + runtime = object.__new__(H2ORuntime) + runtime.config = config + runtime.cache_manager = manager owner_tokens = list(range(6)) owner = Sequence(owner_tokens) owner.seq_id = 0 @@ -277,9 +283,9 @@ def test_h2o_chain_resume_preserves_aligned_scores_and_cleans_side_state(): owner_slots = manager.layer_batch_states[0].slot_mapping.clone().long() _fill_latent_slots(storage, owner_slots, owner_tokens) manager._h2o_scores[(0, owner.seq_id)] = torch.tensor( - [1.0, 9.0, 2.0, 8.0, 0.0, 0.0] + [[1.0, 9.0, 2.0, 8.0, 0.0, 0.0], [0.0, 1.0, 0.0, 2.0, 0.0, 0.0]] ) - manager.evict_after_prefill([owner]) + runtime.finish_step(SparseStepContext([owner], True, None)) assert manager.row_seq_lens[0].tolist() == [4] resident_slots = manager.buffer_req_to_token_slots[0][0, :4].clone().long() payload = storage.layer_payload(0) @@ -291,8 +297,7 @@ def test_h2o_chain_resume_preserves_aligned_scores_and_cleans_side_state(): payload.rope_cache[resident_slots, 0, 0], torch.tensor([101, 103, 104, 105], dtype=torch.bfloat16), ) - assert manager._h2o_scores[(0, owner.seq_id)].tolist() == [9.0, 8.0, 0.0, 0.0] - assert manager._h2o_final_prefill_workspace is None + assert manager._h2o_scores[(0, owner.seq_id)].tolist() == [[9.0, 8.0, 0.0, 0.0], [1.0, 2.0, 0.0, 0.0]] coordinator.index.finish( owner.chain_id, token_ids=owner_tokens, @@ -320,7 +325,7 @@ def test_h2o_chain_resume_preserves_aligned_scores_and_cleans_side_state(): input_ids, positions, _ = manager._prepare_prefill([resumed]) assert input_ids.tolist() == [6, 7] assert positions.tolist() == [6, 7] - assert manager._h2o_scores[(0, resumed.seq_id)].tolist() == [9.0, 8.0, 0.0, 0.0] + assert manager._h2o_scores[(0, resumed.seq_id)].tolist() == [[9.0, 8.0, 0.0, 0.0], [1.0, 2.0, 0.0, 0.0]] resumed_row = manager.seq_id_to_row[0][resumed.seq_id] resumed_slots = manager.buffer_req_to_token_slots[0][ resumed_row, :6 @@ -329,11 +334,11 @@ def test_h2o_chain_resume_preserves_aligned_scores_and_cleans_side_state(): _fill_latent_slots(storage, resumed_slots[4:], [6, 7]) manager._h2o_scores[(0, resumed.seq_id)] = manager._accumulate_score( manager._h2o_scores[(0, resumed.seq_id)], - torch.tensor([0.0, 0.0, 0.0, 0.0, 7.0, 6.0]), + torch.tensor([[0.0, 0.0, 0.0, 0.0, 7.0, 6.0], [3.0, 2.0, 0.0, 0.0, 1.0, 1.0]]), new_len=6, weight=1.0, ) - manager.evict_after_prefill([resumed]) + runtime.finish_step(SparseStepContext([resumed], True, None)) final_slots = manager.buffer_req_to_token_slots[0][ resumed_row, :4 ].clone().long() @@ -345,9 +350,8 @@ def test_h2o_chain_resume_preserves_aligned_scores_and_cleans_side_state(): payload.rope_cache[final_slots, 0, 0], torch.tensor([101, 103, 106, 107], dtype=torch.bfloat16), ) - assert manager._h2o_scores[(0, resumed.seq_id)].tolist() == [9.0, 8.0, 7.0, 6.0] + assert manager._h2o_scores[(0, resumed.seq_id)].tolist() == [[9.0, 8.0, 7.0, 6.0], [4.0, 4.0, 1.0, 1.0]] assert manager._h2o_counters["final_prefill_evictions"] == 2 - assert manager._h2o_final_prefill_workspace is None coordinator.index.finish( owner.chain_id, @@ -358,6 +362,7 @@ def test_h2o_chain_resume_preserves_aligned_scores_and_cleans_side_state(): coordinator.invalidate(owner.chain_id) manager.free_seq(resumed.seq_id) assert manager._h2o_scores == {} + assert manager._h2o_positions == {} assert manager._prefill_attn_score_accumulators == {} assert manager.seq_id_to_row == [{}] assert manager._num_free_slots == [32] diff --git a/tests/test_h2o_cache_manager.py b/tests/test_h2o_cache_manager.py index cdc7b78e..d20e49d5 100644 --- a/tests/test_h2o_cache_manager.py +++ b/tests/test_h2o_cache_manager.py @@ -18,7 +18,6 @@ ) from sparsevllm.engine.cache_manager.h2o import H2OCacheManager from sparsevllm.engine.cache_manager.snapkv import SnapKVCacheManager -from sparsevllm.engine.cache_manager.storage import MlaLatentStorage from sparsevllm.engine.decode_graph_contract import ( DecodeGraphContract, DecodeGraphInputs, @@ -104,7 +103,6 @@ def _manager_with_layer_rows( "decode_evictions": 0, "dropped_tokens": 0, } - manager._h2o_final_prefill_workspace = None manager._uniform_decode_metadata = False manager.seq_id_to_row = [ {idx: idx for idx in range(batch_size)} @@ -188,32 +186,6 @@ def _set_layer_row_slots( ) -def _fill_kv_by_physical_slot(manager: H2OCacheManager): - for layer_idx in manager.kv_transformer_layer_indices(): - k_cache, v_cache = manager.get_layer_kv_cache(layer_idx) - slot_values = torch.arange(k_cache.shape[0], dtype=k_cache.dtype).view(-1, 1, 1) - offsets = torch.arange( - manager.num_kv_heads * manager.head_dim, - dtype=k_cache.dtype, - ).view(1, manager.num_kv_heads, manager.head_dim) - k_cache.copy_(slot_values * 10 + offsets) - v_cache.copy_(-slot_values * 10 - offsets - 1) - - -def _use_kv_layout(manager: H2OCacheManager, layout: str): - if layout == "tensor": - return - if layout != "list": - raise ValueError(f"unknown KV layout: {layout}") - manager.kv_cache = [ - ( - manager.kv_cache[0, layer_idx].clone(), - manager.kv_cache[1, layer_idx].clone(), - ) - for layer_idx in manager.kv_transformer_layer_indices() - ] - - def _assert_scores_match_slot_rows(manager: H2OCacheManager): for layer_idx in manager.kv_transformer_layer_indices(): for seq_id, row_idx in manager.seq_id_to_row[layer_idx].items(): @@ -454,9 +426,8 @@ def test_h2o_prefill_score_ranges_use_compressed_physical_coordinates(): assert ranges[0][2:] == (8, 8, 11) -def test_h2o_logit_prefill_score_window_zero_covers_full_current_chunk(): +def test_h2o_prefill_score_window_zero_covers_full_current_chunk(): manager = _manager_with_rows([11]) - manager.config.sparse_prefill_score_mode = "logits" manager.config.h2o_prefill_score_window = 0 seq = _seq(0, 100, prefilled=64, chunk=6) @@ -465,45 +436,6 @@ def test_h2o_logit_prefill_score_window_zero_covers_full_current_chunk(): assert ranges[0][2:] == (5, 5, 11) -def test_h2o_logit_prefill_score_is_normalized_before_weighted_accumulation(): - logits = torch.tensor([0.0, 1.0, 2.0]) - normalized = H2OCacheManager._normalize_logit_prefill_score(logits, new_len=3) - cumulative = H2OCacheManager._accumulate_score( - torch.tensor([1.0, 2.0]), - normalized, - new_len=3, - weight=4.0, - ) - - assert normalized.sum().item() == pytest.approx(1.0) - assert torch.equal(normalized, torch.softmax(logits, dim=0)) - assert torch.equal( - cumulative, - torch.tensor([1.0, 2.0, 0.0]) + 4.0 * torch.softmax(logits, dim=0), - ) - - -def test_h2o_logit_prefill_score_converts_unscored_minus_inf_to_zero_prob(): - logits = torch.tensor([0.0, -torch.inf, 2.0]) - normalized = H2OCacheManager._normalize_logit_prefill_score(logits, new_len=3) - assert normalized[1].item() == 0.0 - assert torch.isfinite(normalized).all() - assert normalized.sum().item() == pytest.approx(1.0) - - -def test_h2o_logit_prefill_score_rejects_nan_or_all_inf(): - with pytest.raises(RuntimeError, match="invalid non-finite values"): - H2OCacheManager._normalize_logit_prefill_score( - torch.tensor([float("nan"), 1.0]), - new_len=2, - ) - with pytest.raises(RuntimeError, match="invalid non-finite values"): - H2OCacheManager._normalize_logit_prefill_score( - torch.tensor([-torch.inf, -torch.inf]), - new_len=2, - ) - - def _score_runtime(manager): runtime = object.__new__(H2ORuntime) runtime.cache_manager = manager @@ -578,29 +510,6 @@ def test_h2o_missing_score_for_existing_physical_prefix_fails_fast(): manager._require_score_length(0, seq, 8) -def test_h2o_intermediate_and_final_prefill_use_distinct_budgets_and_counters(): - manager = _manager_with_rows([10], decode_budget=4, prefill_budget=8) - seq = _seq(0, 20, prefilled=0, chunk=10) - manager._h2o_scores[(0, 0)] = torch.arange(10, dtype=torch.float32) - - manager.evict_after_prefill([seq]) - assert manager.row_seq_lens[0][0] == 8 - assert manager._h2o_counters["intermediate_prefill_evictions"] == 1 - - manager.buffer_req_to_token_slots[0][0, 8:10] = torch.tensor([20, 21]) - manager.row_seq_lens[0][0] = 10 - manager._h2o_scores[(0, 0)] = manager._expand_score( - manager._h2o_scores[(0, 0)], 10, device=manager.device - ) - seq.num_prefilled_tokens = 10 - seq.current_chunk_size = 10 - manager.evict_after_prefill([seq]) - - assert manager.row_seq_lens[0][0] == 4 - assert manager._h2o_counters["final_prefill_evictions"] == 1 - assert manager._h2o_counters["dropped_tokens"] == 8 - - def test_h2o_decode_score_update_supports_batch_with_different_kv_lengths(): manager = _manager_with_rows([4, 3]) seq0 = _seq(0, 10, prefilled=10, chunk=1) @@ -866,221 +775,6 @@ def tracked_compact(self, layer_indices, batch_seqs, keep_indices, **kwargs): } -def test_h2o_intermediate_prefill_reuses_uniform_batch_fast_path(): - manager = _manager_with_layer_rows( - [[6, 6], [7, 7]], decode_budget=3, prefill_budget=4 - ) - seqs = [ - _seq(0, 20, prefilled=0, chunk=6), - _seq(1, 20, prefilled=0, chunk=6), - ] - _set_scores_from_slot_rows(manager) - calls = [] - original = SnapKVCacheManager.free_part_slots_batch_layers - - def tracked_batch_free(self, layer_indices, batch_seqs, keep_indices, **kwargs): - calls.append((list(layer_indices), keep_indices.clone())) - return original(self, layer_indices, batch_seqs, keep_indices, **kwargs) - - with patch.object( - SnapKVCacheManager, - "free_part_slots_batch_layers", - new=tracked_batch_free, - ): - assert manager._try_batched_evict(seqs, is_prefill=True) - - assert len(calls) == 1 - assert calls[0][0] == [0, 1] - assert calls[0][1].shape == (2, 2, 4) - assert [lengths.tolist() for lengths in manager.row_seq_lens] == [[4, 4], [4, 4]] - _assert_scores_match_slot_rows(manager) - assert manager._h2o_counters == { - "intermediate_prefill_evictions": 4, - "final_prefill_evictions": 0, - "decode_eviction_bursts": 0, - "decode_evictions": 0, - "dropped_tokens": 10, - } - - -def test_h2o_mixed_prefill_budgets_fall_back_with_exact_counters(): - manager = _manager_with_rows([6, 6], decode_budget=3, prefill_budget=4) - seqs = [ - _seq(0, 6, prefilled=0, chunk=6), - _seq(1, 20, prefilled=0, chunk=6), - ] - _set_scores_from_slot_rows(manager) - _fill_kv_by_physical_slot(manager) - k_cache, v_cache = manager.get_layer_kv_cache(0) - selected_slots = torch.tensor([3, 4, 5], dtype=torch.long) - expected_k = k_cache.index_select(0, selected_slots).clone() - expected_v = v_cache.index_select(0, selected_slots).clone() - - assert not manager._try_batched_evict(seqs, is_prefill=True) - manager.evict_after_prefill(seqs) - - assert manager.row_seq_lens[0].tolist() == [3, 4] - assert manager._h2o_scores[(0, 0)].tolist() == [3.0, 4.0, 5.0] - assert manager._h2o_scores[(0, 1)].tolist() == [102.0, 103.0, 104.0, 105.0] - final_slots = manager.buffer_req_to_token_slots[0][0, :3].long() - assert final_slots.tolist() == [0, 1, 2] - assert torch.equal(k_cache.index_select(0, final_slots), expected_k) - assert torch.equal(v_cache.index_select(0, final_slots), expected_v) - assert manager._h2o_counters == { - "intermediate_prefill_evictions": 1, - "final_prefill_evictions": 1, - "decode_eviction_bursts": 0, - "decode_evictions": 0, - "dropped_tokens": 5, - } - - -@pytest.mark.parametrize("kv_layout", ["tensor", "list"]) -def test_h2o_final_prefill_page_table_compaction_preserves_logical_kv_alignment( - kv_layout: str, -): - manager = _manager_with_layer_rows( - [[6, 6], [6, 6]], decode_budget=4, prefill_budget=8 - ) - rows_by_layer = [ - [[9, 2, 7, 1, 6, 4], [29, 22, 27, 21, 26, 24]], - [[109, 102, 107, 101, 106, 104], [129, 122, 127, 121, 126, 124]], - ] - for layer_idx, rows in enumerate(rows_by_layer): - _set_layer_row_slots(manager, layer_idx, rows) - for seq_id in range(2): - manager._h2o_scores[(layer_idx, seq_id)] = torch.tensor( - [1.0, 9.0, 2.0, 8.0, 0.0, 0.0] - ) - _use_kv_layout(manager, kv_layout) - _fill_kv_by_physical_slot(manager) - seqs = [ - _seq(0, 6, prefilled=0, chunk=6), - _seq(1, 6, prefilled=0, chunk=6), - ] - keep = torch.tensor([1, 3, 4, 5], dtype=torch.long) - keep_set = set(keep.tolist()) - expected = {} - for layer_idx in range(2): - k_cache, v_cache = manager.get_layer_kv_cache(layer_idx) - for seq_id, row_slots in enumerate(rows_by_layer[layer_idx]): - selected_slots = torch.tensor(row_slots, dtype=torch.long)[keep] - expected[(layer_idx, seq_id)] = ( - k_cache.index_select(0, selected_slots).clone(), - v_cache.index_select(0, selected_slots).clone(), - ) - - manager.evict_after_prefill(seqs) - - for layer_idx, rows in enumerate(rows_by_layer): - k_cache, v_cache = manager.get_layer_kv_cache(layer_idx) - released = [] - active = [] - for seq_id, row_slots in enumerate(rows): - destination = torch.tensor(row_slots, dtype=torch.long)[keep].tolist() - released.extend( - slot for idx, slot in enumerate(row_slots) if idx not in keep_set - ) - active.extend(destination) - actual_slots = manager.buffer_req_to_token_slots[layer_idx][ - seq_id, :4 - ].long() - assert actual_slots.tolist() == destination - expected_k, expected_v = expected[(layer_idx, seq_id)] - assert torch.equal(k_cache.index_select(0, actual_slots), expected_k) - assert torch.equal(v_cache.index_select(0, actual_slots), expected_v) - assert manager._h2o_scores[(layer_idx, seq_id)].tolist() == [ - 9.0, - 8.0, - 0.0, - 0.0, - ] - assert len(active) == len(set(active)) - assert manager.free_slots_stack[layer_idx][32:36].tolist() == released - assert manager._num_free_slots[layer_idx] == 36 - - assert manager._h2o_final_prefill_workspace is None - - -def test_h2o_final_prefill_compacts_mla_latent_and_rope_slots(): - manager = _manager_with_rows([6], decode_budget=4, prefill_budget=8) - _set_layer_row_slots(manager, 0, [[9, 2, 7, 1, 6, 4]]) - manager._h2o_scores[(0, 0)] = torch.tensor( - [1.0, 9.0, 2.0, 8.0, 0.0, 0.0] - ) - storage = MlaLatentStorage( - kv_lora_rank=512, - rope_dim=64, - dtype=torch.bfloat16, - ) - storage.allocate(num_layers=1, num_slots=64, device=torch.device("cpu")) - manager.attention_cache_storage = storage - manager.kv_cache = None - assert storage.latent_cache is not None - assert storage.rope_cache is not None - for slot in [9, 2, 7, 1, 6, 4]: - storage.latent_cache[0, slot].fill_(slot) - storage.rope_cache[0, slot].fill_(slot + 100) - seq = _seq(0, 6, prefilled=0, chunk=6) - - manager.evict_after_prefill([seq]) - - destination_slots = manager.buffer_req_to_token_slots[0][0, :4].long() - assert destination_slots.tolist() == [2, 1, 6, 4] - payload = storage.layer_payload(0) - expected_sources = torch.tensor([2, 1, 6, 4], dtype=torch.bfloat16) - torch.testing.assert_close( - payload.latent_cache[destination_slots, 0, 0], - expected_sources, - ) - torch.testing.assert_close( - payload.rope_cache[destination_slots, 0, 0], - expected_sources + 100, - ) - assert manager._h2o_scores[(0, 0)].tolist() == [9.0, 8.0, 0.0, 0.0] - assert manager._h2o_final_prefill_workspace is None - - -def test_h2o_intermediate_prefill_does_not_move_kv_payloads(): - manager = _manager_with_rows([6], decode_budget=3, prefill_budget=4) - _set_layer_row_slots(manager, 0, [[9, 2, 7, 1, 6, 4]]) - manager._h2o_scores[(0, 0)] = torch.tensor([1.0, 9.0, 2.0, 8.0, 0.0, 0.0]) - _fill_kv_by_physical_slot(manager) - k_cache, v_cache = manager.get_layer_kv_cache(0) - old_k = k_cache.clone() - old_v = v_cache.clone() - seq = _seq(0, 20, prefilled=0, chunk=6) - - manager.evict_after_prefill([seq]) - - assert manager.buffer_req_to_token_slots[0][0, :4].tolist() == [2, 1, 6, 4] - assert torch.equal(k_cache, old_k) - assert torch.equal(v_cache, old_v) - assert manager._h2o_final_prefill_workspace is None - - -def test_h2o_final_prefill_capacity_preflight_prevents_partial_layer_updates(): - manager = _manager_with_layer_rows([[6], [6]], decode_budget=4, prefill_budget=8) - _set_scores_from_slot_rows(manager) - _fill_kv_by_physical_slot(manager) - manager._num_free_slots[1] = 511 - layer0_slots = manager.buffer_req_to_token_slots[0].clone() - layer0_k, layer0_v = manager.get_layer_kv_cache(0) - expected_k = layer0_k.clone() - expected_v = layer0_v.clone() - seq = _seq(0, 6, prefilled=0, chunk=6) - - with pytest.raises(RuntimeError, match=r"overflow.*layer=1"): - manager.evict_after_prefill([seq]) - - assert torch.equal(manager.buffer_req_to_token_slots[0], layer0_slots) - assert torch.equal(layer0_k, expected_k) - assert torch.equal(layer0_v, expected_v) - assert manager.row_seq_lens[0].tolist() == [6] - assert manager._num_free_slots[0] == 32 - assert manager._h2o_final_prefill_workspace is None - - def test_h2o_decode_waits_for_interval_then_drops_full_burst(): manager = _manager_with_rows( [5], decode_budget=4, decode_eviction_interval=3, prefill_budget=8 @@ -1154,42 +848,6 @@ def test_h2o_decode_pressure_reclaims_over_budget_row_before_interval(): assert preempted == [] -@pytest.mark.parametrize("use_tensor_table", [True, False]) -def test_h2o_final_prefill_pressure_reclaims_unscheduled_active_decode_row( - use_tensor_table: bool, -): - manager = _manager_with_layer_rows( - [[5, 2], [5, 2]], - decode_budget=4, - decode_eviction_interval=3, - prefill_budget=8, - ) - active_decode = _seq(0, 5, prefilled=5, chunk=1) - final_prefill = _seq(1, 2, prefilled=0, chunk=2) - _set_scores_from_slot_rows(manager) - manager.evict_after_decode([active_decode]) - # Post-forward state: this final prefill consumed the last physical slot. - manager._num_free_slots = [0, 0] - if not use_tensor_table: - manager.buffer_req_to_token_slots_tensor = None - - manager.evict_after_prefill([final_prefill]) - - assert [lengths.tolist() for lengths in manager.row_seq_lens] == [ - [4, 2], - [4, 2], - ] - _assert_scores_match_slot_rows(manager) - assert manager._num_free_slots == [1, 1] - assert manager._h2o_counters == { - "intermediate_prefill_evictions": 0, - "final_prefill_evictions": 0, - "decode_eviction_bursts": 1, - "decode_evictions": 2, - "dropped_tokens": 2, - } - - def test_h2o_decode_pressure_reclaims_unscheduled_active_row(): manager = _manager_with_rows( [5, 4], decode_budget=4, decode_eviction_interval=3, prefill_budget=8 diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py index 8dc9216a..30aef110 100644 --- a/tests/test_h2o_per_head_prefill.py +++ b/tests/test_h2o_per_head_prefill.py @@ -193,3 +193,94 @@ def test_chunk_prefill_append_retention_and_final_handoff_follow_logical_positio for head in range(2): assert (manager.attention_cache_storage.cache[:, 0, slot, head] == histories[head][packed] + 100 * head).all() assert manager._num_free_slots[0] + resident == 12 + + +def _multi_request_manager(): + manager = manager_with_storage() + manager.num_layers = manager.num_kv_layers = 2 + manager.runtime_layout.kv_idx_to_layer_idx = (0, 1) + manager.attention_cache_storage.allocate(num_layers=2, num_slots=24, device=manager.device) + cache = manager.attention_cache_storage.cache + cache.copy_(torch.arange(cache.numel()).reshape_as(cache)) + manager.config = SimpleNamespace(h2o_prefill_budget=4, h2o_decode_budget=3, + h2o_recent_ratio=.5, h2o_head_reduction='max') + manager.seq_id_to_row = [{7: 0, 8: 1}, {7: 0, 8: 1}] + manager.row_seq_lens = [np.array([6, 5], dtype=np.int32) for _ in range(2)] + manager.buffer_req_to_token_slots = [torch.zeros(2, 8, dtype=torch.int32) for _ in range(2)] + manager.free_slots_stack = [torch.zeros(24, dtype=torch.int32) for _ in range(2)] + manager._num_free_slots = [13, 13] + manager._h2o_scores.clear() + manager._h2o_positions.clear() + for layer in range(2): + for row, (seq_id, length) in enumerate([(7, 6), (8, 5)]): + manager.buffer_req_to_token_slots[layer][row, :length] = torch.arange(length).flip(0) + row * 12 + manager._h2o_scores[(layer, seq_id)] = torch.stack([ + torch.arange(length).float().roll(layer + head) for head in range(4) + ]) + manager._h2o_positions[(layer, seq_id)] = torch.arange(length)[None].expand(2, -1).clone() + return manager + + +def _mixed_prefill_seqs(): + from sparsevllm.engine.sequence import Sequence + seqs = [Sequence(list(range(6))), Sequence(list(range(20)))] + for seq, seq_id, chunk in zip(seqs, [7, 8], [6, 5]): + seq.seq_id = seq_id + seq.current_chunk_size = chunk + return seqs + + +def test_multilayer_mixed_prefill_budgets_preserve_independent_head_payloads(): + from sparsevllm.engine.sparse_methods.base import SparseStepContext + from sparsevllm.engine.sparse_methods.h2o import H2ORuntime + + manager = _multi_request_manager() + expected = {} + old_cache = manager.attention_cache_storage.cache.clone() + for layer in range(2): + for row, (seq_id, length, budget) in enumerate([(7, 6, 3), (8, 5, 4)]): + scores = manager._h2o_scores[(layer, seq_id)] + recent = max(1, int(budget * .5)) + for group in range(2): + rank = lambda index: max(float(scores[2 * group + head, index]) for head in range(2)) + heavy = sorted(range(length - recent), key=lambda index: (-rank(index), index))[:budget - recent] + selected = sorted(heavy + list(range(length - recent, length))) + slots = manager.buffer_req_to_token_slots[layer][row, selected].long() + expected[layer, seq_id, group] = (selected, old_cache[:, layer, slots, group].clone()) + runtime = object.__new__(H2ORuntime) + runtime.config, runtime.cache_manager = manager.config, manager + runtime.finish_step(SparseStepContext(_mixed_prefill_seqs(), True, None)) + for layer in range(2): + assert manager.row_seq_lens[layer].tolist() == [3, 4] + for row, (seq_id, budget) in enumerate([(7, 3), (8, 4)]): + slots = manager.buffer_req_to_token_slots[layer][row, :budget].long() + for group in range(2): + positions, payload = expected[layer, seq_id, group] + assert manager._h2o_positions[layer, seq_id][group].tolist() == positions + torch.testing.assert_close(manager.attention_cache_storage.cache[:, layer, slots, group], payload) + assert manager._num_free_slots == [17, 17] + assert manager._h2o_counters['final_prefill_evictions'] == 2 + assert manager._h2o_counters['intermediate_prefill_evictions'] == 2 + assert manager._h2o_counters['dropped_tokens'] == 8 + + +def test_later_layer_capacity_failure_does_not_commit_earlier_layers(): + from sparsevllm.engine.sparse_methods.base import SparseStepContext + from sparsevllm.engine.sparse_methods.h2o import H2ORuntime + + manager = _multi_request_manager() + manager._num_free_slots[1] = 23 + old_cache = manager.attention_cache_storage.cache.clone() + old_tables = [table.clone() for table in manager.buffer_req_to_token_slots] + old_scores = {key: score.clone() for key, score in manager._h2o_scores.items()} + runtime = object.__new__(H2ORuntime) + runtime.config, runtime.cache_manager = manager.config, manager + with pytest.raises(RuntimeError, match='overflow.*layer=1'): + runtime.finish_step(SparseStepContext(_mixed_prefill_seqs(), True, None)) + torch.testing.assert_close(manager.attention_cache_storage.cache, old_cache) + for actual, before in zip(manager.buffer_req_to_token_slots, old_tables): + torch.testing.assert_close(actual, before) + for key, before in old_scores.items(): + torch.testing.assert_close(manager._h2o_scores[key], before) + assert manager._num_free_slots == [13, 23] + assert [lengths.tolist() for lengths in manager.row_seq_lens] == [[6, 5], [6, 5]] diff --git a/tests/test_static_eviction_compaction.py b/tests/test_static_eviction_compaction.py index 1fa5aee1..978d1e93 100644 --- a/tests/test_static_eviction_compaction.py +++ b/tests/test_static_eviction_compaction.py @@ -571,39 +571,3 @@ def test_pyramidkv_final_prefill_groups_layers_by_effective_budget(): assert manager.batch_calls[0][0] == 0 assert manager.batch_calls[0][2].shape == (2, 6) assert not manager.scalar_calls - - -def test_h2o_scalar_final_prefill_keeps_dense_relocation_dispatch(): - manager = object.__new__(H2OCacheManager) - manager.config = SimpleNamespace( - h2o_decode_budget=4, - h2o_prefill_budget=8, - h2o_recent_ratio=0.5, - ) - manager.kv_transformer_layer_indices = lambda: [0] - manager._preflight_final_prefill_dense_capacity = Mock() - manager._try_batched_evict = Mock(return_value=False) - manager._physical_row_len = Mock(return_value=6) - score = torch.arange(6, dtype=torch.float32) - manager._require_score_length = Mock(return_value=score) - keep = torch.tensor([1, 3, 4, 5], dtype=torch.long) - manager.select_h2o_indices = Mock(return_value=keep) - manager._compact_final_prefill_dense_batch = Mock() - manager.free_part_slots = Mock() - manager._score_key = lambda layer_idx, seq_id: (int(layer_idx), int(seq_id)) - manager._h2o_scores = {} - manager._h2o_counters = { - "intermediate_prefill_evictions": 0, - "final_prefill_evictions": 0, - "decode_evictions": 0, - "dropped_tokens": 0, - } - seq = _final_prefill_seq(30, 6) - - manager._evict([seq], is_prefill=True) - - manager._compact_final_prefill_dense_batch.assert_called_once() - manager.free_part_slots.assert_not_called() - assert torch.equal(manager._h2o_scores[(0, 30)], score[keep]) - assert manager._h2o_counters["final_prefill_evictions"] == 1 - assert manager._h2o_counters["dropped_tokens"] == 2 From d2c5584123d3631873030811c13949e03ae9d816 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 17:49:20 +0800 Subject: [PATCH 03/11] test: add opt-in h2o model handoff validation --- tests/test_h2o_model_integration.py | 91 +++++++++++++++++++++++++++++ 1 file changed, 91 insertions(+) create mode 100644 tests/test_h2o_model_integration.py diff --git a/tests/test_h2o_model_integration.py b/tests/test_h2o_model_integration.py new file mode 100644 index 00000000..35ac4667 --- /dev/null +++ b/tests/test_h2o_model_integration.py @@ -0,0 +1,91 @@ +"""Opt-in real-model H2O lifecycle check on an exclusively available GPU. + +Set SPARSEVLLM_H2O_TEST_MODEL to a local checkpoint. Artifacts can be retained +with SPARSEVLLM_H2O_TEST_OUTPUT; this is a correctness smoke test, not a quality +or performance benchmark. +""" + +import json +import os +from pathlib import Path +import subprocess + +import pytest + + +@pytest.mark.skipif(not os.environ.get('SPARSEVLLM_H2O_TEST_MODEL'), reason='requires an explicit local model') +def test_model_chunk_prefill_and_native_decode_handoff(tmp_path): + import torch + from sparsevllm import LLM, SamplingParams + from sparsevllm.engine.cache_manager.storage import MlaLatentStorage + + model = os.environ['SPARSEVLLM_H2O_TEST_MODEL'] + output_dir = Path(os.environ.get('SPARSEVLLM_H2O_TEST_OUTPUT', str(tmp_path))) + output_dir.mkdir(parents=True, exist_ok=True) + config = dict(sparse_method='h2o', h2o_head_reduction='max', + h2o_prefill_budget=64, h2o_decode_budget=32, + h2o_prefill_score_window=0, engine_prefill_chunk_size=64, + max_num_batched_tokens=128, max_num_seqs_in_batch=2, + max_num_seqs_in_gpu=2, max_model_len=512, + decode_graph=False, enable_prefix_caching=False, + validate_runtime_invariants=True) + artifact = dict(model=model, config=config, seed=19, + revision=subprocess.check_output(['git', 'rev-parse', 'HEAD'], text=True).strip(), + torch_version=torch.__version__, cuda_version=torch.version.cuda, + samples=[], steps=[]) + engine = None + try: + torch.manual_seed(19) + engine = LLM(model, **config) + manager = engine.model_runner.cache_manager + free_before = list(manager._num_free_slots) + groups = 1 if isinstance(manager.attention_cache_storage, MlaLatentStorage) else manager.num_kv_heads + heads = int(engine.config.hf_config.num_attention_heads) + sample_text = 'The archive contains notes about rivers, mountains, books, and gardens. ' + repeated = engine.tokenizer.encode(sample_text * 50, add_special_tokens=False) + prompts = [repeated[:161], repeated[:99]] + assert [len(prompt) for prompt in prompts] == [161, 99] + for prompt in prompts: + engine.add_request(prompt, SamplingParams(temperature=0, max_tokens=8, ignore_eos=True)) + artifact['samples'].append(dict(status='model_failed', prompt_token_ids=prompt, + token_ids=[], text='')) + outputs = {} + for _ in range(32): + if engine.is_finished(): + break + completed, scheduled_tokens = engine.step() + for seq_id, tokens, *_ in completed: + outputs[seq_id] = tokens + step = dict(scheduled_tokens=scheduled_tokens, scores=[], counters=dict(manager._h2o_counters)) + for (layer, seq_id), scores in manager._h2o_scores.items(): + positions = manager._h2o_positions[layer, seq_id] + row = manager.seq_id_to_row[layer][seq_id] + resident = int(manager.row_seq_lens[layer][row]) + assert scores.ndim == 2 and scores.shape[0] == heads + assert positions.shape == (groups, scores.shape[-1]) + assert resident >= scores.shape[-1] + assert scores.shape[-1] <= config['h2o_prefill_budget'] + assert bool(torch.isfinite(scores).all()) + step['scores'].append(dict(layer=layer, seq_id=seq_id, + score_length=scores.shape[-1], resident=resident)) + artifact['steps'].append(step) + assert engine.is_finished(), 'bounded generation did not finish' + assert len(outputs) == len(prompts) + for sample, tokens in zip(artifact['samples'], [outputs[key] for key in sorted(outputs)]): + sample.update(status='success', token_ids=tokens, + text=engine.tokenizer.decode(tokens, skip_special_tokens=True)) + assert len(tokens) == 8 + assert manager._h2o_counters['intermediate_prefill_evictions'] > 0 + assert manager._h2o_counters['final_prefill_evictions'] > 0 + assert manager._h2o_counters['decode_evictions'] == 0 + assert not manager._h2o_scores and not manager._h2o_positions + assert manager._num_free_slots == free_before + artifact['status'] = 'success' + except Exception as error: + artifact['status'] = 'model_failed' + artifact['error'] = f'{type(error).__name__}: {error}' + raise + finally: + (output_dir / 'model_validation.json').write_text(json.dumps(artifact, indent=2, ensure_ascii=False)) + if engine is not None: + engine.exit() From 24bbd0efb7ad8d277a92572034c218e5266b6743 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 17:57:33 +0800 Subject: [PATCH 04/11] test: validate h2o chunk numerics and memory reservations --- tests/test_h2o_chunk_numerics.py | 157 +++++++++++++++++++++++++++++ tests/test_h2o_per_head_prefill.py | 40 ++++++++ 2 files changed, 197 insertions(+) create mode 100644 tests/test_h2o_chunk_numerics.py diff --git a/tests/test_h2o_chunk_numerics.py b/tests/test_h2o_chunk_numerics.py new file mode 100644 index 00000000..fa11c702 --- /dev/null +++ b/tests/test_h2o_chunk_numerics.py @@ -0,0 +1,157 @@ +"""Real CUDA score/compaction/attention paths against logical-token oracles.""" + +from types import SimpleNamespace + +import numpy as np +import pytest +import torch + +from sparsevllm.engine.cache_manager.base import ( + AttentionViewMeta, ExplicitKVPayload, LayerBatchStates, PrefillComputeView, +) +from sparsevllm.engine.cache_manager.h2o import H2OCacheManager +from sparsevllm.engine.cache_manager.storage import ExplicitKVStorage +from sparsevllm.engine.sequence import Sequence +from sparsevllm.engine.sparse_methods.base import PrefillScoreEvent, SparseStepContext +from sparsevllm.engine.sparse_methods.h2o import H2ORuntime +from sparsevllm.kernels.triton.context_flashattention_nopad import context_attention_fwd +from sparsevllm.kernels.triton.prefill_score import PrefillScoreWorkspace +from sparsevllm.utils.context import reset_context, set_context + +pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason='requires CUDA') + + +def _manager(heads, kv_heads, length, budget, reduction): + manager = object.__new__(H2OCacheManager) + manager.device = torch.device('cuda:0') + manager.num_kv_heads = kv_heads + manager.num_layers = manager.num_kv_layers = 1 + manager.runtime_layout = SimpleNamespace(kv_layer_index=lambda layer: layer, kv_idx_to_layer_idx=(0,)) + manager.max_model_len = length + 1 + manager.attention_cache_storage = ExplicitKVStorage(num_kv_heads=kv_heads, head_dim=32, dtype=torch.bfloat16) + manager.attention_cache_storage.allocate(num_layers=1, num_slots=length + 1, device=manager.device) + manager.seq_id_to_row = [{7: 0}] + manager.row_seq_lens = [np.zeros(1, dtype=np.int32)] + manager.layer_batch_states = [LayerBatchStates()] + manager.buffer_req_to_token_slots = [torch.zeros(1, length + 1, dtype=torch.int32, device=manager.device)] + manager.free_slots_stack = [torch.randperm(length + 1, device=manager.device).int()] + manager._num_free_slots = [length + 1] + manager._h2o_scores, manager._h2o_positions = {}, {} + manager._h2o_counters = dict(final_prefill_evictions=0, intermediate_prefill_evictions=0, dropped_tokens=0) + manager.config = SimpleNamespace(h2o_prefill_budget=budget, h2o_decode_budget=budget, + h2o_prefill_score_window=0, h2o_recent_ratio=.25, + h2o_head_reduction=reduction) + runtime = object.__new__(H2ORuntime) + runtime.config, runtime.cache_manager = manager.config, manager + runtime._prefill_score_workspace = PrefillScoreWorkspace() + runtime._prefill_head_score_buffer = runtime._prefill_head_score_total = None + return manager, runtime + + +def _run_chunks(q, k, v, chunks, budget, reduction): + length, heads, dim = q.shape + kv_heads = k.shape[1] + group_size = heads // kv_heads + manager, runtime = _manager(heads, kv_heads, length, budget, reduction) + seq = Sequence(list(range(length))) + seq.seq_id = 7 + histories = [[] for _ in range(kv_heads)] + score_reference = [dict() for _ in range(heads)] + output_chunks = [] + start = 0 + try: + for chunk in chunks: + seq.num_prefilled_tokens, seq.current_chunk_size = start, chunk + manager._prepare_prefill([seq]) + state = manager.layer_batch_states[0] + new_slots = state.slot_mapping.long() + storage = manager.attention_cache_storage + storage.cache[0, 0, new_slots] = k[start:start + chunk] + storage.cache[1, 0, new_slots] = v[start:start + chunk] + for history in histories: + history.extend(range(start, start + chunk)) + resident = len(histories[0]) + zero = torch.zeros(1, device=q.device, dtype=torch.int32) + chunk_tensor = torch.tensor([chunk], device=q.device, dtype=torch.int32) + prefix = state.context_lens - chunk_tensor + view = PrefillComputeView( + meta=AttentionViewMeta(active_slots=manager.buffer_req_to_token_slots[0], + req_indices=zero, context_lens=state.context_lens, + max_context_len=resident), + payload=ExplicitKVPayload(k_cache=storage.cache[0, 0], v_cache=storage.cache[1, 0]), + ) + query = q[start:start + chunk] + output = torch.empty_like(query) + context_attention_fwd(query, view.payload.k_cache, view.payload.v_cache, output, + zero, zero, state.context_lens, prefix, chunk, view.meta.active_slots) + expected = torch.empty_like(query, dtype=torch.float32) + # The oracle reads original tokens, never the packed production KV. + for head in range(heads): + group = head // group_size + history = histories[group] + keys = k[history, group].float() + logits = query[:, head].float() @ keys.T * dim ** -.5 + causal = torch.tensor(history, device=q.device)[None] > torch.arange(start, start + chunk, device=q.device)[:, None] + probability = logits.masked_fill(causal, -torch.inf).softmax(-1) + expected[:, head] = probability @ v[history, group].float() + for token, mass in zip(history, probability.sum(0).tolist()): + score_reference[head][token] = score_reference[head].get(token, 0.) + mass + torch.testing.assert_close(output.float(), expected, rtol=2e-2, atol=2e-2) + set_context(True, cache_manager=manager, seqs=[seq]) + runtime.collect_prefill_attention_score(PrefillScoreEvent(0, query, view, zero, chunk_tensor, dim ** -.5)) + for head in range(heads): + expected_scores = torch.tensor([score_reference[head][token] for token in histories[head // group_size]], device=q.device) + torch.testing.assert_close(manager._h2o_scores[0, 7][head], expected_scores, rtol=5e-3, atol=5e-3) + runtime.finish_step(SparseStepContext([seq], True, None)) + if resident > budget: + recent = max(1, int(budget * .25)) + for group, history in enumerate(histories): + def importance(token): + scores = [score_reference[h][token] for h in range(group * group_size, (group + 1) * group_size)] + return max(scores) if reduction == 'max' else sum(scores) / group_size + heavy = sorted(history[:-recent], key=lambda token: (-importance(token), token))[:budget - recent] + histories[group] = sorted(heavy + history[-recent:]) + assert manager._h2o_positions[0, 7].tolist() == histories + for head in range(heads): + score_reference[head] = {token: score_reference[head][token] for token in histories[head // group_size]} + torch.testing.assert_close(manager._h2o_scores[0, 7][head], torch.tensor(list(score_reference[head].values()), device=q.device), rtol=5e-3, atol=5e-3) + assert manager._num_free_slots[0] + len(histories[0]) == length + 1 + output_chunks.append(output) + start += chunk + # A one-query score-free read is the causal attention contract used at + # decode handoff; model tests separately exercise the native provider. + resident = len(histories[0]) + context = torch.tensor([resident], device=q.device, dtype=torch.int32) + output = torch.empty_like(q[:1]) + context_attention_fwd(q[:1], storage.cache[0, 0], storage.cache[1, 0], output, + zero, zero, context, context - 1, 1, view.meta.active_slots) + for head in range(heads): + group = head // group_size + keys, values = k[histories[group], group].float(), v[histories[group], group].float() + expected = (q[0, head].float() @ keys.T * dim ** -.5).softmax(-1) @ values + torch.testing.assert_close(output[0, head].float(), expected, rtol=2e-2, atol=2e-2) + return torch.cat(output_chunks), manager._h2o_scores[0, 7].clone() + finally: + reset_context() + + +@pytest.mark.parametrize('kv_heads,reduction', [(4, 'max'), (2, 'max'), (2, 'mean')]) +def test_chunk_attention_and_retention_match_logical_token_oracle(kv_heads, reduction): + torch.manual_seed(57) + q = torch.randn(35, 4, 32, device='cuda', dtype=torch.bfloat16) + k = torch.randn(35, kv_heads, 32, device='cuda', dtype=torch.bfloat16) + v = torch.randn_like(k) + _run_chunks(q, k, v, [11, 7, 17], budget=9, reduction=reduction) + + +@pytest.mark.parametrize('kv_heads', [4, 2]) +def test_no_eviction_full_and_chunked_prefill_are_equivalent(kv_heads): + torch.manual_seed(61) + q = torch.randn(149, 4, 32, device='cuda', dtype=torch.bfloat16) + k = torch.randn(149, kv_heads, 32, device='cuda', dtype=torch.bfloat16) + v = torch.randn_like(k) + full_output, full_scores = _run_chunks(q, k, v, [149], budget=149, reduction='max') + for chunks in ([13, 129, 7], [71, 78]): + output, scores = _run_chunks(q, k, v, chunks, budget=149, reduction='max') + torch.testing.assert_close(output, full_output, rtol=2e-2, atol=2e-2) + torch.testing.assert_close(scores, full_scores, rtol=5e-3, atol=5e-3) diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py index 30aef110..1ac10275 100644 --- a/tests/test_h2o_per_head_prefill.py +++ b/tests/test_h2o_per_head_prefill.py @@ -284,3 +284,43 @@ def test_later_layer_capacity_failure_does_not_commit_earlier_layers(): torch.testing.assert_close(manager._h2o_scores[key], before) assert manager._num_free_slots == [13, 23] assert [lengths.tolist() for lengths in manager.row_seq_lens] == [[6, 5], [6, 5]] + + +@pytest.mark.parametrize('mla,heads,groups,tp', [(False, 8, 8, 1), (False, 32, 2, 1), (False, 16, 2, 2), (True, 32, 1, 1)]) +def test_capacity_reserves_native_payload_and_live_head_metadata(monkeypatch, mla, heads, groups, tp): + from sparsevllm.engine.cache_manager.snapkv import SnapKVCacheManager + + manager = object.__new__(H2OCacheManager) + manager.tp_size, manager.num_kv_heads = tp, groups + manager.hf_config = SimpleNamespace(num_attention_heads=heads * tp) + manager.max_buffer_rows, manager.max_model_len = 3, 2048 + manager.config = SimpleNamespace(max_num_seqs_in_batch=2, engine_prefill_chunk_size=257, + h2o_prefill_budget=129, h2o_prefill_score_window=0, + enable_prefix_caching=False) + storage = (MlaLatentStorage(kv_lora_rank=512, rope_dim=64, dtype=torch.bfloat16) if mla + else ExplicitKVStorage(num_kv_heads=groups, head_dim=128, dtype=torch.bfloat16)) + manager.attention_cache_storage = storage + native_bytes = storage.bytes_per_slot_per_layer() + available = 64 * 1024 * 1024 + monkeypatch.setattr(SnapKVCacheManager, '_get_available_slots_info', lambda self: (available, native_bytes)) + remaining, slot_cost = manager._get_available_slots_info() + # Scores and positions must survive while the newly gathered copies exist. + assert slot_cost >= native_bytes + 2 * (4 * heads + 8 * groups) + assert slot_cost - manager._h2o_metadata_bytes_per_slot == native_bytes + score_buffers = 2 * 4 * 2 * heads * (129 + 257) + copy_buffers = 129 * native_bytes + assert available - remaining >= score_buffers + copy_buffers + assert 0 < remaining < available + monkeypatch.setattr(SnapKVCacheManager, '_get_available_slots_info', lambda self: (1, native_bytes)) + with pytest.raises(RuntimeError, match='Not enough memory'): + manager._get_available_slots_info() + + +def test_mla_rejects_tp_without_cross_rank_head_reduction(monkeypatch): + from sparsevllm.engine.cache_manager.snapkv import SnapKVCacheManager + + manager = manager_with_storage(mla=True) + manager.tp_size = 2 + monkeypatch.setattr(SnapKVCacheManager, '_get_available_slots_info', lambda self: (2**30, 1152)) + with pytest.raises(ValueError, match='requires TP1'): + manager._get_available_slots_info() From f5a73ae8767fa1303a94ba3c78e392c3f4478e6c Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 18:00:01 +0800 Subject: [PATCH 05/11] fix: preserve h2o history on rejected prefill restart --- src/sparsevllm/engine/cache_manager/h2o.py | 15 ++-- tests/test_h2o_chunk_numerics.py | 91 ++++++++++++++++++++++ tests/test_h2o_per_head_prefill.py | 17 ++++ 3 files changed, 118 insertions(+), 5 deletions(-) diff --git a/src/sparsevllm/engine/cache_manager/h2o.py b/src/sparsevllm/engine/cache_manager/h2o.py index 087c3057..4415984b 100644 --- a/src/sparsevllm/engine/cache_manager/h2o.py +++ b/src/sparsevllm/engine/cache_manager/h2o.py @@ -705,6 +705,16 @@ def _prepare_prefill(self, seqs: list[Sequence]): logical_start = int(seq.num_prefilled_tokens) logical_end = logical_start + chunk_size if logical_start == 0: + # Reject a stale restart before discarding the resident + # score/position history needed to resume or release it. + for layer_idx in layer_ids: + row = self.seq_id_to_row[layer_idx].get(int(seq.seq_id)) + if row is not None and int(self.row_seq_lens[layer_idx][row]) != 0: + raise RuntimeError( + "H2O first prefill chunk found a non-empty physical row: " + f"layer={layer_idx} seq_id={seq.seq_id} " + f"physical_len={int(self.row_seq_lens[layer_idx][row])}." + ) for layer_idx in score_layer_ids: self._h2o_scores.pop(self._score_key(layer_idx, seq.seq_id), None) self._h2o_positions.pop(self._score_key(layer_idx, seq.seq_id), None) @@ -712,11 +722,6 @@ def _prepare_prefill(self, seqs: list[Sequence]): for layer_idx in layer_ids: row_idx = self._get_free_row(layer_idx, int(seq.seq_id)) physical_start = int(self.row_seq_lens[layer_idx][row_idx]) - if logical_start == 0 and physical_start != 0: - raise RuntimeError( - "H2O first prefill chunk found a non-empty physical row: " - f"layer={layer_idx} seq_id={seq.seq_id} physical_len={physical_start}." - ) if logical_start > 0 and layer_idx in score_layer_ids: self._require_score_length(layer_idx, seq, physical_start) key = self._score_key(layer_idx, seq.seq_id) diff --git a/tests/test_h2o_chunk_numerics.py b/tests/test_h2o_chunk_numerics.py index fa11c702..f19b8745 100644 --- a/tests/test_h2o_chunk_numerics.py +++ b/tests/test_h2o_chunk_numerics.py @@ -155,3 +155,94 @@ def test_no_eviction_full_and_chunked_prefill_are_equivalent(kv_heads): output, scores = _run_chunks(q, k, v, chunks, budget=149, reduction='max') torch.testing.assert_close(output, full_output, rtol=2e-2, atol=2e-2) torch.testing.assert_close(scores, full_scores, rtol=5e-3, atol=5e-3) + + +@pytest.mark.parametrize('budget', [9, 37]) +def test_mla_chunk_scores_and_outputs_follow_shared_latent_retention(budget): + from sparsevllm.engine.cache_manager.base import MlaLatentPayload + from sparsevllm.engine.cache_manager.storage import MlaLatentStorage + from test_mla_attention_layer import _attention + + torch.manual_seed(73) + attention = _attention(device='cuda', tp_size=1) + heads, length = attention.spec.local_q_heads, 37 + latent = torch.randn(length, 512, device='cuda', dtype=torch.bfloat16) + rope = torch.randn(length, 64, device='cuda', dtype=torch.bfloat16) + weights = torch.randn(512, heads * 448, device='cuda', dtype=torch.bfloat16) * .025 + original_projection = (latent @ weights).reshape(length, heads, 448) + q = torch.randn(length, heads, 256, device='cuda', dtype=torch.bfloat16) + + def run(chunks): + manager, runtime = _manager(heads, 1, length, budget, 'max') + storage = MlaLatentStorage(kv_lora_rank=512, rope_dim=64, dtype=torch.bfloat16) + storage.allocate(num_layers=1, num_slots=length + 1, device=manager.device) + manager.attention_cache_storage = storage + seq = Sequence(list(range(length))) + seq.seq_id = 7 + history, outputs = [], [] + cumulative = torch.zeros(heads, length, device='cuda') + start = 0 + try: + for chunk in chunks: + seq.num_prefilled_tokens, seq.current_chunk_size = start, chunk + manager._prepare_prefill([seq]) + state = manager.layer_batch_states[0] + slots = state.slot_mapping.long() + storage.latent_cache[0, slots, 0] = latent[start:start + chunk] + storage.rope_cache[0, slots, 0] = rope[start:start + chunk] + history.extend(range(start, start + chunk)) + zero = torch.zeros(1, device='cuda', dtype=torch.int32) + chunk_tensor = torch.tensor([chunk], device='cuda', dtype=torch.int32) + set_context(True, cache_manager=manager, seqs=[seq]) + view = PrefillComputeView( + meta=AttentionViewMeta(active_slots=manager.buffer_req_to_token_slots[0], + req_indices=zero, context_lens=state.context_lens, + max_context_len=len(history)), + payload=MlaLatentPayload(latent_cache=storage.latent_cache[0], rope_cache=storage.rope_cache[0]), + ) + gathered = attention.prepare_prefill_history(view, query_tokens=chunk) + projected = (gathered.gathered_latent @ weights).reshape(len(history), heads, 448) + workset = attention.bind_prefill_kv( + gathered, + expanded_k=torch.cat((projected[..., :192], gathered.gathered_rope[:, None].expand(-1, heads, -1)), -1), + expanded_v=projected[..., 192:].contiguous(), + ) + query = q[start:start + chunk] + output = attention.run_prefill(query, workset, b_start_loc=zero, chunk_lens=chunk_tensor) + expected = torch.empty_like(query, dtype=torch.float32) + mask = torch.tensor(history, device='cuda')[None] > torch.arange(start, start + chunk, device='cuda')[:, None] + for head in range(heads): + # Original latent projection and separate RoPE term form + # an oracle independent of the physical gather/compaction. + logits = (query[:, head, :192].float() @ original_projection[history, head, :192].float().T + + query[:, head, 192:].float() @ rope[history].float().T) * attention.spec.softmax_scale + probability = logits.masked_fill(mask, -torch.inf).softmax(-1) + expected[:, head] = probability @ original_projection[history, head, 192:].float() + cumulative[head, history] += probability.sum(0) + torch.testing.assert_close(output.float(), expected, rtol=3e-2, atol=3e-2) + runtime.collect_prefill_attention_score(PrefillScoreEvent( + 0, query, attention.build_prefill_explicit_view(workset), zero, chunk_tensor, attention.spec.softmax_scale, + )) + torch.testing.assert_close(manager._h2o_scores[0, 7], cumulative[:, history], rtol=5e-3, atol=5e-3) + runtime.finish_step(SparseStepContext([seq], True, None)) + if len(history) > budget: + recent = max(1, int(budget * .25)) + heavy = sorted(history[:-recent], key=lambda token: (-float(cumulative[:, token].max()), token))[:budget - recent] + history = sorted(heavy + history[-recent:]) + assert manager._h2o_positions[0, 7].tolist() == [history] + torch.testing.assert_close(manager._h2o_scores[0, 7], cumulative[:, history], rtol=5e-3, atol=5e-3) + retained = manager.buffer_req_to_token_slots[0][0, :len(history)].long() + torch.testing.assert_close(storage.latent_cache[0, retained, 0], latent[history]) + torch.testing.assert_close(storage.rope_cache[0, retained, 0], rope[history]) + assert manager._num_free_slots[0] + len(history) == length + 1 + outputs.append(output) + start += chunk + return torch.cat(outputs), manager._h2o_scores[0, 7].clone() + finally: + reset_context() + + output, scores = run([13, 17, 7]) + if budget >= length: + full_output, full_scores = run([length]) + torch.testing.assert_close(output, full_output, rtol=3e-2, atol=3e-2) + torch.testing.assert_close(scores, full_scores, rtol=5e-3, atol=5e-3) diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py index 1ac10275..d4da30a4 100644 --- a/tests/test_h2o_per_head_prefill.py +++ b/tests/test_h2o_per_head_prefill.py @@ -324,3 +324,20 @@ def test_mla_rejects_tp_without_cross_rank_head_reduction(monkeypatch): monkeypatch.setattr(SnapKVCacheManager, '_get_available_slots_info', lambda self: (2**30, 1152)) with pytest.raises(ValueError, match='requires TP1'): manager._get_available_slots_info() + + +def test_rejected_first_chunk_preserves_resident_head_history(): + from sparsevllm.engine.sequence import Sequence + + manager = manager_with_storage() + manager.num_layers = manager.num_kv_layers = 1 + manager.runtime_layout.kv_idx_to_layer_idx = (0,) + seq = Sequence(list(range(6))) + seq.seq_id, seq.current_chunk_size = 7, 3 + scores, positions = manager._h2o_scores[0, 7], manager._h2o_positions[0, 7] + with pytest.raises(RuntimeError, match='non-empty physical row'): + manager._prepare_prefill([seq]) + assert manager._h2o_scores[0, 7] is scores + assert manager._h2o_positions[0, 7] is positions + assert manager.row_seq_lens[0][0] == 6 + assert manager._num_free_slots == [6] From f5007cc634724764c208054c1ced8db5db33d45c Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 18:04:07 +0800 Subject: [PATCH 06/11] test: check native decode reads after h2o head compaction --- tests/test_h2o_chunk_numerics.py | 20 ++++++++++++++------ 1 file changed, 14 insertions(+), 6 deletions(-) diff --git a/tests/test_h2o_chunk_numerics.py b/tests/test_h2o_chunk_numerics.py index f19b8745..426bb69d 100644 --- a/tests/test_h2o_chunk_numerics.py +++ b/tests/test_h2o_chunk_numerics.py @@ -7,7 +7,7 @@ import torch from sparsevllm.engine.cache_manager.base import ( - AttentionViewMeta, ExplicitKVPayload, LayerBatchStates, PrefillComputeView, + AttentionViewMeta, DecodeComputeView, ExplicitKVPayload, LayerBatchStates, PrefillComputeView, ) from sparsevllm.engine.cache_manager.h2o import H2OCacheManager from sparsevllm.engine.cache_manager.storage import ExplicitKVStorage @@ -16,6 +16,7 @@ from sparsevllm.engine.sparse_methods.h2o import H2ORuntime from sparsevllm.kernels.triton.context_flashattention_nopad import context_attention_fwd from sparsevllm.kernels.triton.prefill_score import PrefillScoreWorkspace +from sparsevllm.operators.decode_attention import DecodeAttentionOpSpec, prepare_decode_attention_op from sparsevllm.utils.context import reset_context, set_context pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason='requires CUDA') @@ -118,13 +119,20 @@ def importance(token): assert manager._num_free_slots[0] + len(histories[0]) == length + 1 output_chunks.append(output) start += chunk - # A one-query score-free read is the causal attention contract used at - # decode handoff; model tests separately exercise the native provider. + # Exercise the registered native decode provider on the compressed KV. resident = len(histories[0]) context = torch.tensor([resident], device=q.device, dtype=torch.int32) - output = torch.empty_like(q[:1]) - context_attention_fwd(q[:1], storage.cache[0, 0], storage.cache[1, 0], output, - zero, zero, context, context - 1, 1, view.meta.active_slots) + set_context(False, cache_manager=manager, seqs=[seq]) + decode = prepare_decode_attention_op(DecodeAttentionOpSpec( + num_query_heads=heads, num_kv_heads=kv_heads, head_dim=dim, + activation_dtype=q.dtype, softmax_scale=dim ** -.5, + max_batch_size=1, cuda_graph=False, layer_varying_page_table=True, + ), device_index=0) + output = decode.run(q[:1], DecodeComputeView( + meta=AttentionViewMeta(active_slots=view.meta.active_slots, req_indices=zero, + context_lens=context, max_context_len=resident), + payload=view.payload, + )) for head in range(heads): group = head // group_size keys, values = k[histories[group], group].float(), v[histories[group], group].float() From 4e75dafd0de54a873ecc29079eb0bc13bae165e8 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 18:04:44 +0800 Subject: [PATCH 07/11] test: release prepared decode operators after validation --- tests/test_h2o_chunk_numerics.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/tests/test_h2o_chunk_numerics.py b/tests/test_h2o_chunk_numerics.py index 426bb69d..c324aef1 100644 --- a/tests/test_h2o_chunk_numerics.py +++ b/tests/test_h2o_chunk_numerics.py @@ -60,6 +60,7 @@ def _run_chunks(q, k, v, chunks, budget, reduction): score_reference = [dict() for _ in range(heads)] output_chunks = [] start = 0 + decode = None try: for chunk in chunks: seq.num_prefilled_tokens, seq.current_chunk_size = start, chunk @@ -140,6 +141,8 @@ def importance(token): torch.testing.assert_close(output[0, head].float(), expected, rtol=2e-2, atol=2e-2) return torch.cat(output_chunks), manager._h2o_scores[0, 7].clone() finally: + if decode is not None: + decode.close() reset_context() From cd4b4abc15567d96f2f69ce71970bdec6d3baecc Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Sun, 6 Sep 2026 18:09:37 +0800 Subject: [PATCH 08/11] fix: reject unsupported flashinfer decode merge widths --- src/sparsevllm/operators/decode_attention.py | 3 +++ tests/test_decode_attention_provider.py | 16 ++++++++++++++++ 2 files changed, 19 insertions(+) diff --git a/src/sparsevllm/operators/decode_attention.py b/src/sparsevllm/operators/decode_attention.py index 605e2109..9a4f2ee0 100644 --- a/src/sparsevllm/operators/decode_attention.py +++ b/src/sparsevllm/operators/decode_attention.py @@ -506,6 +506,9 @@ class FlashInferPagedDecodeAttentionProvider(DecodeAttentionProvider): capabilities = AttentionKernelCapabilities( platforms=frozenset({PlatformEnum.CUDA}), activation_dtypes=frozenset({torch.bfloat16, torch.float16}), + # Split-KV merge uses FlashInfer's DISPATCH_HEAD_DIM, even when the + # attention JIT itself can compile another width for a short context. + head_dims=frozenset({64, 128, 256, 512}), page_sizes=None, score_outputs=frozenset({AttentionScoreKind.NONE}), returns_softmax_lse=True, diff --git a/tests/test_decode_attention_provider.py b/tests/test_decode_attention_provider.py index a50cd813..25b9332d 100644 --- a/tests/test_decode_attention_provider.py +++ b/tests/test_decode_attention_provider.py @@ -405,6 +405,22 @@ def test_flashinfer_lse_decode_accepts_cuda_graph_contract(): support.assert_called_once_with() +def test_flashinfer_rejects_width_that_fails_split_kv_merge(): + # Observed with a 149-token H2O handoff: the 32-wide attention JIT runs + # for short rows, but VariableLengthMergeStates rejects this head width. + with patch( + "sparsevllm.operators.decode_attention.flashinfer_paged_decode_support", + return_value=(True, "available"), + ) as dependency: + result = FlashInferPagedDecodeAttentionProvider.supports( + _spec(head_dim=32, softmax_scale=32**-.5, cuda_graph=False), + _cuda_caps(device_name="NVIDIA H20", compute_capability=(9, 0)), + ) + assert not result.supported + assert 'head_dim' in result.reason + dependency.assert_not_called() + + def test_prepared_h2o_decode_applies_fixed_probability_scorer(): spec = _spec( may_require_attention_scores=True, From 3e175201466eb11b71098423ba717f07a9c91b76 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Mon, 7 Sep 2026 21:03:18 +0800 Subject: [PATCH 09/11] refactor: consolidate h2o code and restore shared cache sizing --- src/sparsevllm/engine/cache_manager/h2o.py | 147 ++++++++++++++---- .../engine/cache_manager/h2o_retention.py | 128 --------------- src/sparsevllm/engine/sparse_methods/h2o.py | 49 +++++- .../engine/sparse_methods/h2o_selection.py | 42 ----- tests/test_h2o_per_head_prefill.py | 34 +--- 5 files changed, 163 insertions(+), 237 deletions(-) delete mode 100644 src/sparsevllm/engine/cache_manager/h2o_retention.py delete mode 100644 src/sparsevllm/engine/sparse_methods/h2o_selection.py diff --git a/src/sparsevllm/engine/cache_manager/h2o.py b/src/sparsevllm/engine/cache_manager/h2o.py index 4415984b..b8239f65 100644 --- a/src/sparsevllm/engine/cache_manager/h2o.py +++ b/src/sparsevllm/engine/cache_manager/h2o.py @@ -1,6 +1,7 @@ from __future__ import annotations from collections import deque +from dataclasses import dataclass from typing import NamedTuple import numpy as np @@ -16,7 +17,17 @@ from sparsevllm.utils.profiler import profiler from .snapkv import SnapKVCacheManager -from .h2o_retention import H2OPrefillRetentionMixin, H2O_PREFILL_QUERY_TILE + + +@dataclass(frozen=True) +class H2ORetention: + """One layer/request selection in the current packed cache coordinates.""" + + layer_idx: int + seq_id: int + source_length: int + keep: torch.Tensor # [native KV heads, budget], or [1, budget] for MLA + final_prefill: bool class _H2ORowRef(NamedTuple): @@ -25,7 +36,7 @@ class _H2ORowRef(NamedTuple): seq_id: int -class H2OCacheManager(H2OPrefillRetentionMixin, SnapKVCacheManager): +class H2OCacheManager(SnapKVCacheManager): """Native KV storage with cumulative per-query-head prefill probabilities.""" def __init__( @@ -91,7 +102,6 @@ def h2o_decode_enabled(self) -> bool: def _get_available_slots_info(self) -> tuple[int, int]: from .storage import MlaLatentStorage - available, payload_bytes = super()._get_available_slots_info() mla = isinstance(self.attention_cache_storage, MlaLatentStorage) if mla and self.tp_size != 1: raise ValueError("H2O MLA prefill currently requires TP1 for layer-wide head reduction.") @@ -100,31 +110,112 @@ def _get_available_slots_info(self) -> tuple[int, int]: if heads <= 0 or groups <= 0 or heads % groups: raise ValueError("H2O requires complete local query-to-KV head groups.") - # Account for score/position retention copies and selection indices in - # addition to native payload. Decode may grow the cache but adds no scores. - metadata_per_slot = 20 * heads + 64 * groups + 64 - batch = min(int(self.max_buffer_rows), int(self.config.max_num_seqs_in_batch)) - chunk = min(int(self.config.engine_prefill_chunk_size), int(self.max_model_len)) - width = min(int(self.max_model_len), self.h2o_prefill_budget + chunk) - if bool(getattr(self.config, "enable_prefix_caching", False)): - width = int(self.max_model_len) - window = int(self.config.h2o_prefill_score_window) - queries = min(H2O_PREFILL_QUERY_TILE, chunk, window or chunk) - # Probability QK stats use at most two padded head rows per real head - # and at most one bounded query tile at a time. - padded_queries = max(16, 1 << (queries - 1).bit_length()) - stats = 8 * batch * (2 * heads) * padded_queries * ((width + 63) // 64 + 1) - scores = 8 * batch * heads * width - # Include replacement allocation while the previous workspace is live. - workspace_bytes = 2 * (stats + scores) + 2 * self.h2o_prefill_budget * payload_bytes - if workspace_bytes >= available: - raise RuntimeError( - "Not enough memory for H2O prefill score/retention workspaces: " - f"required={workspace_bytes} available={available}." + return super()._get_available_slots_info() + + @staticmethod + def _assert_retention_tensor(condition: torch.Tensor, message: str) -> None: + if condition.is_cuda: + torch._assert_async(condition) + elif not bool(condition.item()): + raise RuntimeError(message) + + @property + def h2o_selection_groups(self) -> int: + from .storage import MlaLatentStorage + + return ( + 1 if isinstance(self.attention_cache_storage, MlaLatentStorage) + else self.num_kv_heads + ) + + def commit_h2o_retention(self, requests: list[H2ORetention]) -> None: + from .storage import ExplicitKVStorage + + prepared = [] + seen = set() + seen_rows = set() + release_counts = {} + storage = self.attention_cache_storage + for request in requests: + layer, seq_id, length = request.layer_idx, request.seq_id, request.source_length + key = (layer, seq_id) + if key in seen: + raise ValueError("Duplicate H2O retention request.") + seen.add(key) + row = self.seq_id_to_row[layer][seq_id] + if (layer, row) in seen_rows: + raise ValueError("H2O retention requests share a physical row.") + seen_rows.add((layer, row)) + if int(self.row_seq_lens[layer][row]) != length: + raise ValueError("H2O retention refers to a stale physical row.") + keep = request.keep + groups = self.h2o_selection_groups + if ( + keep.ndim != 2 or keep.shape[0] != groups + or keep.dtype != torch.long or keep.device != self.device + or not 0 < keep.shape[-1] < length + ): + raise ValueError("H2O retention requires [selection_groups, budget] int64 indices.") + budget = int(keep.shape[-1]) + self._assert_retention_tensor( + ((keep >= 0) & (keep < length)).all() + & (keep[:, 1:] > keep[:, :-1]).all(), + "H2O retention indices must be in bounds and strictly increasing.", ) - self._h2o_reserved_workspace_bytes = workspace_bytes - self._h2o_metadata_bytes_per_slot = metadata_per_slot - return available - workspace_bytes, payload_bytes + metadata_per_slot + score = self._h2o_scores[key] + positions = self._h2o_positions[key] + if ( + score.ndim != 2 or score.shape[-1] != length + or score.shape[0] % groups or tuple(positions.shape) != (groups, length) + ): + raise ValueError("H2O retention score/position metadata is not aligned.") + slots = self.buffer_req_to_token_slots[layer][row, :length].long().clone() + ordered_slots = slots.sort().values + self._assert_retention_tensor( + ((slots >= 0) & (slots < storage.slot_capacity())).all() + & (ordered_slots[1:] > ordered_slots[:-1]).all(), + "H2O retention physical slots must be valid and unique.", + ) + release_counts[layer] = release_counts.get(layer, 0) + length - budget + pointer = int(self._num_free_slots[layer]) + end = pointer + release_counts[layer] + if pointer < 0 or end > self.free_slots_stack[layer].numel(): + raise RuntimeError(f"H2O retention would overflow the free-slot stack: layer={layer}.") + # Every query head keeps its own history, including heads that did + # not supply the group's maximum on this step. + score_keep = keep.repeat_interleave(score.shape[0] // groups, dim=0) + kept_score = score.gather(1, score_keep).contiguous() + kept_positions = positions.gather(1, keep).contiguous() + prepared.append((request, row, slots, ordered_slots, kept_score, kept_positions)) + + # Validate the complete submission before publishing any row mutation. + for request, row, slots, ordered_slots, score, positions in prepared: + layer, seq_id, length = request.layer_idx, request.seq_id, request.source_length + keep = request.keep + budget = int(keep.shape[-1]) + destination = ordered_slots[:budget] + released = ordered_slots[budget:] + kv_layer = self.kv_layer_index(layer) + selected = slots[keep] + if isinstance(storage, ExplicitKVStorage): + storage.copy_head_slots(kv_layer, selected, destination) + else: + storage.copy_slots(kv_layer, selected[0], destination) + ptr = int(self._num_free_slots[layer]) + self.free_slots_stack[layer][ptr:ptr + released.numel()] = released + self._num_free_slots[layer] = ptr + released.numel() + self.buffer_req_to_token_slots[layer][row, :budget] = destination + self.buffer_req_to_token_slots[layer][row, budget:length] = 0 + self.row_seq_lens[layer][row] = budget + self._h2o_scores[(layer, seq_id)] = score + self._h2o_positions[(layer, seq_id)] = positions + counter = "final_prefill_evictions" if request.final_prefill else "intermediate_prefill_evictions" + self._h2o_counters[counter] += 1 + self._h2o_counters["dropped_tokens"] += length - budget + if prepared: + self._uniform_decode_metadata = False + self._decode_static_state_binding_key = None + self._invalidate_h2o_decode_score_workspace() def _iter_accounting_tensors(self): yield from super()._iter_accounting_tensors() @@ -1290,7 +1381,5 @@ def debug_state_summary(self) -> dict[str, object]: f"{layer_idx}:{seq_id}": int(score.shape[-1]) for (layer_idx, seq_id), score in sorted(self._h2o_scores.items()) }, - "reserved_workspace_bytes": getattr(self, "_h2o_reserved_workspace_bytes", 0), - "metadata_bytes_per_slot": getattr(self, "_h2o_metadata_bytes_per_slot", 0), } return summary diff --git a/src/sparsevllm/engine/cache_manager/h2o_retention.py b/src/sparsevllm/engine/cache_manager/h2o_retention.py deleted file mode 100644 index 19ad2134..00000000 --- a/src/sparsevllm/engine/cache_manager/h2o_retention.py +++ /dev/null @@ -1,128 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass - -import torch - - -H2O_PREFILL_QUERY_TILE = 128 - - -@dataclass(frozen=True) -class H2ORetention: - """One layer/request selection in the current packed cache coordinates.""" - - layer_idx: int - seq_id: int - source_length: int - keep: torch.Tensor # [native KV heads, budget], or [1, budget] for MLA - final_prefill: bool - - -class H2OPrefillRetentionMixin: - """Cache-owned per-head state and overlap-safe physical retention.""" - - @staticmethod - def _assert_retention_tensor(condition: torch.Tensor, message: str) -> None: - if condition.is_cuda: - torch._assert_async(condition) - elif not bool(condition.item()): - raise RuntimeError(message) - - @property - def h2o_selection_groups(self) -> int: - from .storage import MlaLatentStorage - - return ( - 1 if isinstance(self.attention_cache_storage, MlaLatentStorage) - else self.num_kv_heads - ) - - def commit_h2o_retention(self, requests: list[H2ORetention]) -> None: - from .storage import ExplicitKVStorage - - prepared = [] - seen = set() - seen_rows = set() - release_counts = {} - storage = self.attention_cache_storage - for request in requests: - layer, seq_id, length = request.layer_idx, request.seq_id, request.source_length - key = (layer, seq_id) - if key in seen: - raise ValueError("Duplicate H2O retention request.") - seen.add(key) - row = self.seq_id_to_row[layer][seq_id] - if (layer, row) in seen_rows: - raise ValueError("H2O retention requests share a physical row.") - seen_rows.add((layer, row)) - if int(self.row_seq_lens[layer][row]) != length: - raise ValueError("H2O retention refers to a stale physical row.") - keep = request.keep - groups = self.h2o_selection_groups - if ( - keep.ndim != 2 or keep.shape[0] != groups - or keep.dtype != torch.long or keep.device != self.device - or not 0 < keep.shape[-1] < length - ): - raise ValueError("H2O retention requires [selection_groups, budget] int64 indices.") - budget = int(keep.shape[-1]) - self._assert_retention_tensor( - ((keep >= 0) & (keep < length)).all() - & (keep[:, 1:] > keep[:, :-1]).all(), - "H2O retention indices must be in bounds and strictly increasing.", - ) - score = self._h2o_scores[key] - positions = self._h2o_positions[key] - if ( - score.ndim != 2 or score.shape[-1] != length - or score.shape[0] % groups or tuple(positions.shape) != (groups, length) - ): - raise ValueError("H2O retention score/position metadata is not aligned.") - slots = self.buffer_req_to_token_slots[layer][row, :length].long().clone() - ordered_slots = slots.sort().values - self._assert_retention_tensor( - ((slots >= 0) & (slots < storage.slot_capacity())).all() - & (ordered_slots[1:] > ordered_slots[:-1]).all(), - "H2O retention physical slots must be valid and unique.", - ) - release_counts[layer] = release_counts.get(layer, 0) + length - budget - pointer = int(self._num_free_slots[layer]) - end = pointer + release_counts[layer] - if pointer < 0 or end > self.free_slots_stack[layer].numel(): - raise RuntimeError(f"H2O retention would overflow the free-slot stack: layer={layer}.") - # Every query head keeps its own history, including heads that did - # not supply the group's maximum on this step. - score_keep = keep.repeat_interleave(score.shape[0] // groups, dim=0) - kept_score = score.gather(1, score_keep).contiguous() - kept_positions = positions.gather(1, keep).contiguous() - prepared.append((request, row, slots, ordered_slots, kept_score, kept_positions)) - - # Validate the complete submission before publishing any row mutation. - for request, row, slots, ordered_slots, score, positions in prepared: - layer, seq_id, length = request.layer_idx, request.seq_id, request.source_length - keep = request.keep - budget = int(keep.shape[-1]) - destination = ordered_slots[:budget] - released = ordered_slots[budget:] - kv_layer = self.kv_layer_index(layer) - selected = slots[keep] - if isinstance(storage, ExplicitKVStorage): - storage.copy_head_slots(kv_layer, selected, destination) - else: - storage.copy_slots(kv_layer, selected[0], destination) - ptr = int(self._num_free_slots[layer]) - self.free_slots_stack[layer][ptr:ptr + released.numel()] = released - self._num_free_slots[layer] = ptr + released.numel() - self.buffer_req_to_token_slots[layer][row, :budget] = destination - self.buffer_req_to_token_slots[layer][row, budget:length] = 0 - self.row_seq_lens[layer][row] = budget - self._h2o_scores[(layer, seq_id)] = score - self._h2o_positions[(layer, seq_id)] = positions - counter = "final_prefill_evictions" if request.final_prefill else "intermediate_prefill_evictions" - self._h2o_counters[counter] += 1 - self._h2o_counters["dropped_tokens"] += length - budget - if prepared: - self._uniform_decode_metadata = False - self._decode_static_state_binding_key = None - self._invalidate_h2o_decode_score_workspace() diff --git a/src/sparsevllm/engine/sparse_methods/h2o.py b/src/sparsevllm/engine/sparse_methods/h2o.py index 61dcb5b2..662042de 100644 --- a/src/sparsevllm/engine/sparse_methods/h2o.py +++ b/src/sparsevllm/engine/sparse_methods/h2o.py @@ -3,11 +3,7 @@ import torch from sparsevllm.engine.sequence import Sequence -from sparsevllm.method_registry import ( - normalize_sparse_method, - resolve_prefill_sparse_method, -) -from sparsevllm.engine.cache_manager.h2o_retention import H2ORetention, H2O_PREFILL_QUERY_TILE +from sparsevllm.engine.cache_manager.h2o import H2ORetention from sparsevllm.utils.profiler import profiler from sparsevllm.utils.context import get_context from sparsevllm.engine.cache_manager.base import ExplicitKVPayload @@ -15,7 +11,48 @@ from .base import SparseStepContext, PrefillScoreEvent from .passthrough import PassThroughRuntime -from .h2o_selection import select_h2o_heads + + +H2O_PREFILL_QUERY_TILE = 128 + + +def select_h2o_heads( + cumulative: torch.Tensor, + *, + selection_groups: int, + budget: int, + recent_ratio: float, + reduction: str, +) -> torch.Tensor: + """Rank cumulative [Hq, L] probabilities, returning [groups, min(B,L)]. + + A group is one native KV head for explicit KV, or the whole layer for + shared MLA latent storage. Reduction happens only after time accumulation. + Equal scores prefer older positions, independently of the device top-k. + """ + if cumulative.ndim != 2 or cumulative.shape[0] == 0: + raise ValueError("H2O cumulative scores must have shape [query_heads, length].") + heads, length = cumulative.shape + if selection_groups <= 0 or heads % selection_groups: + raise ValueError("H2O query heads must divide into complete selection groups.") + if budget <= 0 or not 0 < recent_ratio < 1: + raise ValueError("H2O requires a positive budget and recent_ratio in (0, 1).") + if reduction not in {"max", "mean"}: + raise ValueError("H2O head reduction must be 'max' or 'mean'.") + if length <= budget: + return torch.arange(length, device=cumulative.device).expand(selection_groups, -1) + + grouped = cumulative.reshape(selection_groups, heads // selection_groups, length) + ranks = grouped.amax(dim=1) if reduction == "max" else grouped.mean(dim=1) + recent_count = min(budget, max(1, int(budget * recent_ratio))) + recent_start = length - recent_count + heavy = torch.argsort( + ranks[:, :recent_start], dim=-1, descending=True, stable=True, + )[:, :budget - recent_count] + recent = torch.arange(recent_start, length, device=cumulative.device).expand( + selection_groups, -1, + ) + return torch.cat((heavy, recent), dim=-1).sort(dim=-1).values class H2ORuntime(PassThroughRuntime): diff --git a/src/sparsevllm/engine/sparse_methods/h2o_selection.py b/src/sparsevllm/engine/sparse_methods/h2o_selection.py deleted file mode 100644 index d6b78072..00000000 --- a/src/sparsevllm/engine/sparse_methods/h2o_selection.py +++ /dev/null @@ -1,42 +0,0 @@ -from __future__ import annotations - -import torch - - -def select_h2o_heads( - cumulative: torch.Tensor, - *, - selection_groups: int, - budget: int, - recent_ratio: float, - reduction: str, -) -> torch.Tensor: - """Rank cumulative [Hq, L] probabilities, returning [groups, min(B,L)]. - - A group is one native KV head for explicit KV, or the whole layer for - shared MLA latent storage. Reduction happens only after time accumulation. - Equal scores prefer older positions, independently of the device top-k. - """ - if cumulative.ndim != 2 or cumulative.shape[0] == 0: - raise ValueError("H2O cumulative scores must have shape [query_heads, length].") - heads, length = cumulative.shape - if selection_groups <= 0 or heads % selection_groups: - raise ValueError("H2O query heads must divide into complete selection groups.") - if budget <= 0 or not 0 < recent_ratio < 1: - raise ValueError("H2O requires a positive budget and recent_ratio in (0, 1).") - if reduction not in {"max", "mean"}: - raise ValueError("H2O head reduction must be 'max' or 'mean'.") - if length <= budget: - return torch.arange(length, device=cumulative.device).expand(selection_groups, -1) - - grouped = cumulative.reshape(selection_groups, heads // selection_groups, length) - ranks = grouped.amax(dim=1) if reduction == "max" else grouped.mean(dim=1) - recent_count = min(budget, max(1, int(budget * recent_ratio))) - recent_start = length - recent_count - heavy = torch.argsort( - ranks[:, :recent_start], dim=-1, descending=True, stable=True, - )[:, :budget - recent_count] - recent = torch.arange(recent_start, length, device=cumulative.device).expand( - selection_groups, -1, - ) - return torch.cat((heavy, recent), dim=-1).sort(dim=-1).values diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py index d4da30a4..2c53970e 100644 --- a/tests/test_h2o_per_head_prefill.py +++ b/tests/test_h2o_per_head_prefill.py @@ -5,9 +5,9 @@ import torch from sparsevllm.engine.cache_manager.h2o import H2OCacheManager -from sparsevllm.engine.cache_manager.h2o_retention import H2ORetention +from sparsevllm.engine.cache_manager.h2o import H2ORetention from sparsevllm.engine.cache_manager.storage import ExplicitKVStorage, MlaLatentStorage -from sparsevllm.engine.sparse_methods.h2o_selection import select_h2o_heads +from sparsevllm.engine.sparse_methods.h2o import select_h2o_heads def test_late_max_preserves_head_history_across_chunks(): @@ -286,36 +286,6 @@ def test_later_layer_capacity_failure_does_not_commit_earlier_layers(): assert [lengths.tolist() for lengths in manager.row_seq_lens] == [[6, 5], [6, 5]] -@pytest.mark.parametrize('mla,heads,groups,tp', [(False, 8, 8, 1), (False, 32, 2, 1), (False, 16, 2, 2), (True, 32, 1, 1)]) -def test_capacity_reserves_native_payload_and_live_head_metadata(monkeypatch, mla, heads, groups, tp): - from sparsevllm.engine.cache_manager.snapkv import SnapKVCacheManager - - manager = object.__new__(H2OCacheManager) - manager.tp_size, manager.num_kv_heads = tp, groups - manager.hf_config = SimpleNamespace(num_attention_heads=heads * tp) - manager.max_buffer_rows, manager.max_model_len = 3, 2048 - manager.config = SimpleNamespace(max_num_seqs_in_batch=2, engine_prefill_chunk_size=257, - h2o_prefill_budget=129, h2o_prefill_score_window=0, - enable_prefix_caching=False) - storage = (MlaLatentStorage(kv_lora_rank=512, rope_dim=64, dtype=torch.bfloat16) if mla - else ExplicitKVStorage(num_kv_heads=groups, head_dim=128, dtype=torch.bfloat16)) - manager.attention_cache_storage = storage - native_bytes = storage.bytes_per_slot_per_layer() - available = 64 * 1024 * 1024 - monkeypatch.setattr(SnapKVCacheManager, '_get_available_slots_info', lambda self: (available, native_bytes)) - remaining, slot_cost = manager._get_available_slots_info() - # Scores and positions must survive while the newly gathered copies exist. - assert slot_cost >= native_bytes + 2 * (4 * heads + 8 * groups) - assert slot_cost - manager._h2o_metadata_bytes_per_slot == native_bytes - score_buffers = 2 * 4 * 2 * heads * (129 + 257) - copy_buffers = 129 * native_bytes - assert available - remaining >= score_buffers + copy_buffers - assert 0 < remaining < available - monkeypatch.setattr(SnapKVCacheManager, '_get_available_slots_info', lambda self: (1, native_bytes)) - with pytest.raises(RuntimeError, match='Not enough memory'): - manager._get_available_slots_info() - - def test_mla_rejects_tp_without_cross_rank_head_reduction(monkeypatch): from sparsevllm.engine.cache_manager.snapkv import SnapKVCacheManager From e53b9da9220e4b307b353a6646f190899cf94e55 Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Mon, 7 Sep 2026 21:38:58 +0800 Subject: [PATCH 10/11] refactor: use max as the sole h2o head reduction --- docs/en/features/sparse-methods.md | 3 +-- docs/zh/features/sparse-methods.md | 2 +- src/sparsevllm/configs/groups.py | 1 - src/sparsevllm/configs/sparse.py | 4 ---- src/sparsevllm/engine/chain_cache.py | 1 - src/sparsevllm/engine/llm_engine.py | 1 - src/sparsevllm/engine/sparse_methods/h2o.py | 8 ++------ tests/test_chain_prefix_cache.py | 2 -- tests/test_glm_mla_prefix_cache.py | 1 - tests/test_h2o_chunk_numerics.py | 21 ++++++++++----------- tests/test_h2o_model_integration.py | 2 +- tests/test_h2o_per_head_prefill.py | 18 ++++++++---------- 12 files changed, 23 insertions(+), 41 deletions(-) diff --git a/docs/en/features/sparse-methods.md b/docs/en/features/sparse-methods.md index 4dff645b..5197e348 100644 --- a/docs/en/features/sparse-methods.md +++ b/docs/en/features/sparse-methods.md @@ -50,8 +50,7 @@ SnapKV defaults `sparse_prefill_score_mode` to `logits`; `probability` remains an explicit reproducibility option because its additional normalized QK sweep is substantially more expensive in measured long-context prefill. PyramidKV and H2O use `probability`. H2O accumulates FP32 probability sums per query -head across prefill chunks. At eviction, `h2o_head_reduction=max` (default) or -`mean` combines the cumulative scores within each GQA KV group, or across the +head across prefill chunks. At eviction, max combines the cumulative scores within each GQA KV group, or across the whole layer for MLA. MHA heads select independently. Each selection retains heavy hitters plus recent tokens within its token budget; native GQA KV sharing and MLA latent storage are preserved. MLA H2O prefill currently requires TP1. diff --git a/docs/zh/features/sparse-methods.md b/docs/zh/features/sparse-methods.md index 0ea1fb87..d346eb23 100644 --- a/docs/zh/features/sparse-methods.md +++ b/docs/zh/features/sparse-methods.md @@ -39,7 +39,7 @@ prefill attention 计算。它们是同一条轴上的备选项,可以分别 SnapKV 的 `sparse_prefill_score_mode` 默认值改为 `logits`;`probability` 仍可显式启用以复现实验,但它需要额外执行归一化 QK sweep,在已测长上下文 prefill 中开销明显更高。PyramidKV 和 H2O 使用 `probability`。H2O 逐 query head 跨 prefill chunks -累计 FP32 概率和;驱逐时通过 `h2o_head_reduction=max`(默认)或 `mean`, +累计 FP32 概率和;驱逐时通过 max, 在 GQA 的每个 KV 组内或 MLA 的整个 layer 内归约累计分数。MHA 各 head 独立选择。 每套选择在预算内保留 heavy hitters 和 recent tokens,保持 GQA 的原生 KV 共享 以及 MLA 的原生 latent 存储。MLA H2O prefill 当前要求 TP1。 diff --git a/src/sparsevllm/configs/groups.py b/src/sparsevllm/configs/groups.py index 6c1846ac..4dcf9e23 100644 --- a/src/sparsevllm/configs/groups.py +++ b/src/sparsevllm/configs/groups.py @@ -60,7 +60,6 @@ class SparseMethodConfig: h2o_prefill_budget: int = 8192 h2o_recent_ratio: float = 0.5 h2o_prefill_score_window: int = 0 - h2o_head_reduction: str = "max" rkv_compression_interval: int = 128 rkv_observation_tokens: int = 8 diff --git a/src/sparsevllm/configs/sparse.py b/src/sparsevllm/configs/sparse.py index b86959ee..39ea0c91 100644 --- a/src/sparsevllm/configs/sparse.py +++ b/src/sparsevllm/configs/sparse.py @@ -148,10 +148,6 @@ def _normalize_snapkv(config) -> None: def _normalize_h2o(config) -> None: - reduction = str(getattr(config, "h2o_head_reduction", "max")).strip().lower() - if reduction not in {"max", "mean"}: - raise ValueError("h2o_head_reduction must be 'max' or 'mean'.") - config.h2o_head_reduction = reduction if getattr(config, "sparse_prefill_score_mode", "probability") != "probability": raise ValueError("H2O per-head accumulation requires sparse_prefill_score_mode='probability'.") _normalize_positive_int(config, "h2o_decode_budget", fallback=0) diff --git a/src/sparsevllm/engine/chain_cache.py b/src/sparsevllm/engine/chain_cache.py index fac6306a..c0569447 100644 --- a/src/sparsevllm/engine/chain_cache.py +++ b/src/sparsevllm/engine/chain_cache.py @@ -225,7 +225,6 @@ def build_chain_cache_fingerprint(config: Any) -> bytes: "h2o_prefill_budget", "h2o_recent_ratio", "h2o_prefill_score_window", - "h2o_head_reduction", "sparse_prefill_score_mode", "sparse_attn_score_dtype", ), diff --git a/src/sparsevllm/engine/llm_engine.py b/src/sparsevllm/engine/llm_engine.py index 8ebc1922..e0489dac 100644 --- a/src/sparsevllm/engine/llm_engine.py +++ b/src/sparsevllm/engine/llm_engine.py @@ -1294,7 +1294,6 @@ def worker_info( "h2o_prefill_budget", "h2o_recent_ratio", "h2o_prefill_score_window", - "h2o_head_reduction", "pool_kernel_size", "sparse_attn_score_dtype", "pyramid_layer_ratios", diff --git a/src/sparsevllm/engine/sparse_methods/h2o.py b/src/sparsevllm/engine/sparse_methods/h2o.py index 662042de..b5f3051e 100644 --- a/src/sparsevllm/engine/sparse_methods/h2o.py +++ b/src/sparsevllm/engine/sparse_methods/h2o.py @@ -22,12 +22,11 @@ def select_h2o_heads( selection_groups: int, budget: int, recent_ratio: float, - reduction: str, ) -> torch.Tensor: """Rank cumulative [Hq, L] probabilities, returning [groups, min(B,L)]. A group is one native KV head for explicit KV, or the whole layer for - shared MLA latent storage. Reduction happens only after time accumulation. + shared MLA latent storage. Group-wise max follows time accumulation. Equal scores prefer older positions, independently of the device top-k. """ if cumulative.ndim != 2 or cumulative.shape[0] == 0: @@ -37,13 +36,11 @@ def select_h2o_heads( raise ValueError("H2O query heads must divide into complete selection groups.") if budget <= 0 or not 0 < recent_ratio < 1: raise ValueError("H2O requires a positive budget and recent_ratio in (0, 1).") - if reduction not in {"max", "mean"}: - raise ValueError("H2O head reduction must be 'max' or 'mean'.") if length <= budget: return torch.arange(length, device=cumulative.device).expand(selection_groups, -1) grouped = cumulative.reshape(selection_groups, heads // selection_groups, length) - ranks = grouped.amax(dim=1) if reduction == "max" else grouped.mean(dim=1) + ranks = grouped.amax(dim=1) recent_count = min(budget, max(1, int(budget * recent_ratio))) recent_start = length - recent_count heavy = torch.argsort( @@ -224,7 +221,6 @@ def finish_step(self, step: SparseStepContext) -> None: selection_groups=manager.h2o_selection_groups, budget=budget, recent_ratio=float(self.config.h2o_recent_ratio), - reduction=self.config.h2o_head_reduction, ) requests.append(H2ORetention(layer_idx, int(seq.seq_id), length, keep, final)) manager.commit_h2o_retention(requests) diff --git a/tests/test_chain_prefix_cache.py b/tests/test_chain_prefix_cache.py index 62d8a992..5b273d6d 100644 --- a/tests/test_chain_prefix_cache.py +++ b/tests/test_chain_prefix_cache.py @@ -1000,7 +1000,6 @@ def _h2o_fingerprint_config(**overrides): "h2o_prefill_budget": 8, "h2o_recent_ratio": 0.5, "h2o_prefill_score_window": 4, - "h2o_head_reduction": "max", "sparse_attn_score_dtype": "float32", } values.update(overrides) @@ -1015,7 +1014,6 @@ def _h2o_fingerprint_config(**overrides): ("h2o_prefill_budget", 9), ("h2o_recent_ratio", 0.25), ("h2o_prefill_score_window", 8), - ("h2o_head_reduction", "mean"), ("sparse_attn_score_dtype", "float16"), ], ) diff --git a/tests/test_glm_mla_prefix_cache.py b/tests/test_glm_mla_prefix_cache.py index 4750c196..fc203c9d 100644 --- a/tests/test_glm_mla_prefix_cache.py +++ b/tests/test_glm_mla_prefix_cache.py @@ -56,7 +56,6 @@ def _latent_chain_manager(manager_type, method: str): h2o_prefill_budget=8, h2o_recent_ratio=0.5, h2o_prefill_score_window=2, - h2o_head_reduction="max", rkv_compression_interval=2, rkv_observation_tokens=2, rkv_alpha=0.5, diff --git a/tests/test_h2o_chunk_numerics.py b/tests/test_h2o_chunk_numerics.py index c324aef1..ee0a0cb8 100644 --- a/tests/test_h2o_chunk_numerics.py +++ b/tests/test_h2o_chunk_numerics.py @@ -22,7 +22,7 @@ pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason='requires CUDA') -def _manager(heads, kv_heads, length, budget, reduction): +def _manager(heads, kv_heads, length, budget): manager = object.__new__(H2OCacheManager) manager.device = torch.device('cuda:0') manager.num_kv_heads = kv_heads @@ -40,8 +40,7 @@ def _manager(heads, kv_heads, length, budget, reduction): manager._h2o_scores, manager._h2o_positions = {}, {} manager._h2o_counters = dict(final_prefill_evictions=0, intermediate_prefill_evictions=0, dropped_tokens=0) manager.config = SimpleNamespace(h2o_prefill_budget=budget, h2o_decode_budget=budget, - h2o_prefill_score_window=0, h2o_recent_ratio=.25, - h2o_head_reduction=reduction) + h2o_prefill_score_window=0, h2o_recent_ratio=.25) runtime = object.__new__(H2ORuntime) runtime.config, runtime.cache_manager = manager.config, manager runtime._prefill_score_workspace = PrefillScoreWorkspace() @@ -49,11 +48,11 @@ def _manager(heads, kv_heads, length, budget, reduction): return manager, runtime -def _run_chunks(q, k, v, chunks, budget, reduction): +def _run_chunks(q, k, v, chunks, budget): length, heads, dim = q.shape kv_heads = k.shape[1] group_size = heads // kv_heads - manager, runtime = _manager(heads, kv_heads, length, budget, reduction) + manager, runtime = _manager(heads, kv_heads, length, budget) seq = Sequence(list(range(length))) seq.seq_id = 7 histories = [[] for _ in range(kv_heads)] @@ -110,7 +109,7 @@ def _run_chunks(q, k, v, chunks, budget, reduction): for group, history in enumerate(histories): def importance(token): scores = [score_reference[h][token] for h in range(group * group_size, (group + 1) * group_size)] - return max(scores) if reduction == 'max' else sum(scores) / group_size + return max(scores) heavy = sorted(history[:-recent], key=lambda token: (-importance(token), token))[:budget - recent] histories[group] = sorted(heavy + history[-recent:]) assert manager._h2o_positions[0, 7].tolist() == histories @@ -146,13 +145,13 @@ def importance(token): reset_context() -@pytest.mark.parametrize('kv_heads,reduction', [(4, 'max'), (2, 'max'), (2, 'mean')]) -def test_chunk_attention_and_retention_match_logical_token_oracle(kv_heads, reduction): +@pytest.mark.parametrize('kv_heads', [4, 2]) +def test_chunk_attention_and_retention_match_logical_token_oracle(kv_heads): torch.manual_seed(57) q = torch.randn(35, 4, 32, device='cuda', dtype=torch.bfloat16) k = torch.randn(35, kv_heads, 32, device='cuda', dtype=torch.bfloat16) v = torch.randn_like(k) - _run_chunks(q, k, v, [11, 7, 17], budget=9, reduction=reduction) + _run_chunks(q, k, v, [11, 7, 17], budget=9) @pytest.mark.parametrize('kv_heads', [4, 2]) @@ -161,9 +160,9 @@ def test_no_eviction_full_and_chunked_prefill_are_equivalent(kv_heads): q = torch.randn(149, 4, 32, device='cuda', dtype=torch.bfloat16) k = torch.randn(149, kv_heads, 32, device='cuda', dtype=torch.bfloat16) v = torch.randn_like(k) - full_output, full_scores = _run_chunks(q, k, v, [149], budget=149, reduction='max') + full_output, full_scores = _run_chunks(q, k, v, [149], budget=149) for chunks in ([13, 129, 7], [71, 78]): - output, scores = _run_chunks(q, k, v, chunks, budget=149, reduction='max') + output, scores = _run_chunks(q, k, v, chunks, budget=149) torch.testing.assert_close(output, full_output, rtol=2e-2, atol=2e-2) torch.testing.assert_close(scores, full_scores, rtol=5e-3, atol=5e-3) diff --git a/tests/test_h2o_model_integration.py b/tests/test_h2o_model_integration.py index 35ac4667..a63f918f 100644 --- a/tests/test_h2o_model_integration.py +++ b/tests/test_h2o_model_integration.py @@ -22,7 +22,7 @@ def test_model_chunk_prefill_and_native_decode_handoff(tmp_path): model = os.environ['SPARSEVLLM_H2O_TEST_MODEL'] output_dir = Path(os.environ.get('SPARSEVLLM_H2O_TEST_OUTPUT', str(tmp_path))) output_dir.mkdir(parents=True, exist_ok=True) - config = dict(sparse_method='h2o', h2o_head_reduction='max', + config = dict(sparse_method='h2o', h2o_prefill_budget=64, h2o_decode_budget=32, h2o_prefill_score_window=0, engine_prefill_chunk_size=64, max_num_batched_tokens=128, max_num_seqs_in_batch=2, diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py index 2c53970e..115f33a9 100644 --- a/tests/test_h2o_per_head_prefill.py +++ b/tests/test_h2o_per_head_prefill.py @@ -22,25 +22,24 @@ def test_late_max_preserves_head_history_across_chunks(): cumulative, chunk, new_len=3, weight=1, ) torch.testing.assert_close(cumulative, p.sum(0)) - keep = select_h2o_heads(cumulative, selection_groups=1, budget=2, recent_ratio=.5, reduction='max') + keep = select_h2o_heads(cumulative, selection_groups=1, budget=2, recent_ratio=.5) assert keep.tolist() == [[1, 2]] assert p.amax(1).sum(0)[0] > p.amax(1).sum(0)[1] -def test_groups_select_independently_and_mean_changes_only_ranking(): +def test_groups_select_independently_with_max_ranking(): cumulative = torch.tensor([[9., 6., 0., 0.], [0., 6., 0., 0.], [0., 0., 8., 0.], [0., 0., 7., 0.]]) args = dict(selection_groups=2, budget=2, recent_ratio=.5) - assert select_h2o_heads(cumulative, reduction='max', **args).tolist() == [[0, 3], [2, 3]] - assert select_h2o_heads(cumulative, reduction='mean', **args).tolist() == [[1, 3], [2, 3]] - mha = select_h2o_heads(cumulative, selection_groups=4, budget=2, recent_ratio=.5, reduction='max') + assert select_h2o_heads(cumulative, **args).tolist() == [[0, 3], [2, 3]] + mha = select_h2o_heads(cumulative, selection_groups=4, budget=2, recent_ratio=.5) assert mha.tolist() == [[0, 3], [1, 3], [2, 3], [2, 3]] def test_ties_and_recent_suffix_fill_budget_in_logical_order(): scores = torch.zeros(3, 9) - keep = select_h2o_heads(scores, selection_groups=3, budget=4, recent_ratio=.5, reduction='max') + keep = select_h2o_heads(scores, selection_groups=3, budget=4, recent_ratio=.5) assert keep.tolist() == [[0, 1, 7, 8]] * 3 - all_tokens = select_h2o_heads(scores, selection_groups=3, budget=12, recent_ratio=.5, reduction='mean') + all_tokens = select_h2o_heads(scores, selection_groups=3, budget=12, recent_ratio=.5) assert all_tokens.tolist() == [list(range(9))] * 3 @@ -139,8 +138,7 @@ def test_chunk_prefill_append_retention_and_final_handoff_follow_logical_positio manager._h2o_scores.clear() manager._h2o_positions.clear() manager.config = SimpleNamespace(h2o_prefill_budget=4, h2o_decode_budget=3, - h2o_prefill_score_window=0, h2o_recent_ratio=.5, - h2o_head_reduction='max') + h2o_prefill_score_window=0, h2o_recent_ratio=.5) runtime = object.__new__(H2ORuntime) runtime.config = manager.config runtime.cache_manager = manager @@ -203,7 +201,7 @@ def _multi_request_manager(): cache = manager.attention_cache_storage.cache cache.copy_(torch.arange(cache.numel()).reshape_as(cache)) manager.config = SimpleNamespace(h2o_prefill_budget=4, h2o_decode_budget=3, - h2o_recent_ratio=.5, h2o_head_reduction='max') + h2o_recent_ratio=.5) manager.seq_id_to_row = [{7: 0, 8: 1}, {7: 0, 8: 1}] manager.row_seq_lens = [np.array([6, 5], dtype=np.int32) for _ in range(2)] manager.buffer_req_to_token_slots = [torch.zeros(2, 8, dtype=torch.int32) for _ in range(2)] From c29c58e3c08b92b66b9fefcbdc1cbd7a862d721d Mon Sep 17 00:00:00 2001 From: QuanshengGu Date: Mon, 7 Sep 2026 21:45:23 +0800 Subject: [PATCH 11/11] fix: align per-head h2o retention with independent phase configuration --- docs/en/features/sparse-methods.md | 9 ++-- docs/zh/features/sparse-methods.md | 3 +- src/sparsevllm/method_registry.py | 4 +- tests/test_h2o_cache_manager.py | 58 ------------------------ tests/test_h2o_chunk_numerics.py | 4 +- tests/test_h2o_model_integration.py | 2 +- tests/test_h2o_per_head_prefill.py | 34 +++++++++++++- tests/test_prefill_attention_provider.py | 2 +- 8 files changed, 46 insertions(+), 70 deletions(-) diff --git a/docs/en/features/sparse-methods.md b/docs/en/features/sparse-methods.md index 5197e348..45b5a055 100644 --- a/docs/en/features/sparse-methods.md +++ b/docs/en/features/sparse-methods.md @@ -50,8 +50,8 @@ SnapKV defaults `sparse_prefill_score_mode` to `logits`; `probability` remains an explicit reproducibility option because its additional normalized QK sweep is substantially more expensive in measured long-context prefill. PyramidKV and H2O use `probability`. H2O accumulates FP32 probability sums per query -head across prefill chunks. At eviction, max combines the cumulative scores within each GQA KV group, or across the -whole layer for MLA. MHA heads select independently. Each selection retains +head across prefill chunks. At eviction, max combines the cumulative scores +within each GQA KV group, or across the whole layer for MLA. MHA heads select independently. Each selection retains heavy hitters plus recent tokens within its token budget; native GQA KV sharing and MLA latent storage are preserved. MLA H2O prefill currently requires TP1. @@ -61,8 +61,9 @@ reduced logit vector cannot represent per-head cumulative probabilities. When attention provides its softmax LSE, scoring reuses it; otherwise it recomputes normalization from the same visible keys. Intermediate chunk eviction changes subsequent attention, so results can depend on chunk size and budgets. -Decode scoring and eviction remain disabled; `h2o_decode_budget` determines the -final prefill retention budget, and the cache grows during generation. +Decode scoring and eviction remain disabled. With `sparse_method=h2o`, +`h2o_decode_budget` determines final-prefill retention; prefill-only H2O skips +this final compaction. The cache then grows during generation. ## Prefill Scheduling Policies diff --git a/docs/zh/features/sparse-methods.md b/docs/zh/features/sparse-methods.md index d346eb23..76874b3f 100644 --- a/docs/zh/features/sparse-methods.md +++ b/docs/zh/features/sparse-methods.md @@ -48,7 +48,8 @@ prefill 中开销明显更高。PyramidKV 和 H2O 使用 `probability`。H2O 逐 H2O 拒绝 `logits` 模式,因为归约后的 logits 无法表示逐 head 累计概率。 Attention 提供 softmax LSE 时复用该结果,否则使用同一可见 KV 集合重新计算归一化。 中间 chunk 的实际驱逐会改变后续 attention,结果因此可能随 chunk size 和预算变化。 -Decode 评分和驱逐仍保持关闭;`h2o_decode_budget` 用于最后一个 prefill chunk 的保留预算, +Decode 评分和驱逐仍保持关闭;仅当 `sparse_method=h2o` 时, +`h2o_decode_budget` 用于最后一个 prefill chunk 的保留预算。仅启用 H2O prefill 时不做这次最终压缩; 之后缓存随生成增长。 ## Prefill Scheduling Policy diff --git a/src/sparsevllm/method_registry.py b/src/sparsevllm/method_registry.py index d1ef39b3..03029447 100644 --- a/src/sparsevllm/method_registry.py +++ b/src/sparsevllm/method_registry.py @@ -239,7 +239,9 @@ def sparse_prefill_attention_contract( prefill_sparse_method, sparse_method=normalized, ) - cache_method = resolve_cache_sparse_method(normalized, prefill_sparse_method=resolved_prefill_method) + cache_method = resolve_cache_sparse_method( + normalized, prefill_sparse_method=resolved_prefill_method, + ) layer_varying_page_table = _PREFILL_LAYER_VARYING_PAGE_TABLE[cache_method] if cache_method == "h2o": if resolve_sparse_prefill_score_mode(normalized, sparse_prefill_score_mode) != "probability": diff --git a/tests/test_h2o_cache_manager.py b/tests/test_h2o_cache_manager.py index d20e49d5..ea5a065c 100644 --- a/tests/test_h2o_cache_manager.py +++ b/tests/test_h2o_cache_manager.py @@ -264,64 +264,6 @@ def test_h2o_decode_does_not_request_scores_or_run_eviction(): runtime.cache_manager.evict_after_decode.assert_not_called() -def test_h2o_flashprefill_uses_posthoc_scoring_and_preserves_eviction(): - runtime = object.__new__(H2ORuntime) - runtime.config = SimpleNamespace( - sparse_method="h2o", - prefill_sparse_method="flashprefill_v2", - sparse_prefill_score_mode="logits", - h2o_prefill_score_window=0, - ) - runtime.cache_manager = Mock() - prefill = SparseStepContext( - seqs=[], - is_prefill=True, - forward_context=SimpleNamespace(is_prefill=True, is_long_text=True), - ) - - assert runtime.needs_attention_score(0, prefill) is False - runtime.finish_step(prefill) - - runtime.cache_manager.evict_after_intermediate_prefill.assert_not_called() - runtime.cache_manager.compact_final_prefill_for_decode.assert_called_once_with([]) - - -@pytest.mark.parametrize( - ("sparse_method", "prefill_sparse_method", "intermediate", "final"), - [ - ("", "h2o_prefill", True, False), - ("h2o", "", False, True), - ("h2o", "h2o_prefill", True, True), - ("", "", False, False), - ], -) -def test_h2o_runtime_triggers_prefill_and_decode_boundaries_independently( - sparse_method, - prefill_sparse_method, - intermediate, - final, -): - runtime = object.__new__(H2ORuntime) - runtime.config = SimpleNamespace( - sparse_method=sparse_method, - prefill_sparse_method=prefill_sparse_method, - ) - runtime.cache_manager = Mock() - step = SparseStepContext( - seqs=[], - is_prefill=True, - forward_context=SimpleNamespace(is_prefill=True, is_long_text=True), - ) - - runtime.finish_step(step) - - assert ( - runtime.cache_manager.evict_after_intermediate_prefill.called - is intermediate - ) - assert runtime.cache_manager.compact_final_prefill_for_decode.called is final - - def test_h2o_cache_manager_factory_routes_first_class_method(): expected = object() config = SimpleNamespace( diff --git a/tests/test_h2o_chunk_numerics.py b/tests/test_h2o_chunk_numerics.py index ee0a0cb8..c71dd047 100644 --- a/tests/test_h2o_chunk_numerics.py +++ b/tests/test_h2o_chunk_numerics.py @@ -39,7 +39,7 @@ def _manager(heads, kv_heads, length, budget): manager._num_free_slots = [length + 1] manager._h2o_scores, manager._h2o_positions = {}, {} manager._h2o_counters = dict(final_prefill_evictions=0, intermediate_prefill_evictions=0, dropped_tokens=0) - manager.config = SimpleNamespace(h2o_prefill_budget=budget, h2o_decode_budget=budget, + manager.config = SimpleNamespace(sparse_method='h2o', prefill_sparse_method='h2o_prefill', h2o_prefill_budget=budget, h2o_decode_budget=budget, h2o_prefill_score_window=0, h2o_recent_ratio=.25) runtime = object.__new__(H2ORuntime) runtime.config, runtime.cache_manager = manager.config, manager @@ -183,7 +183,7 @@ def test_mla_chunk_scores_and_outputs_follow_shared_latent_retention(budget): q = torch.randn(length, heads, 256, device='cuda', dtype=torch.bfloat16) def run(chunks): - manager, runtime = _manager(heads, 1, length, budget, 'max') + manager, runtime = _manager(heads, 1, length, budget) storage = MlaLatentStorage(kv_lora_rank=512, rope_dim=64, dtype=torch.bfloat16) storage.allocate(num_layers=1, num_slots=length + 1, device=manager.device) manager.attention_cache_storage = storage diff --git a/tests/test_h2o_model_integration.py b/tests/test_h2o_model_integration.py index a63f918f..3c56c9b3 100644 --- a/tests/test_h2o_model_integration.py +++ b/tests/test_h2o_model_integration.py @@ -22,7 +22,7 @@ def test_model_chunk_prefill_and_native_decode_handoff(tmp_path): model = os.environ['SPARSEVLLM_H2O_TEST_MODEL'] output_dir = Path(os.environ.get('SPARSEVLLM_H2O_TEST_OUTPUT', str(tmp_path))) output_dir.mkdir(parents=True, exist_ok=True) - config = dict(sparse_method='h2o', + config = dict(sparse_method='h2o', prefill_sparse_method='h2o_prefill', h2o_prefill_budget=64, h2o_decode_budget=32, h2o_prefill_score_window=0, engine_prefill_chunk_size=64, max_num_batched_tokens=128, max_num_seqs_in_batch=2, diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py index 115f33a9..7bc0efb3 100644 --- a/tests/test_h2o_per_head_prefill.py +++ b/tests/test_h2o_per_head_prefill.py @@ -137,7 +137,7 @@ def test_chunk_prefill_append_retention_and_final_handoff_follow_logical_positio manager.row_seq_lens[0][0] = 0 manager._h2o_scores.clear() manager._h2o_positions.clear() - manager.config = SimpleNamespace(h2o_prefill_budget=4, h2o_decode_budget=3, + manager.config = SimpleNamespace(sparse_method='h2o', prefill_sparse_method='h2o_prefill', h2o_prefill_budget=4, h2o_decode_budget=3, h2o_prefill_score_window=0, h2o_recent_ratio=.5) runtime = object.__new__(H2ORuntime) runtime.config = manager.config @@ -200,7 +200,7 @@ def _multi_request_manager(): manager.attention_cache_storage.allocate(num_layers=2, num_slots=24, device=manager.device) cache = manager.attention_cache_storage.cache cache.copy_(torch.arange(cache.numel()).reshape_as(cache)) - manager.config = SimpleNamespace(h2o_prefill_budget=4, h2o_decode_budget=3, + manager.config = SimpleNamespace(sparse_method='h2o', prefill_sparse_method='h2o_prefill', h2o_prefill_budget=4, h2o_decode_budget=3, h2o_recent_ratio=.5) manager.seq_id_to_row = [{7: 0, 8: 1}, {7: 0, 8: 1}] manager.row_seq_lens = [np.array([6, 5], dtype=np.int32) for _ in range(2)] @@ -309,3 +309,33 @@ def test_rejected_first_chunk_preserves_resident_head_history(): assert manager._h2o_positions[0, 7] is positions assert manager.row_seq_lens[0][0] == 6 assert manager._num_free_slots == [6] + + +@pytest.mark.parametrize( + 'sparse_method,prefill_method,expected_lengths', + [ + ('', 'h2o_prefill', [6, 4]), + ('h2o', '', [3, 5]), + ('h2o', 'h2o_prefill', [3, 4]), + ('h2o', 'flashprefill_v2', [3, 5]), + ('', '', [6, 5]), + ], +) +def test_independent_phase_switches_apply_per_head_retention( + sparse_method, prefill_method, expected_lengths, +): + from sparsevllm.engine.sparse_methods.base import SparseStepContext + from sparsevllm.engine.sparse_methods.h2o import H2ORuntime + + manager = _multi_request_manager() + manager.config.sparse_method = sparse_method + manager.config.prefill_sparse_method = prefill_method + runtime = object.__new__(H2ORuntime) + runtime.config, runtime.cache_manager = manager.config, manager + step = SparseStepContext(_mixed_prefill_seqs(), True, None) + assert not runtime.needs_attention_score(0, step) + runtime.finish_step(step) + for layer in range(manager.num_layers): + assert manager.row_seq_lens[layer].tolist() == expected_lengths + for row, seq in enumerate(step.seqs): + assert manager._h2o_scores[layer, seq.seq_id].shape[-1] == expected_lengths[row] diff --git a/tests/test_prefill_attention_provider.py b/tests/test_prefill_attention_provider.py index 2b0a0b00..459ca418 100644 --- a/tests/test_prefill_attention_provider.py +++ b/tests/test_prefill_attention_provider.py @@ -251,7 +251,7 @@ def test_independent_h2o_axes_keep_prompt_score_and_page_table_contract( assert contract.main_score_kind is AttentionScoreKind.NONE assert ( contract.score_collection - is PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_REDUCED + is PrefillScoreCollectionKind.METHOD_OWNED_POSTHOC_PER_HEAD )