diff --git a/docs/en/features/sparse-methods.md b/docs/en/features/sparse-methods.md index d203de4f..45b5a055 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, 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. + +`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. 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 f06fef69..76874b3f 100644 --- a/docs/zh/features/sparse-methods.md +++ b/docs/zh/features/sparse-methods.md @@ -38,15 +38,19 @@ 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 概率和;驱逐时通过 max, +在 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 评分和驱逐仍保持关闭;仅当 `sparse_method=h2o` 时, +`h2o_decode_budget` 用于最后一个 prefill chunk 的保留预算。仅启用 H2O prefill 时不做这次最终压缩; +之后缓存随生成增长。 ## Prefill Scheduling Policy diff --git a/src/sparsevllm/configs/sparse.py b/src/sparsevllm/configs/sparse.py index d4990f95..39ea0c91 100644 --- a/src/sparsevllm/configs/sparse.py +++ b/src/sparsevllm/configs/sparse.py @@ -148,6 +148,8 @@ def _normalize_snapkv(config) -> None: def _normalize_h2o(config) -> None: + 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 +164,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..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 @@ -15,9 +16,18 @@ from sparsevllm.utils.context import get_context from sparsevllm.utils.profiler import profiler -from .base import ExplicitKVPayload, PrefillComputeView from .snapkv import SnapKVCacheManager -from .storage import ExplicitKVStorage + + +@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): @@ -27,15 +37,7 @@ class _H2ORowRef(NamedTuple): 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. - """ + """Native KV storage with cumulative per-query-head prefill probabilities.""" def __init__( self, @@ -50,6 +52,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. @@ -64,7 +67,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 @@ -97,6 +99,132 @@ 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 + + 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.") + + 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.", + ) + 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() + 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 +712,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 +728,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,44 +746,16 @@ 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 - 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: @@ -703,20 +796,40 @@ 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) 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) + 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 +903,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 +915,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 @@ -1203,473 +1177,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], @@ -1814,6 +1321,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 +1351,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 @@ -1854,22 +1375,11 @@ 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.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": ( - None - if workspace is None - else { - "shape": list(workspace.shape), - "dtype": str(workspace.dtype), - "device": str(workspace.device), - "nbytes": int(workspace.untyped_storage().nbytes()), - } - ), } return summary 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/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..b5f3051e 100644 --- a/src/sparsevllm/engine/sparse_methods/h2o.py +++ b/src/sparsevllm/engine/sparse_methods/h2o.py @@ -3,20 +3,61 @@ import torch 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 import H2ORetention 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 +H2O_PREFILL_QUERY_TILE = 128 + + +def select_h2o_heads( + cumulative: torch.Tensor, + *, + selection_groups: int, + budget: int, + recent_ratio: float, +) -> 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. 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: + 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 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) + 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): 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 +79,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 +88,142 @@ 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), + ) + 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/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..03029447 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() @@ -239,25 +240,15 @@ def sparse_prefill_attention_contract( sparse_method=normalized, ) cache_method = resolve_cache_sparse_method( - normalized, - prefill_sparse_method=resolved_prefill_method, + normalized, prefill_sparse_method=resolved_prefill_method, ) - h2o_score_collection = cache_method == "h2o" 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 +263,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/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, diff --git a/tests/test_glm_mla_prefix_cache.py b/tests/test_glm_mla_prefix_cache.py index 92551639..fc203c9d 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 @@ -105,6 +107,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 +116,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 +261,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 +282,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 +296,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 +324,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 +333,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 +349,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 +361,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 e13a226a..ea5a065c 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, @@ -28,6 +27,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 +94,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, @@ -101,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)} @@ -185,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(): @@ -283,70 +258,12 @@ 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() -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( @@ -451,9 +368,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) @@ -462,49 +378,20 @@ 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 + runtime.config = manager.config + runtime._prefill_score_workspace = PrefillScoreWorkspace() + runtime._prefill_head_score_buffer = None + return runtime 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]) + 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], @@ -512,131 +399,24 @@ def test_h2o_prefill_score_collection_accumulates_in_physical_coordinates(): 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)), - ), + 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]]) - 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 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) - -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 test_h2o_logit_prefill_score_collection_consumes_fused_main_score(): - 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) - 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)), - ), - ) - set_context(is_prefill=True, cache_manager=manager, seqs=[seq]) - - 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), - ) - - 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,45 +438,20 @@ 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) -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) @@ -962,221 +717,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 @@ -1250,42 +790,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 @@ -2290,3 +1794,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_chunk_numerics.py b/tests/test_h2o_chunk_numerics.py new file mode 100644 index 00000000..c71dd047 --- /dev/null +++ b/tests/test_h2o_chunk_numerics.py @@ -0,0 +1,258 @@ +"""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, DecodeComputeView, 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.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') + + +def _manager(heads, kv_heads, length, budget): + 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(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 + 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): + length, heads, dim = q.shape + kv_heads = k.shape[1] + group_size = heads // kv_heads + manager, runtime = _manager(heads, kv_heads, length, budget) + 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 + decode = None + 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) + 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 + # Exercise the registered native decode provider on the compressed KV. + resident = len(histories[0]) + context = torch.tensor([resident], device=q.device, dtype=torch.int32) + 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() + 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: + if decode is not None: + decode.close() + reset_context() + + +@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) + + +@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) + for chunks in ([13, 129, 7], [71, 78]): + 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) + + +@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) + 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_model_integration.py b/tests/test_h2o_model_integration.py new file mode 100644 index 00000000..3c56c9b3 --- /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', 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, + 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() diff --git a/tests/test_h2o_per_head_prefill.py b/tests/test_h2o_per_head_prefill.py new file mode 100644 index 00000000..7bc0efb3 --- /dev/null +++ b/tests/test_h2o_per_head_prefill.py @@ -0,0 +1,341 @@ +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 import H2ORetention +from sparsevllm.engine.cache_manager.storage import ExplicitKVStorage, MlaLatentStorage +from sparsevllm.engine.sparse_methods.h2o 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) + assert keep.tolist() == [[1, 2]] + assert p.amax(1).sum(0)[0] > p.amax(1).sum(0)[1] + + +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, **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) + assert keep.tolist() == [[0, 1, 7, 8]] * 3 + all_tokens = select_h2o_heads(scores, selection_groups=3, budget=12, recent_ratio=.5) + 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(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 + 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 + + +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(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)] + 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]] + + +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() + + +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] + + +@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_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..459ca418 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( @@ -260,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 ) @@ -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( 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