Add the KV-canary perturb framework for fault-injection self-tests (#26816)
This commit is contained in:
@@ -755,6 +755,7 @@ class Envs:
|
|||||||
SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024)
|
SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024)
|
||||||
SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT = EnvBool(False)
|
SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT = EnvBool(False)
|
||||||
SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS = EnvInt(50)
|
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_PERTURB_NEXT_TOKEN_SWAP_PROB = EnvFloat(0.0)
|
||||||
SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE = EnvBool(False)
|
SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE = EnvBool(False)
|
||||||
SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False)
|
SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False)
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import torch
|
|||||||
|
|
||||||
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
|
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
|
||||||
from sglang.srt.kv_canary.config import CanaryConfig, CanaryMode
|
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.api import attach_canary_buffers
|
||||||
from sglang.srt.kv_canary.pool_patcher.utils import wrap_method
|
from sglang.srt.kv_canary.pool_patcher.utils import wrap_method
|
||||||
from sglang.srt.kv_canary.runner.canary_manager import CanaryManager
|
from sglang.srt.kv_canary.runner.canary_manager import CanaryManager
|
||||||
@@ -36,6 +37,7 @@ def install_canary(
|
|||||||
"when canary is enabled"
|
"when canary is enabled"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
perturb_config = PerturbConfig.from_env()
|
||||||
device = torch.device(model_runner.device)
|
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``;
|
# 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.
|
# 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)
|
speculative_num_steps = int(server_args.speculative_num_steps or 1)
|
||||||
manager = CanaryManager(
|
manager = CanaryManager(
|
||||||
config=config,
|
config=config,
|
||||||
|
perturb_config=perturb_config,
|
||||||
buffer_groups=buffer_groups,
|
buffer_groups=buffer_groups,
|
||||||
device=device,
|
device=device,
|
||||||
req_to_token_pool=model_runner.req_to_token_pool,
|
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.
|
# 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.
|
# Disaggregation mode is included so PD logs are unambiguous about which side this is.
|
||||||
logger.info(
|
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 "
|
"launch_capacities=%s n_buffer_groups=%d buffer_group_kinds=%s "
|
||||||
"swa_window_size=%d speculative_num_steps=%d",
|
"swa_window_size=%d speculative_num_steps=%d",
|
||||||
server_args.disaggregation_mode,
|
server_args.disaggregation_mode,
|
||||||
config,
|
config,
|
||||||
|
perturb_config,
|
||||||
launch_capacities,
|
launch_capacities,
|
||||||
len(buffer_groups),
|
len(buffer_groups),
|
||||||
[g.kind.name for g in buffer_groups],
|
[g.kind.name for g in buffer_groups],
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
@@ -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]
|
||||||
@@ -16,6 +16,8 @@ from sglang.srt.kv_canary.endpoint import (
|
|||||||
CanaryEndpoint,
|
CanaryEndpoint,
|
||||||
build_endpoints_from_group,
|
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.sweep import SweepOrchestrator
|
||||||
from sglang.srt.kv_canary.runner.violation_manager import ViolationManager
|
from sglang.srt.kv_canary.runner.violation_manager import ViolationManager
|
||||||
from sglang.srt.kv_canary.single_forward_manager.manager import (
|
from sglang.srt.kv_canary.single_forward_manager.manager import (
|
||||||
@@ -38,6 +40,7 @@ class CanaryManager:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
config: CanaryConfig,
|
config: CanaryConfig,
|
||||||
|
perturb_config: PerturbConfig,
|
||||||
buffer_groups: tuple[CanaryBufferGroup, ...],
|
buffer_groups: tuple[CanaryBufferGroup, ...],
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
req_to_token_pool: "ReqToTokenPool",
|
req_to_token_pool: "ReqToTokenPool",
|
||||||
@@ -96,6 +99,13 @@ class CanaryManager:
|
|||||||
swa_window_size=self._swa_window_size,
|
swa_window_size=self._swa_window_size,
|
||||||
outer_step_counter_getter=self._get_outer_step_counter,
|
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)
|
num_sfms = max(1, speculative_num_steps - 1)
|
||||||
self._single_forward_managers: tuple[SingleForwardManager, ...] = tuple(
|
self._single_forward_managers: tuple[SingleForwardManager, ...] = tuple(
|
||||||
SingleForwardManager(
|
SingleForwardManager(
|
||||||
@@ -183,6 +193,9 @@ class CanaryManager:
|
|||||||
self._single_forward_managers[idx].pre_ops_outside_graph(
|
self._single_forward_managers[idx].pre_ops_outside_graph(
|
||||||
maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch
|
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(
|
def _post_ops_outside_graph(
|
||||||
self,
|
self,
|
||||||
@@ -203,6 +216,7 @@ class CanaryManager:
|
|||||||
|
|
||||||
def attach_radix_cache(self, radix_cache: "BasePrefixCache") -> None:
|
def attach_radix_cache(self, radix_cache: "BasePrefixCache") -> None:
|
||||||
self._sweep_orchestrator.attach_radix_cache(radix_cache)
|
self._sweep_orchestrator.attach_radix_cache(radix_cache)
|
||||||
|
self._perturb_manager.attach_radix_cache(radix_cache)
|
||||||
|
|
||||||
def _get_outer_step_counter(self) -> int:
|
def _get_outer_step_counter(self) -> int:
|
||||||
return self._outer_step_counter
|
return self._outer_step_counter
|
||||||
|
|||||||
@@ -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.buffer_group import CanaryBufferGroup
|
||||||
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
|
from sglang.srt.kv_canary.capacities import CanaryLaunchCapacities
|
||||||
from sglang.srt.kv_canary.config import CanaryConfig, CanaryMode
|
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 import kernel_launcher as kernel_launcher_module
|
||||||
from sglang.srt.kv_canary.runner.canary_manager import CanaryManager
|
from sglang.srt.kv_canary.runner.canary_manager import CanaryManager
|
||||||
from sglang.test.kv_canary.fixtures import (
|
from sglang.test.kv_canary.fixtures import (
|
||||||
@@ -44,10 +45,20 @@ class RecordingEndpoint:
|
|||||||
self.calls.append(kwargs)
|
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(
|
def make_manager(
|
||||||
*,
|
*,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
config: CanaryConfig | None = None,
|
config: CanaryConfig | None = None,
|
||||||
|
perturb_config: PerturbConfig | None = None,
|
||||||
group: CanaryBufferGroup | None = None,
|
group: CanaryBufferGroup | None = None,
|
||||||
req_pool: SimpleNamespace | None = None,
|
req_pool: SimpleNamespace | None = None,
|
||||||
per_forward_verify_capacity: int = 16,
|
per_forward_verify_capacity: int = 16,
|
||||||
@@ -55,12 +66,15 @@ def make_manager(
|
|||||||
) -> CanaryManager:
|
) -> CanaryManager:
|
||||||
if config is None:
|
if config is None:
|
||||||
config = make_config()
|
config = make_config()
|
||||||
|
if perturb_config is None:
|
||||||
|
perturb_config = make_perturb_config()
|
||||||
if group is None:
|
if group is None:
|
||||||
group = make_buffer_group(device=device)
|
group = make_buffer_group(device=device)
|
||||||
if req_pool is None:
|
if req_pool is None:
|
||||||
req_pool = make_req_to_token_pool(device=device, max_reqs=4, max_seq_len=8)
|
req_pool = make_req_to_token_pool(device=device, max_reqs=4, max_seq_len=8)
|
||||||
return CanaryManager(
|
return CanaryManager(
|
||||||
config=config,
|
config=config,
|
||||||
|
perturb_config=perturb_config,
|
||||||
buffer_groups=(group,),
|
buffer_groups=(group,),
|
||||||
device=device,
|
device=device,
|
||||||
req_to_token_pool=req_pool,
|
req_to_token_pool=req_pool,
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from __future__ import annotations
|
|||||||
import time
|
import time
|
||||||
from typing import Literal, Optional
|
from typing import Literal, Optional
|
||||||
|
|
||||||
|
from sglang.srt.kv_canary.perturb.config import TargetGroupKind
|
||||||
from sglang.test.kv_canary.violation_log_utils import (
|
from sglang.test.kv_canary.violation_log_utils import (
|
||||||
assert_no_violation_in_log,
|
assert_no_violation_in_log,
|
||||||
find_violation_in_log,
|
find_violation_in_log,
|
||||||
@@ -19,16 +20,35 @@ class CanaryViolationAssertMixin:
|
|||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
fail_reason: str,
|
fail_reason: str,
|
||||||
|
target_group: Optional[TargetGroupKind] = None,
|
||||||
side: _Side = None,
|
side: _Side = None,
|
||||||
flush_wait_seconds: float = 2.0,
|
flush_wait_seconds: float = 2.0,
|
||||||
) -> None:
|
) -> None:
|
||||||
|
suffix = "" if target_group is None else f"_{target_group.name}"
|
||||||
self.assert_violation_logged_any(
|
self.assert_violation_logged_any(
|
||||||
launch_tag_patterns=("HEAD_*", "TAIL_*"),
|
launch_tag_patterns=(f"HEAD_*{suffix}", f"TAIL_*{suffix}"),
|
||||||
fail_reason=fail_reason,
|
fail_reason=fail_reason,
|
||||||
side=side,
|
side=side,
|
||||||
flush_wait_seconds=flush_wait_seconds,
|
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(
|
def assert_any_launch_tag_violation_reported(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
|
|||||||
@@ -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()
|
||||||
Reference in New Issue
Block a user