From cdee16e14415b8f104bfb889d9f85d51d14e5a22 Mon Sep 17 00:00:00 2001 From: fzyzcjy <5236035+fzyzcjy@users.noreply.github.com> Date: Sun, 31 May 2026 09:58:14 +0800 Subject: [PATCH] Add the KV-canary perturb framework for fault-injection self-tests (#26816) --- python/sglang/srt/environ.py | 1 + python/sglang/srt/kv_canary/api.py | 6 +- python/sglang/srt/kv_canary/perturb/config.py | 67 +++++++++++ .../sglang/srt/kv_canary/perturb/manager.py | 43 +++++++ python/sglang/srt/kv_canary/perturb/utils.py | 107 ++++++++++++++++++ .../srt/kv_canary/runner/canary_manager.py | 14 +++ .../sglang/test/kv_canary/runner_test_base.py | 14 +++ .../test/kv_canary/violation_assert_mixin.py | 22 +++- .../kv_canary/test_self_unit_perturb.py | 59 ++++++++++ 9 files changed, 331 insertions(+), 2 deletions(-) create mode 100644 python/sglang/srt/kv_canary/perturb/config.py create mode 100644 python/sglang/srt/kv_canary/perturb/manager.py create mode 100644 python/sglang/srt/kv_canary/perturb/utils.py create mode 100644 test/registered/kv_canary/test_self_unit_perturb.py diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 547ed453c..97b2ddd4c 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -755,6 +755,7 @@ class Envs: SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024) SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT = EnvBool(False) SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS = EnvInt(50) + SGLANG_KV_CANARY_PERTURB_TARGET_GROUP = EnvStr(None) SGLANG_KV_CANARY_PERTURB_NEXT_TOKEN_SWAP_PROB = EnvFloat(0.0) SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE = EnvBool(False) SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False) diff --git a/python/sglang/srt/kv_canary/api.py b/python/sglang/srt/kv_canary/api.py index 2d1eaab90..95cc1d961 100644 --- a/python/sglang/srt/kv_canary/api.py +++ b/python/sglang/srt/kv_canary/api.py @@ -7,6 +7,7 @@ import torch from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities from sglang.srt.kv_canary.config import CanaryConfig, CanaryMode +from sglang.srt.kv_canary.perturb.config import PerturbConfig from sglang.srt.kv_canary.pool_patcher.api import attach_canary_buffers from sglang.srt.kv_canary.pool_patcher.utils import wrap_method from sglang.srt.kv_canary.runner.canary_manager import CanaryManager @@ -36,6 +37,7 @@ def install_canary( "when canary is enabled" ) + perturb_config = PerturbConfig.from_env() device = torch.device(model_runner.device) # EAGLE draft worker pools rotate input_ids so slot ``p`` stores K/V for the token at position ``p+1``; # target pools have no such shift. Threaded into the plan-side expected-token gather kernel. @@ -56,6 +58,7 @@ def install_canary( speculative_num_steps = int(server_args.speculative_num_steps or 1) manager = CanaryManager( config=config, + perturb_config=perturb_config, buffer_groups=buffer_groups, device=device, req_to_token_pool=model_runner.req_to_token_pool, @@ -71,11 +74,12 @@ def install_canary( # Single-line summary of every knob that controls canary behavior at boot time. # Disaggregation mode is included so PD logs are unambiguous about which side this is. logger.info( - "install_canary: disaggregation_mode=%s config=%s " + "install_canary: disaggregation_mode=%s config=%s perturb_config=%s " "launch_capacities=%s n_buffer_groups=%d buffer_group_kinds=%s " "swa_window_size=%d speculative_num_steps=%d", server_args.disaggregation_mode, config, + perturb_config, launch_capacities, len(buffer_groups), [g.kind.name for g in buffer_groups], diff --git a/python/sglang/srt/kv_canary/perturb/config.py b/python/sglang/srt/kv_canary/perturb/config.py new file mode 100644 index 000000000..289b1414a --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/config.py @@ -0,0 +1,67 @@ +from __future__ import annotations + +from dataclasses import dataclass +from enum import IntEnum + +from sglang.srt.environ import envs +from sglang.srt.kv_canary.buffer_group import PoolKind + + +class TargetGroupKind(IntEnum): + FULL = PoolKind.FULL.value + SWA = PoolKind.SWA.value + + def __str__(self) -> str: + return self.name.lower() + + +@dataclass(frozen=True, slots=True, kw_only=True) +class PerturbConfig: + target_group_kind: TargetGroupKind | None + warmup_steps: int + + @classmethod + def from_env(cls) -> "PerturbConfig": + return cls( + target_group_kind=_parse_target_group_kind_from_env( + raw=envs.SGLANG_KV_CANARY_PERTURB_TARGET_GROUP.get(), + ), + warmup_steps=envs.SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS.get(), + ) + + +def _parse_target_group_kind_from_env( + *, + raw: str | None, +) -> TargetGroupKind | None: + if raw is not None and raw.strip(): + return _parse_target_group_kind(raw) + return None + + +def require_target_group_kind( + *, target_group_kind: TargetGroupKind | None, perturb_name: str +) -> TargetGroupKind: + if target_group_kind is None: + raise ValueError( + "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP must be explicitly set to " + f"'full' or 'swa' when {perturb_name} perturbation is enabled" + ) + return target_group_kind + + +def _parse_target_group_kind(raw: str | None) -> TargetGroupKind: + if raw is None or not raw.strip(): + raise ValueError( + "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP must be explicitly set to " + "'full' or 'swa'" + ) + + value = raw.strip().lower() + try: + return TargetGroupKind[value.upper()] + except KeyError: + raise ValueError( + "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP must be one of 'full' / " + f"'swa', got {raw!r}" + ) from None diff --git a/python/sglang/srt/kv_canary/perturb/manager.py b/python/sglang/srt/kv_canary/perturb/manager.py new file mode 100644 index 000000000..81e71e848 --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/manager.py @@ -0,0 +1,43 @@ +from __future__ import annotations + +from collections.abc import Callable +from typing import TYPE_CHECKING, Optional + +from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup +from sglang.srt.kv_canary.perturb.config import PerturbConfig +from sglang.srt.kv_canary.perturb.utils import WarmupGate + +if TYPE_CHECKING: + from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + + +class PerturbManager: + def __init__( + self, + *, + config: PerturbConfig, + buffer_groups: tuple[CanaryBufferGroup, ...], + outer_step_counter_getter: Callable[[], int], + swa_window_size: int = 0, + sweep_interval: int = 0, + ) -> None: + self._config = config + self._buffer_groups = buffer_groups + self._outer_step_counter_getter = outer_step_counter_getter + self._swa_window_size = swa_window_size + self._sweep_interval = sweep_interval + self._radix_cache: Optional["BasePrefixCache"] = None + self._warmup_gate = WarmupGate( + config=config, outer_step_counter_getter=outer_step_counter_getter + ) + + def attach_radix_cache(self, radix_cache: "BasePrefixCache") -> None: + self._radix_cache = radix_cache + + def perturb( + self, + *, + maybe_inaccurate_forward_batch: Optional["ForwardBatch"], + ) -> None: + pass diff --git a/python/sglang/srt/kv_canary/perturb/utils.py b/python/sglang/srt/kv_canary/perturb/utils.py new file mode 100644 index 000000000..db9dc17d7 --- /dev/null +++ b/python/sglang/srt/kv_canary/perturb/utils.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import logging +import random +from collections.abc import Callable +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind +from sglang.srt.kv_canary.perturb.config import PerturbConfig, TargetGroupKind + +logger = logging.getLogger(__name__) + +if TYPE_CHECKING: + from sglang.srt.model_executor.forward_batch_info import ForwardBatch + + +class WarmupGate: + """Per-hook warmup window check + once-per-lifetime disable/enable log emission. + + Shared across the four perturb-point hooks so warmup state is decided in one place + rather than duplicated per hook. + """ + + def __init__( + self, + *, + config: PerturbConfig, + outer_step_counter_getter: Callable[[], int], + ) -> None: + self._config = config + self._outer_step_counter_getter = outer_step_counter_getter + self._warmup_disable_logged: bool = False + self._warmup_enable_logged: bool = False + + def is_in_warmup(self) -> bool: + step = self._outer_step_counter_getter() + warmup_steps = self._config.warmup_steps + + if step < warmup_steps: + self._log_warmup_disabled_once(warmup_steps) + return True + + self._log_warmup_enabled_once(step) + return False + + def _log_warmup_disabled_once(self, warmup_steps: int) -> None: + if self._warmup_disable_logged: + return + + logger.info( + "kv_canary perturb: disabled during warmup window " + "(first %d forward steps)", + warmup_steps, + ) + self._warmup_disable_logged = True + + def _log_warmup_enabled_once(self, step: int) -> None: + if self._warmup_enable_logged: + return + + logger.info("kv_canary perturb: enabled after warmup window at step=%d", step) + self._warmup_enable_logged = True + + +def should_run_perturbation( + *, + perturb_name: str, + probability: float, + warmup_gate: WarmupGate, + maybe_inaccurate_forward_batch: Optional["ForwardBatch"], + require_forward_batch: bool = True, +) -> bool: + if probability <= 0.0: + return False + if warmup_gate.is_in_warmup(): + return False + if require_forward_batch and maybe_inaccurate_forward_batch is None: + logger.info( + "kv_canary perturb %s: skipped because maybe_inaccurate_forward_batch is unavailable", + perturb_name, + ) + return False + return torch.rand((), device="cpu").item() < probability + + +def pick_target_group( + *, + buffer_groups: tuple[CanaryBufferGroup, ...], + target_kind: TargetGroupKind, +) -> Optional[CanaryBufferGroup]: + """Filter buffer_groups by target_kind. + + Returns None if no group matches. + """ + if target_kind == TargetGroupKind.FULL: + want = PoolKind.FULL + elif target_kind == TargetGroupKind.SWA: + want = PoolKind.SWA + else: + raise ValueError(f"Unsupported target_group_kind: {target_kind!r}") + filtered = [group for group in buffer_groups if group.kind == want] + if not filtered: + return None + pick = random.randrange(len(filtered)) + return filtered[pick] diff --git a/python/sglang/srt/kv_canary/runner/canary_manager.py b/python/sglang/srt/kv_canary/runner/canary_manager.py index bff3026be..8dfa809a7 100644 --- a/python/sglang/srt/kv_canary/runner/canary_manager.py +++ b/python/sglang/srt/kv_canary/runner/canary_manager.py @@ -16,6 +16,8 @@ from sglang.srt.kv_canary.endpoint import ( CanaryEndpoint, build_endpoints_from_group, ) +from sglang.srt.kv_canary.perturb.config import PerturbConfig +from sglang.srt.kv_canary.perturb.manager import PerturbManager from sglang.srt.kv_canary.runner.sweep import SweepOrchestrator from sglang.srt.kv_canary.runner.violation_manager import ViolationManager from sglang.srt.kv_canary.single_forward_manager.manager import ( @@ -38,6 +40,7 @@ class CanaryManager: self, *, config: CanaryConfig, + perturb_config: PerturbConfig, buffer_groups: tuple[CanaryBufferGroup, ...], device: torch.device, req_to_token_pool: "ReqToTokenPool", @@ -96,6 +99,13 @@ class CanaryManager: swa_window_size=self._swa_window_size, outer_step_counter_getter=self._get_outer_step_counter, ) + self._perturb_manager = PerturbManager( + config=perturb_config, + buffer_groups=self._buffer_groups, + outer_step_counter_getter=self._get_outer_step_counter, + swa_window_size=self._swa_window_size, + sweep_interval=config.sweep_interval, + ) num_sfms = max(1, speculative_num_steps - 1) self._single_forward_managers: tuple[SingleForwardManager, ...] = tuple( SingleForwardManager( @@ -183,6 +193,9 @@ class CanaryManager: self._single_forward_managers[idx].pre_ops_outside_graph( maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch ) + self._perturb_manager.perturb( + maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch + ) def _post_ops_outside_graph( self, @@ -203,6 +216,7 @@ class CanaryManager: def attach_radix_cache(self, radix_cache: "BasePrefixCache") -> None: self._sweep_orchestrator.attach_radix_cache(radix_cache) + self._perturb_manager.attach_radix_cache(radix_cache) def _get_outer_step_counter(self) -> int: return self._outer_step_counter diff --git a/python/sglang/test/kv_canary/runner_test_base.py b/python/sglang/test/kv_canary/runner_test_base.py index d2f3f48b8..0e0ac942c 100644 --- a/python/sglang/test/kv_canary/runner_test_base.py +++ b/python/sglang/test/kv_canary/runner_test_base.py @@ -10,6 +10,7 @@ from sglang.srt.kv_canary import endpoint as endpoint_module from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities from sglang.srt.kv_canary.config import CanaryConfig, CanaryMode +from sglang.srt.kv_canary.perturb.config import PerturbConfig from sglang.srt.kv_canary.runner import kernel_launcher as kernel_launcher_module from sglang.srt.kv_canary.runner.canary_manager import CanaryManager from sglang.test.kv_canary.fixtures import ( @@ -44,10 +45,20 @@ class RecordingEndpoint: self.calls.append(kwargs) +def make_perturb_config() -> PerturbConfig: + """Build a PerturbConfig with every probability pinned to 0 so the + perturb hooks do nothing during unit tests.""" + return PerturbConfig( + target_group_kind=None, + warmup_steps=0, + ) + + def make_manager( *, device: torch.device, config: CanaryConfig | None = None, + perturb_config: PerturbConfig | None = None, group: CanaryBufferGroup | None = None, req_pool: SimpleNamespace | None = None, per_forward_verify_capacity: int = 16, @@ -55,12 +66,15 @@ def make_manager( ) -> CanaryManager: if config is None: config = make_config() + if perturb_config is None: + perturb_config = make_perturb_config() if group is None: group = make_buffer_group(device=device) if req_pool is None: req_pool = make_req_to_token_pool(device=device, max_reqs=4, max_seq_len=8) return CanaryManager( config=config, + perturb_config=perturb_config, buffer_groups=(group,), device=device, req_to_token_pool=req_pool, diff --git a/python/sglang/test/kv_canary/violation_assert_mixin.py b/python/sglang/test/kv_canary/violation_assert_mixin.py index 808203acc..621cb25ef 100644 --- a/python/sglang/test/kv_canary/violation_assert_mixin.py +++ b/python/sglang/test/kv_canary/violation_assert_mixin.py @@ -3,6 +3,7 @@ from __future__ import annotations import time from typing import Literal, Optional +from sglang.srt.kv_canary.perturb.config import TargetGroupKind from sglang.test.kv_canary.violation_log_utils import ( assert_no_violation_in_log, find_violation_in_log, @@ -19,16 +20,35 @@ class CanaryViolationAssertMixin: self, *, fail_reason: str, + target_group: Optional[TargetGroupKind] = None, side: _Side = None, flush_wait_seconds: float = 2.0, ) -> None: + suffix = "" if target_group is None else f"_{target_group.name}" self.assert_violation_logged_any( - launch_tag_patterns=("HEAD_*", "TAIL_*"), + launch_tag_patterns=(f"HEAD_*{suffix}", f"TAIL_*{suffix}"), fail_reason=fail_reason, side=side, flush_wait_seconds=flush_wait_seconds, ) + def assert_sweep_violation_reported( + self, + *, + fail_reason: str, + target_group: TargetGroupKind, + side: _Side = None, + flush_wait_seconds: float = 2.0, + max_retries: int = 4, + ) -> None: + self.assert_violation_logged_any( + launch_tag_patterns=(f"SWEEP_*_{target_group.name}",), + fail_reason=fail_reason, + side=side, + flush_wait_seconds=flush_wait_seconds, + max_retries=max_retries, + ) + def assert_any_launch_tag_violation_reported( self, *, diff --git a/test/registered/kv_canary/test_self_unit_perturb.py b/test/registered/kv_canary/test_self_unit_perturb.py new file mode 100644 index 000000000..d808c981e --- /dev/null +++ b/test/registered/kv_canary/test_self_unit_perturb.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +import os +import unittest +from unittest.mock import patch + +from sglang.srt.kv_canary.perturb.config import ( + PerturbConfig, + TargetGroupKind, + _parse_target_group_kind, +) +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.test_utils import CustomTestCase + +register_cuda_ci(est_time=10, stage="extra-a", runner_config="1-gpu-small") + + +class TestParseTargetGroupKind(CustomTestCase): + def test_parse_target_group_kind_accepts_valid_values_case_insensitively( + self, + ) -> None: + """Verify target group kind parsing accepts valid names case-insensitively.""" + cases = [ + ("full", TargetGroupKind.FULL), + ("FULL", TargetGroupKind.FULL), + (" swa ", TargetGroupKind.SWA), + ] + + for raw, expected in cases: + with self.subTest(raw=raw): + self.assertEqual(_parse_target_group_kind(raw), expected) + + def test_parse_target_group_kind_rejects_invalid_value(self) -> None: + """Verify target group kind parsing rejects unknown names.""" + with self.assertRaisesRegex(ValueError, "must be one of"): + _parse_target_group_kind("prefix") + + def test_parse_target_group_kind_rejects_missing_or_any(self) -> None: + """Verify target group kind parsing requires an explicit concrete group.""" + for raw in [None, "", "any", " Any "]: + with self.subTest(raw=raw): + with self.assertRaisesRegex( + ValueError, "SGLANG_KV_CANARY_PERTURB_TARGET_GROUP" + ): + _parse_target_group_kind(raw) + + def test_from_env_allows_missing_target( + self, + ) -> None: + """Verify normal canary startup does not require a perturb target group.""" + with patch.dict(os.environ, {}, clear=False): + os.environ.pop("SGLANG_KV_CANARY_PERTURB_TARGET_GROUP", None) + config = PerturbConfig.from_env() + + self.assertIsNone(config.target_group_kind) + + +if __name__ == "__main__": + unittest.main()