[HiSparse] Clarify decode token usage logs (#22331)

This commit is contained in:
huangtingwei
2026-04-13 18:03:25 -07:00
committed by GitHub
parent c456cba7fd
commit 945d73824f
2 changed files with 54 additions and 0 deletions
@@ -27,6 +27,13 @@ class HiSparseAct(NamedTuple):
req: Req
class HiSparseTokenStats(NamedTuple):
device_tokens: int
device_token_usage: float
host_tokens: int
host_token_usage: float
class HiSparseCoordinator:
def __init__(
self,
@@ -125,6 +132,23 @@ class HiSparseCoordinator:
def set_decode_producer_stream(self, stream) -> None:
self.decode_producer_stream = stream
def get_token_stats(self) -> HiSparseTokenStats:
device_allocator = self.token_to_kv_pool_allocator.hisparse_attn_allocator
device_capacity = device_allocator.size
device_tokens = device_capacity - device_allocator.available_size()
host_capacity = self.mem_pool_host.size
host_tokens = host_capacity - self.mem_pool_host.available_size()
return HiSparseTokenStats(
device_tokens=device_tokens,
device_token_usage=(
device_tokens / device_capacity if device_capacity > 0 else 0.0
),
host_tokens=host_tokens,
host_token_usage=(
host_tokens / host_capacity if host_capacity > 0 else 0.0
),
)
def admit_request_into_staging(self, req: Req) -> None:
req.hisparse_staging = True
logical_indices = self.req_to_token_pool.req_to_token[
@@ -31,6 +31,7 @@ class PoolStats:
is_hybrid_swa: bool = False
is_hybrid_ssm: bool = False
is_hisparse: bool = False
# For hybrid-swa pools
swa_num_used: Optional[int] = None
@@ -44,6 +45,12 @@ class PoolStats:
mamba_available_size: Optional[int] = None
mamba_evictable_size: Optional[int] = None
# HiSparse device/host breakdown for decode logs (plain KV pool only)
hisparse_device_tokens: Optional[int] = None
hisparse_device_token_usage: Optional[float] = None
hisparse_host_tokens: Optional[int] = None
hisparse_host_token_usage: Optional[float] = None
def get_kv_token_stats(self) -> Tuple[int, float]:
# NOTE: mamba pool is not included in the "token usage" calculation.
if self.is_hybrid_swa:
@@ -98,6 +105,13 @@ class PoolStats:
f"mamba num: {self.mamba_num_used}",
f"mamba usage: {self.mamba_usage:.2f}",
]
if self.is_hisparse:
parts += [
f"#gpu token: {self.hisparse_device_tokens}",
f"gpu token usage: {self.hisparse_device_token_usage:.2f}",
f"#cpu token: {self.hisparse_host_tokens}",
f"cpu token usage: {self.hisparse_host_token_usage:.2f}",
]
if not parts:
parts.append(
f"#token: {self.full_num_used}, token usage: {self.full_token_usage:.2f}"
@@ -149,6 +163,8 @@ class SchedulerRuntimeCheckerMixin:
pool_stats = self._get_swa_token_info()
elif self.is_hybrid_ssm:
return self._get_mamba_token_info()
elif self.enable_hisparse:
return self._get_hisparse_token_info()
else:
return self._get_token_info()
@@ -175,6 +191,20 @@ class SchedulerRuntimeCheckerMixin:
full_evictable_size=evictable_size,
)
def _get_hisparse_token_info(self: Scheduler) -> PoolStats:
pool_stats = self._get_token_info()
if self.enable_hisparse and self.hisparse_coordinator is not None:
h = self.hisparse_coordinator.get_token_stats()
return dataclasses.replace(
pool_stats,
is_hisparse=True,
hisparse_device_tokens=h.device_tokens,
hisparse_device_token_usage=h.device_token_usage,
hisparse_host_tokens=h.host_tokens,
hisparse_host_token_usage=h.host_token_usage,
)
return pool_stats
def _get_mamba_token_info(self: Scheduler):
is_mamba_radix_cache = (
self.tree_cache.supports_mamba() and self.tree_cache.is_tree_cache()