File size: 1,992 Bytes
31dc8dc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 | import torch
from typing import Callable
from dataclasses import dataclass, field
@dataclass
class AttnMetaDataBase:
is_prefill: list[bool] = field(default_factory=lambda: [False])
enforce_eager: bool = False
cu_seqlens_q: torch.Tensor | None = None
cu_seqlens_k: torch.Tensor | None = None
max_seqlen_q: int = 0
max_seqlen_k: int = 0
slot_mapping: torch.Tensor | None = None
context_lens: torch.Tensor | None = None
page_tables: torch.Tensor | None = None
status_table: torch.Tensor | None = None
page_size: int = 32
block_size: int = 32
kv_cache_layout: str = "unified"
token_merge_enabled: bool = False
token_merge_mask: torch.Tensor | None = None
token_merge_topk_ids: torch.Tensor | None = None
token_merge_topk_probs: torch.Tensor | None = None
token_merge_residual_probs: torch.Tensor | None = None
token_merge_mask_token_id: int | None = None
token_merge_renormalize: bool = True
token_merge_mode: str = "dmax_topk"
token_merge_weight: float = 1.0
@property
def num_reqs(self) -> int:
return len(self.cu_seqlens_q) - 1
@property
def chunk_size(self) -> int:
return self.block_size * self.buffer_size
@property
def need_kv_cache_store(self) -> bool:
if is_warming_up() and self.slot_mapping.numel() > 0:
return True
return (self.slot_mapping >= 0).any() if self.enforce_eager else self.slot_mapping.numel() > 0
FN_TYPE_AttnMetaDataFetch = Callable[[], AttnMetaDataBase]
fetch_attn_metadata: FN_TYPE_AttnMetaDataFetch = ...
def set_fetch_fn_for_attn_metadata(fn: FN_TYPE_AttnMetaDataFetch) -> None:
global fetch_attn_metadata
fetch_attn_metadata = fn
WARMING_UP = False
def set_warming_up(is_warming_up: bool) -> None:
global WARMING_UP
WARMING_UP = is_warming_up
def is_warming_up() -> bool:
return WARMING_UP
def reset_warming_up() -> None:
global WARMING_UP
WARMING_UP = False
|