Add periodic KV-canary stats logging and kernel-run-counter health check (#26821)

This commit is contained in:
fzyzcjy
2026-05-31 10:00:19 +08:00
committed by GitHub
parent 7dd19ae3d8
commit f220c72929
9 changed files with 304 additions and 0 deletions
+1
View File
@@ -753,6 +753,7 @@ class Envs:
# KV-Canary / Token-Oracle (testing-only)
# ===================================================================
SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024)
SGLANG_KV_CANARY_STATS_PRINT_EVERY_N_STEPS = EnvInt(100)
SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT = EnvBool(False)
SGLANG_KV_CANARY_PERTURB_REQ_TO_TOKEN_PROB = EnvFloat(0.0)
SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS = EnvInt(50)
+4
View File
@@ -47,6 +47,8 @@ class CanaryConfig:
expected_tokens from each req's ``origin_input_ids + output_ids`` (snapshotted at
ForwardBatch.init_new) and compare against the canary's stored tokens at verify time.
Independent of ``enable_write_input_assert``.
stats_print_every_n_steps: 0 disables periodic stats logging; positive N prints
"canary protected N tokens, ran M sweep passes, K violations so far" every N forward steps.
"""
mode: CanaryMode
@@ -55,6 +57,7 @@ class CanaryConfig:
real_kv_hash_mode: RealKvHashMode
enable_write_input_assert: bool
enable_verify_token_assert: bool
stats_print_every_n_steps: int
@classmethod
def from_env(cls, server_args: "ServerArgs") -> "CanaryConfig":
@@ -73,4 +76,5 @@ class CanaryConfig:
real_kv_hash_mode=RealKvHashMode[real_kv_raw],
enable_write_input_assert=envs.SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT.get(),
enable_verify_token_assert=envs.SGLANG_KV_CANARY_ENABLE_VERIFY_TOKEN_ASSERT.get(),
stats_print_every_n_steps=envs.SGLANG_KV_CANARY_STATS_PRINT_EVERY_N_STEPS.get(),
)
@@ -18,6 +18,8 @@ from sglang.srt.kv_canary.endpoint import (
)
from sglang.srt.kv_canary.perturb.config import PerturbConfig
from sglang.srt.kv_canary.perturb.manager import PerturbManager
from sglang.srt.kv_canary.runner.health_checker import KernelRunCounterHealthChecker
from sglang.srt.kv_canary.runner.stats_logger import PeriodicCanaryStatsLogger
from sglang.srt.kv_canary.runner.swa_divergence import SwaDivergenceReporter
from sglang.srt.kv_canary.runner.sweep import SweepOrchestrator
from sglang.srt.kv_canary.runner.violation_manager import ViolationManager
@@ -127,6 +129,22 @@ class CanaryManager:
swa_window_size=self._swa_window_size,
sweep_interval=config.sweep_interval,
)
self._health_checker = KernelRunCounterHealthChecker(
config=config,
device_state=self._device_state,
active_tags=self._active_tags,
outer_step_counter_getter=self._get_outer_step_counter,
d2h_stream=self._d2h_stream,
)
self._stats_logger = PeriodicCanaryStatsLogger(
config=config,
device_state=self._device_state,
active_tags=self._active_tags,
outer_step_counter_getter=self._get_outer_step_counter,
sweep_orchestrator=self._sweep_orchestrator,
d2h_stream=self._d2h_stream,
)
num_sfms = max(1, speculative_num_steps - 1)
self._single_forward_managers: tuple[SingleForwardManager, ...] = tuple(
SingleForwardManager(
@@ -233,6 +251,8 @@ class CanaryManager:
self._sweep_orchestrator.maybe_run_sweep()
self._outer_step_counter += 1
self._violation_manager.step()
self._health_checker.step()
self._stats_logger.step()
if self._swa_divergence_report is not None:
self._swa_divergence_report.step(
outer_step_counter=self._outer_step_counter,
@@ -0,0 +1,81 @@
from __future__ import annotations
import logging
from collections.abc import Callable
from typing import Optional
import torch
from sglang.jit_kernel.kv_canary.verify import CanaryLaunchTag
from sglang.srt.kv_canary.config import CanaryConfig
from sglang.srt.kv_canary.runner.future_tensor import DelayedDeviceHostHandler
from sglang.srt.kv_canary.runner.kernel_launcher import passes_v_half_gate
from sglang.srt.kv_canary.state import CanaryDeviceState
logger = logging.getLogger(__name__)
_HEALTH_CHECK_EVERY_N_STEPS: int = 100
_HEALTH_CHECK_WARMUP_STEPS: int = 100
_SWEEP_TAGS: frozenset[CanaryLaunchTag] = frozenset(
(
CanaryLaunchTag.SWEEP_K_FULL,
CanaryLaunchTag.SWEEP_V_FULL,
CanaryLaunchTag.SWEEP_K_SWA,
CanaryLaunchTag.SWEEP_V_SWA,
)
)
class KernelRunCounterHealthChecker:
def __init__(
self,
*,
config: CanaryConfig,
device_state: CanaryDeviceState,
active_tags: tuple[CanaryLaunchTag, ...],
outer_step_counter_getter: Callable[[], int],
d2h_stream: torch.cuda.Stream,
) -> None:
self._config = config
self._device_state = device_state
self._active_tags = active_tags
self._outer_step_counter_getter = outer_step_counter_getter
self._handler = DelayedDeviceHostHandler(d2h_stream=d2h_stream)
self._prev_counters_host: torch.Tensor = torch.zeros_like(
device_state.kernel_run_counters, device="cpu"
)
def step(self) -> None:
self._handler.step(
compute_on_device=self._compute_on_device,
postprocess_on_host=self._postprocess_on_host,
)
def _compute_on_device(self) -> Optional[torch.Tensor]:
outer_step_counter = self._outer_step_counter_getter()
if outer_step_counter < _HEALTH_CHECK_WARMUP_STEPS:
return None
if outer_step_counter % _HEALTH_CHECK_EVERY_N_STEPS != 0:
return None
if not self._active_tags:
return None
return self._device_state.kernel_run_counters
def _postprocess_on_host(self, new_counter_host: torch.Tensor) -> None:
delta = new_counter_host - self._prev_counters_host
self._prev_counters_host = new_counter_host
expected_tags = self._expected_active_tags_for_health_check()
stalled = [tag for tag in expected_tags if int(delta[tag.value]) == 0]
if stalled:
names = ", ".join(tag.name for tag in stalled)
raise RuntimeError(
f"kv-canary: kernel_run_counter did not increase since previous check "
f"for tags=[{names}] at step={self._outer_step_counter_getter()}; "
f"canary path is not executing"
)
def _expected_active_tags_for_health_check(self) -> tuple[CanaryLaunchTag, ...]:
tags = self._active_tags
if self._config.sweep_interval <= 0:
tags = tuple(tag for tag in tags if tag not in _SWEEP_TAGS)
return tuple(tag for tag in tags if passes_v_half_gate(tag))
@@ -0,0 +1,66 @@
from __future__ import annotations
import logging
from collections.abc import Callable
from typing import Any, Optional
import torch
from sglang.jit_kernel.kv_canary.verify import CanaryLaunchTag
from sglang.srt.kv_canary.config import CanaryConfig
from sglang.srt.kv_canary.runner.future_tensor import DelayedDeviceHostHandler
from sglang.srt.kv_canary.runner.sweep import SweepOrchestrator
from sglang.srt.kv_canary.state import CanaryDeviceState
logger = logging.getLogger(__name__)
class PeriodicCanaryStatsLogger:
def __init__(
self,
*,
config: CanaryConfig,
device_state: CanaryDeviceState,
active_tags: tuple[CanaryLaunchTag, ...],
outer_step_counter_getter: Callable[[], int],
sweep_orchestrator: SweepOrchestrator,
d2h_stream: torch.cuda.Stream,
) -> None:
self._config = config
self._device_state = device_state
self._active_tags = active_tags
self._outer_step_counter_getter = outer_step_counter_getter
self._sweep_orchestrator = sweep_orchestrator
self._handler = DelayedDeviceHostHandler(d2h_stream=d2h_stream)
def step(self) -> None:
self._handler.step(
compute_on_device=self._compute_on_device,
postprocess_on_host=self._postprocess_on_host,
)
def _compute_on_device(self) -> Optional[dict[str, Any]]:
period = self._config.stats_print_every_n_steps
if period <= 0:
return None
outer_step_counter = self._outer_step_counter_getter()
if outer_step_counter == 0 or outer_step_counter % period != 0:
return None
device_state = self._device_state
return {
"step": outer_step_counter,
"slot_sum": device_state.slot_run_counters.sum().view(1),
"write_index": device_state.violation_log.violation_write_index,
}
def _postprocess_on_host(self, host_data: dict[str, Any]) -> None:
logger.info(
"[canary] step=%d protected_tokens=%d sweep_passes=%d violations=%d "
"launch_tags_active=%d/%d",
int(host_data["step"]),
int(host_data["slot_sum"].item()),
self._sweep_orchestrator.sweep_passes,
int(host_data["write_index"].item()),
len(self._active_tags),
len(CanaryLaunchTag),
)
+1
View File
@@ -140,6 +140,7 @@ def make_base_config() -> CanaryConfig:
real_kv_hash_mode=consts.RealKvHashMode.NONE,
enable_write_input_assert=False,
enable_verify_token_assert=True,
stats_print_every_n_steps=100,
)
@@ -30,6 +30,7 @@ def make_config(
real_kv_hash_mode: RealKvHashMode = RealKvHashMode.NONE,
enable_write_input_assert: bool = False,
enable_verify_token_assert: bool = True,
stats_print_every_n_steps: int = 100,
) -> CanaryConfig:
return CanaryConfig(
mode=mode,
@@ -38,6 +39,7 @@ def make_config(
real_kv_hash_mode=real_kv_hash_mode,
enable_write_input_assert=enable_write_input_assert,
enable_verify_token_assert=enable_verify_token_assert,
stats_print_every_n_steps=stats_print_every_n_steps,
)