Add a sliding-window-attention divergence reporter for the KV-canary (#26820)
This commit is contained in:
@@ -763,6 +763,7 @@ class Envs:
|
|||||||
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_VERIFY_TOKEN_ASSERT = EnvBool(False)
|
SGLANG_KV_CANARY_ENABLE_VERIFY_TOKEN_ASSERT = EnvBool(False)
|
||||||
|
SGLANG_KV_CANARY_SWA_DIVERGENCE_STATS_INTERVAL = EnvInt(0)
|
||||||
SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False)
|
SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ 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
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -48,6 +49,10 @@ def install_canary(
|
|||||||
device=device,
|
device=device,
|
||||||
kv_token_id_vs_position_offset=kv_token_id_vs_position_offset,
|
kv_token_id_vs_position_offset=kv_token_id_vs_position_offset,
|
||||||
)
|
)
|
||||||
|
allocator = model_runner.token_to_kv_pool_allocator
|
||||||
|
swa_allocator = (
|
||||||
|
allocator if isinstance(allocator, SWATokenToKVPoolAllocator) else None
|
||||||
|
)
|
||||||
launch_capacities = CanaryLaunchCapacities.from_args(
|
launch_capacities = CanaryLaunchCapacities.from_args(
|
||||||
server_args=model_runner.server_args,
|
server_args=model_runner.server_args,
|
||||||
req_to_token_pool_size=model_runner.req_to_token_pool.size,
|
req_to_token_pool_size=model_runner.req_to_token_pool.size,
|
||||||
@@ -65,6 +70,7 @@ def install_canary(
|
|||||||
launch_capacities=launch_capacities,
|
launch_capacities=launch_capacities,
|
||||||
swa_window_size=swa_window_size,
|
swa_window_size=swa_window_size,
|
||||||
token_oracle_manager=token_oracle_manager,
|
token_oracle_manager=token_oracle_manager,
|
||||||
|
swa_allocator=swa_allocator,
|
||||||
speculative_num_steps=speculative_num_steps,
|
speculative_num_steps=speculative_num_steps,
|
||||||
is_eagle_draft_decode=model_runner.is_draft_worker,
|
is_eagle_draft_decode=model_runner.is_draft_worker,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from sglang.srt.kv_canary.endpoint import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.kv_canary.perturb.config import PerturbConfig
|
from sglang.srt.kv_canary.perturb.config import PerturbConfig
|
||||||
from sglang.srt.kv_canary.perturb.manager import PerturbManager
|
from sglang.srt.kv_canary.perturb.manager import PerturbManager
|
||||||
|
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.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 (
|
||||||
@@ -30,6 +31,7 @@ from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
|
|||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -47,12 +49,14 @@ class CanaryManager:
|
|||||||
launch_capacities: CanaryLaunchCapacities,
|
launch_capacities: CanaryLaunchCapacities,
|
||||||
swa_window_size: int = 0,
|
swa_window_size: int = 0,
|
||||||
token_oracle_manager: Optional[TokenOracleManager] = None,
|
token_oracle_manager: Optional[TokenOracleManager] = None,
|
||||||
|
swa_allocator: Optional["SWATokenToKVPoolAllocator"] = None,
|
||||||
speculative_num_steps: int = 1,
|
speculative_num_steps: int = 1,
|
||||||
is_eagle_draft_decode: bool = False,
|
is_eagle_draft_decode: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.config = config
|
self.config = config
|
||||||
self._req_to_token_pool = req_to_token_pool
|
self._req_to_token_pool = req_to_token_pool
|
||||||
self._swa_window_size = swa_window_size
|
self._swa_window_size = swa_window_size
|
||||||
|
self._swa_allocator: Optional["SWATokenToKVPoolAllocator"] = swa_allocator
|
||||||
self._outer_step_counter: int = 0
|
self._outer_step_counter: int = 0
|
||||||
self._active_single_forward_manager_index: Optional[int] = None
|
self._active_single_forward_manager_index: Optional[int] = None
|
||||||
|
|
||||||
@@ -85,6 +89,22 @@ class CanaryManager:
|
|||||||
|
|
||||||
self._d2h_stream: torch.cuda.Stream = torch.cuda.Stream(device=device)
|
self._d2h_stream: torch.cuda.Stream = torch.cuda.Stream(device=device)
|
||||||
|
|
||||||
|
swa_divergence_interval = (
|
||||||
|
envs.SGLANG_KV_CANARY_SWA_DIVERGENCE_STATS_INTERVAL.get()
|
||||||
|
)
|
||||||
|
if swa_divergence_interval > 0:
|
||||||
|
self._swa_divergence_report: Optional[SwaDivergenceReporter] = (
|
||||||
|
SwaDivergenceReporter(
|
||||||
|
device=device,
|
||||||
|
d2h_stream=self._d2h_stream,
|
||||||
|
interval=swa_divergence_interval,
|
||||||
|
swa_allocator=self._swa_allocator,
|
||||||
|
req_to_token_pool=self._req_to_token_pool,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self._swa_divergence_report = None
|
||||||
|
|
||||||
self._violation_manager = ViolationManager(
|
self._violation_manager = ViolationManager(
|
||||||
config=config,
|
config=config,
|
||||||
device_state=self._device_state,
|
device_state=self._device_state,
|
||||||
@@ -122,6 +142,7 @@ class CanaryManager:
|
|||||||
per_forward_write_entry_capacity=launch_capacities.per_forward_write_entry_capacity,
|
per_forward_write_entry_capacity=launch_capacities.per_forward_write_entry_capacity,
|
||||||
d2h_stream=self._d2h_stream,
|
d2h_stream=self._d2h_stream,
|
||||||
token_oracle_manager=token_oracle_manager,
|
token_oracle_manager=token_oracle_manager,
|
||||||
|
swa_divergence_report=self._swa_divergence_report,
|
||||||
is_eagle_draft_decode=is_eagle_draft_decode,
|
is_eagle_draft_decode=is_eagle_draft_decode,
|
||||||
)
|
)
|
||||||
for _ in range(num_sfms)
|
for _ in range(num_sfms)
|
||||||
@@ -212,6 +233,11 @@ class CanaryManager:
|
|||||||
self._sweep_orchestrator.maybe_run_sweep()
|
self._sweep_orchestrator.maybe_run_sweep()
|
||||||
self._outer_step_counter += 1
|
self._outer_step_counter += 1
|
||||||
self._violation_manager.step()
|
self._violation_manager.step()
|
||||||
|
if self._swa_divergence_report is not None:
|
||||||
|
self._swa_divergence_report.step(
|
||||||
|
outer_step_counter=self._outer_step_counter,
|
||||||
|
maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch,
|
||||||
|
)
|
||||||
|
|
||||||
def mark_init_finished(self) -> None:
|
def mark_init_finished(self) -> None:
|
||||||
for single_forward_manager in self._single_forward_managers:
|
for single_forward_manager in self._single_forward_managers:
|
||||||
|
|||||||
@@ -0,0 +1,202 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
import logging
|
||||||
|
import re
|
||||||
|
from dataclasses import asdict, dataclass
|
||||||
|
from typing import TYPE_CHECKING, Any, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.jit_kernel.kv_canary.verify import VerifyPlan
|
||||||
|
from sglang.srt.kv_canary.buffer_group import CanaryBufferGroup, PoolKind
|
||||||
|
from sglang.srt.kv_canary.runner.future_tensor import DelayedDeviceHostHandler
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
|
from sglang.srt.mem_cache.swa_memory_pool import SWATokenToKVPoolAllocator
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_SWA_DIVERGENCE_LOG_PREFIX: str = "kv_canary_swa_divergence="
|
||||||
|
_SWA_DIVERGENCE_LINE_RE = re.compile(re.escape(_SWA_DIVERGENCE_LOG_PREFIX) + r"(\S+)")
|
||||||
|
_FULL_IDX = 0
|
||||||
|
_SWA_IDX = 1
|
||||||
|
|
||||||
|
|
||||||
|
class SwaDivergenceReporter:
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
device: torch.device,
|
||||||
|
d2h_stream: torch.cuda.Stream,
|
||||||
|
interval: int,
|
||||||
|
swa_allocator: Optional["SWATokenToKVPoolAllocator"] = None,
|
||||||
|
req_to_token_pool: Optional["ReqToTokenPool"] = None,
|
||||||
|
) -> None:
|
||||||
|
self._interval = interval
|
||||||
|
self._swa_allocator = swa_allocator
|
||||||
|
self._req_to_token_pool = req_to_token_pool
|
||||||
|
self._forward_ct: int = 0
|
||||||
|
# Per-group running total of verify entries (shape ``[2]``, int32).
|
||||||
|
self.verify_total_count_device: torch.Tensor = torch.zeros(
|
||||||
|
2, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
self._handler = DelayedDeviceHostHandler(d2h_stream=d2h_stream)
|
||||||
|
|
||||||
|
def observe_after_invoke_plan(
|
||||||
|
self, *, group: CanaryBufferGroup, verify_plan: VerifyPlan
|
||||||
|
) -> None:
|
||||||
|
idx = _FULL_IDX if group.kind is PoolKind.FULL else _SWA_IDX
|
||||||
|
# verify_num_valid is shape [1]; slice to a length-1 view so the in-place add
|
||||||
|
# has matching ranks (else torch refuses the broadcast into shape []).
|
||||||
|
self.verify_total_count_device[idx : idx + 1].add_(verify_plan.verify_num_valid)
|
||||||
|
|
||||||
|
def step(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
outer_step_counter: int,
|
||||||
|
maybe_inaccurate_forward_batch: Optional["ForwardBatch"],
|
||||||
|
) -> None:
|
||||||
|
self._forward_ct += 1
|
||||||
|
self._handler.step(
|
||||||
|
compute_on_device=lambda: self._compute_on_device(
|
||||||
|
outer_step_counter=outer_step_counter,
|
||||||
|
maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch,
|
||||||
|
),
|
||||||
|
postprocess_on_host=self._postprocess_on_host,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _compute_on_device(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
outer_step_counter: int,
|
||||||
|
maybe_inaccurate_forward_batch: Optional["ForwardBatch"],
|
||||||
|
) -> Optional[dict[str, Any]]:
|
||||||
|
if outer_step_counter == 0 or outer_step_counter % self._interval != 0:
|
||||||
|
return None
|
||||||
|
|
||||||
|
result: dict[str, Any] = {
|
||||||
|
"forward_ct": self._forward_ct,
|
||||||
|
"verify_total_count": self.verify_total_count_device,
|
||||||
|
}
|
||||||
|
if (
|
||||||
|
self._swa_allocator is not None
|
||||||
|
and maybe_inaccurate_forward_batch is not None
|
||||||
|
):
|
||||||
|
result["swa_full_idx_divergence"] = compute_swa_full_idx_divergence(
|
||||||
|
swa_allocator=self._swa_allocator,
|
||||||
|
req_to_token_pool=self._req_to_token_pool,
|
||||||
|
maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch,
|
||||||
|
)
|
||||||
|
result["swa_out_of_window_tokens"] = compute_swa_out_of_window_tokens(
|
||||||
|
swa_allocator=self._swa_allocator,
|
||||||
|
req_to_token_pool=self._req_to_token_pool,
|
||||||
|
maybe_inaccurate_forward_batch=maybe_inaccurate_forward_batch,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
def _postprocess_on_host(self, host_data: dict[str, Any]) -> None:
|
||||||
|
verify_totals = host_data["verify_total_count"].tolist()
|
||||||
|
swa_full_idx_divergence = (
|
||||||
|
int(x.item())
|
||||||
|
if (x := host_data.get("swa_full_idx_divergence")) is not None
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
swa_out_of_window_tokens = (
|
||||||
|
int(x.item())
|
||||||
|
if (x := host_data.get("swa_out_of_window_tokens")) is not None
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
logger.info(
|
||||||
|
SwaDivergenceLog(
|
||||||
|
forward_ct=host_data["forward_ct"],
|
||||||
|
verify_full=int(verify_totals[_FULL_IDX]),
|
||||||
|
verify_swa=int(verify_totals[_SWA_IDX]),
|
||||||
|
swa_full_idx_divergence=swa_full_idx_divergence,
|
||||||
|
swa_out_of_window_tokens=swa_out_of_window_tokens,
|
||||||
|
).format()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||||
|
class SwaDivergenceLog:
|
||||||
|
forward_ct: int
|
||||||
|
verify_full: int
|
||||||
|
verify_swa: int
|
||||||
|
swa_full_idx_divergence: int
|
||||||
|
swa_out_of_window_tokens: int = 0
|
||||||
|
|
||||||
|
def format(self) -> str:
|
||||||
|
return _SWA_DIVERGENCE_LOG_PREFIX + json.dumps(
|
||||||
|
asdict(self), separators=(",", ":")
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def parse(cls, line: str) -> Optional["SwaDivergenceLog"]:
|
||||||
|
match = _SWA_DIVERGENCE_LINE_RE.search(line)
|
||||||
|
if match is None:
|
||||||
|
return None
|
||||||
|
return cls(**json.loads(match.group(1)))
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def find_last(cls, text: str) -> Optional[tuple["SwaDivergenceLog", str]]:
|
||||||
|
last_match: Optional[re.Match] = None
|
||||||
|
for match in _SWA_DIVERGENCE_LINE_RE.finditer(text):
|
||||||
|
last_match = match
|
||||||
|
if last_match is None:
|
||||||
|
return None
|
||||||
|
return cls(**json.loads(last_match.group(1))), last_match.group(0)
|
||||||
|
|
||||||
|
|
||||||
|
def compute_swa_out_of_window_tokens(
|
||||||
|
*,
|
||||||
|
swa_allocator: "SWATokenToKVPoolAllocator",
|
||||||
|
req_to_token_pool: "ReqToTokenPool",
|
||||||
|
maybe_inaccurate_forward_batch: "ForwardBatch",
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Count tokens in the live req_to_token range whose SWA mapping is 0 (out-of-window)."""
|
||||||
|
full_to_swa_index_mapping = swa_allocator.full_to_swa_index_mapping
|
||||||
|
device = full_to_swa_index_mapping.device
|
||||||
|
req_pool_indices = maybe_inaccurate_forward_batch.req_pool_indices
|
||||||
|
seq_lens = maybe_inaccurate_forward_batch.seq_lens
|
||||||
|
if req_pool_indices.numel() == 0:
|
||||||
|
return torch.zeros(1, dtype=torch.int32, device=device)
|
||||||
|
req_to_token = req_to_token_pool.req_to_token
|
||||||
|
rows = req_to_token[req_pool_indices]
|
||||||
|
positions = torch.arange(rows.shape[1], device=rows.device)
|
||||||
|
mask = positions[None, :] < seq_lens[:, None]
|
||||||
|
swa_indices = full_to_swa_index_mapping[rows]
|
||||||
|
return ((swa_indices == 0) & mask).sum().to(torch.int32).view(1)
|
||||||
|
|
||||||
|
|
||||||
|
def compute_swa_full_idx_divergence(
|
||||||
|
*,
|
||||||
|
swa_allocator: "SWATokenToKVPoolAllocator",
|
||||||
|
req_to_token_pool: "ReqToTokenPool",
|
||||||
|
maybe_inaccurate_forward_batch: "ForwardBatch",
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Count non-identity (full, swa) index pairs in the live req_to_token range."""
|
||||||
|
full_to_swa_index_mapping = swa_allocator.full_to_swa_index_mapping
|
||||||
|
device = full_to_swa_index_mapping.device
|
||||||
|
req_pool_indices = maybe_inaccurate_forward_batch.req_pool_indices
|
||||||
|
seq_lens = maybe_inaccurate_forward_batch.seq_lens
|
||||||
|
|
||||||
|
if req_pool_indices.numel() == 0:
|
||||||
|
return torch.zeros(1, dtype=torch.int32, device=device)
|
||||||
|
|
||||||
|
req_to_token = req_to_token_pool.req_to_token
|
||||||
|
rows = req_to_token[req_pool_indices]
|
||||||
|
positions = torch.arange(rows.shape[1], device=rows.device)
|
||||||
|
mask = positions[None, :] < seq_lens[:, None]
|
||||||
|
swa_indices = full_to_swa_index_mapping[rows]
|
||||||
|
# FULL pool slots beyond the sliding window have their SWA mapping written
|
||||||
|
# to 0 (see SWATokenToKVPoolAllocator.alloc_extend); skip those so they
|
||||||
|
# don't get counted as divergence.
|
||||||
|
return (
|
||||||
|
((swa_indices != rows) & mask & (swa_indices != 0))
|
||||||
|
.sum()
|
||||||
|
.to(torch.int32)
|
||||||
|
.view(1)
|
||||||
|
)
|
||||||
@@ -11,6 +11,7 @@ class PostOpsInsideGraphOutputBuffer:
|
|||||||
kernel_run_counters: torch.Tensor
|
kernel_run_counters: torch.Tensor
|
||||||
slot_run_counters: torch.Tensor
|
slot_run_counters: torch.Tensor
|
||||||
violation_write_index: torch.Tensor
|
violation_write_index: torch.Tensor
|
||||||
|
swa_verify_total_count: torch.Tensor | None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def allocate(
|
def allocate(
|
||||||
@@ -18,6 +19,7 @@ class PostOpsInsideGraphOutputBuffer:
|
|||||||
*,
|
*,
|
||||||
num_kernel_tags: int,
|
num_kernel_tags: int,
|
||||||
num_slot_tags: int,
|
num_slot_tags: int,
|
||||||
|
swa_verify_total_count_shape: tuple[int, ...] | None,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
) -> "PostOpsInsideGraphOutputBuffer":
|
) -> "PostOpsInsideGraphOutputBuffer":
|
||||||
return cls(
|
return cls(
|
||||||
@@ -29,6 +31,13 @@ class PostOpsInsideGraphOutputBuffer:
|
|||||||
num_slot_tags, dtype=torch.int64, device=device
|
num_slot_tags, dtype=torch.int64, device=device
|
||||||
),
|
),
|
||||||
violation_write_index=torch.zeros(1, dtype=torch.int32, device=device),
|
violation_write_index=torch.zeros(1, dtype=torch.int32, device=device),
|
||||||
|
swa_verify_total_count=(
|
||||||
|
None
|
||||||
|
if swa_verify_total_count_shape is None
|
||||||
|
else torch.zeros(
|
||||||
|
swa_verify_total_count_shape, dtype=torch.int32, device=device
|
||||||
|
)
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def copy_from(
|
def copy_from(
|
||||||
@@ -38,8 +47,14 @@ class PostOpsInsideGraphOutputBuffer:
|
|||||||
kernel_run_counters: torch.Tensor,
|
kernel_run_counters: torch.Tensor,
|
||||||
slot_run_counters: torch.Tensor,
|
slot_run_counters: torch.Tensor,
|
||||||
violation_write_index: torch.Tensor,
|
violation_write_index: torch.Tensor,
|
||||||
|
swa_verify_total_count: torch.Tensor | None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.verify_plan_enable.copy_(verify_plan_enable)
|
self.verify_plan_enable.copy_(verify_plan_enable)
|
||||||
self.kernel_run_counters.copy_(kernel_run_counters)
|
self.kernel_run_counters.copy_(kernel_run_counters)
|
||||||
self.slot_run_counters.copy_(slot_run_counters)
|
self.slot_run_counters.copy_(slot_run_counters)
|
||||||
self.violation_write_index.copy_(violation_write_index)
|
self.violation_write_index.copy_(violation_write_index)
|
||||||
|
assert (self.swa_verify_total_count is not None) == (
|
||||||
|
swa_verify_total_count is not None
|
||||||
|
)
|
||||||
|
if self.swa_verify_total_count is not None:
|
||||||
|
self.swa_verify_total_count.copy_(swa_verify_total_count)
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ from sglang.srt.kv_canary.runner.kernel_launcher import (
|
|||||||
invoke_plan,
|
invoke_plan,
|
||||||
launch_endpoints_per_forward,
|
launch_endpoints_per_forward,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.kv_canary.runner.swa_divergence import SwaDivergenceReporter
|
||||||
from sglang.srt.kv_canary.single_forward_manager.data import (
|
from sglang.srt.kv_canary.single_forward_manager.data import (
|
||||||
PostOpsInsideGraphOutputBuffer,
|
PostOpsInsideGraphOutputBuffer,
|
||||||
)
|
)
|
||||||
@@ -72,6 +73,7 @@ class SingleForwardManager:
|
|||||||
per_forward_write_entry_capacity: int,
|
per_forward_write_entry_capacity: int,
|
||||||
d2h_stream: torch.cuda.Stream,
|
d2h_stream: torch.cuda.Stream,
|
||||||
token_oracle_manager: Optional[TokenOracleManager],
|
token_oracle_manager: Optional[TokenOracleManager],
|
||||||
|
swa_divergence_report: Optional[SwaDivergenceReporter],
|
||||||
is_eagle_draft_decode: bool,
|
is_eagle_draft_decode: bool,
|
||||||
) -> None:
|
) -> None:
|
||||||
self._config = config
|
self._config = config
|
||||||
@@ -83,6 +85,9 @@ class SingleForwardManager:
|
|||||||
self._swa_window_size = swa_window_size
|
self._swa_window_size = swa_window_size
|
||||||
self._d2h_stream = d2h_stream
|
self._d2h_stream = d2h_stream
|
||||||
self._token_oracle_manager: Optional[TokenOracleManager] = token_oracle_manager
|
self._token_oracle_manager: Optional[TokenOracleManager] = token_oracle_manager
|
||||||
|
self._swa_divergence_report: Optional[SwaDivergenceReporter] = (
|
||||||
|
swa_divergence_report
|
||||||
|
)
|
||||||
self._is_eagle_draft_decode: bool = is_eagle_draft_decode
|
self._is_eagle_draft_decode: bool = is_eagle_draft_decode
|
||||||
|
|
||||||
self._write_req_capacity = per_forward_write_req_capacity
|
self._write_req_capacity = per_forward_write_req_capacity
|
||||||
@@ -101,6 +106,11 @@ class SingleForwardManager:
|
|||||||
self._output_buffer = PostOpsInsideGraphOutputBuffer.allocate(
|
self._output_buffer = PostOpsInsideGraphOutputBuffer.allocate(
|
||||||
num_kernel_tags=int(device_state.kernel_run_counters.shape[0]),
|
num_kernel_tags=int(device_state.kernel_run_counters.shape[0]),
|
||||||
num_slot_tags=int(device_state.slot_run_counters.shape[0]),
|
num_slot_tags=int(device_state.slot_run_counters.shape[0]),
|
||||||
|
swa_verify_total_count_shape=(
|
||||||
|
None
|
||||||
|
if swa_divergence_report is None
|
||||||
|
else tuple(swa_divergence_report.verify_total_count_device.shape)
|
||||||
|
),
|
||||||
device=device,
|
device=device,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -201,6 +211,11 @@ class SingleForwardManager:
|
|||||||
swa_window_size=self._swa_window_size,
|
swa_window_size=self._swa_window_size,
|
||||||
req_to_verify_expected_tokens=self._device_state.req_to_verify_expected_tokens,
|
req_to_verify_expected_tokens=self._device_state.req_to_verify_expected_tokens,
|
||||||
)
|
)
|
||||||
|
if self._swa_divergence_report is not None:
|
||||||
|
self._swa_divergence_report.observe_after_invoke_plan(
|
||||||
|
group=group,
|
||||||
|
verify_plan=verify_plan,
|
||||||
|
)
|
||||||
launch_endpoints_per_forward(
|
launch_endpoints_per_forward(
|
||||||
endpoints=self._endpoints,
|
endpoints=self._endpoints,
|
||||||
group=group,
|
group=group,
|
||||||
@@ -261,6 +276,11 @@ class SingleForwardManager:
|
|||||||
kernel_run_counters=self._device_state.kernel_run_counters,
|
kernel_run_counters=self._device_state.kernel_run_counters,
|
||||||
slot_run_counters=self._device_state.slot_run_counters,
|
slot_run_counters=self._device_state.slot_run_counters,
|
||||||
violation_write_index=self._device_state.violation_log.violation_write_index,
|
violation_write_index=self._device_state.violation_log.violation_write_index,
|
||||||
|
swa_verify_total_count=(
|
||||||
|
None
|
||||||
|
if self._swa_divergence_report is None
|
||||||
|
else self._swa_divergence_report.verify_total_count_device
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def post_ops_outside_graph(self) -> None:
|
def post_ops_outside_graph(self) -> None:
|
||||||
|
|||||||
@@ -3,9 +3,11 @@ from __future__ import annotations
|
|||||||
import io
|
import io
|
||||||
import os
|
import os
|
||||||
import string
|
import string
|
||||||
|
import time
|
||||||
from typing import ClassVar, Literal, Optional
|
from typing import ClassVar, Literal, Optional
|
||||||
|
|
||||||
from sglang.srt.kv_canary.config import CanaryMode
|
from sglang.srt.kv_canary.config import CanaryMode
|
||||||
|
from sglang.srt.kv_canary.runner.swa_divergence import SwaDivergenceLog
|
||||||
from sglang.srt.utils import kill_process_tree
|
from sglang.srt.utils import kill_process_tree
|
||||||
from sglang.test.kv_canary.mode_config import _MODE_CONFIGS, _ModeConfig
|
from sglang.test.kv_canary.mode_config import _MODE_CONFIGS, _ModeConfig
|
||||||
from sglang.test.kv_canary.utils import build_canary_server_args, post_parallel_generate
|
from sglang.test.kv_canary.utils import build_canary_server_args, post_parallel_generate
|
||||||
@@ -81,6 +83,9 @@ class CanaryE2EBase(CapturedServerE2EBase):
|
|||||||
server_env.setdefault("SGLANG_KV_CANARY_ENABLE_VERIFY_TOKEN_ASSERT", "1")
|
server_env.setdefault("SGLANG_KV_CANARY_ENABLE_VERIFY_TOKEN_ASSERT", "1")
|
||||||
server_env.update(cls.extra_env)
|
server_env.update(cls.extra_env)
|
||||||
if cls.model_mode == "swa":
|
if cls.model_mode == "swa":
|
||||||
|
server_env.setdefault(
|
||||||
|
"SGLANG_KV_CANARY_SWA_DIVERGENCE_STATS_INTERVAL", "20"
|
||||||
|
)
|
||||||
# SWA mode uses google/gemma-4-E2B-it, whose forward does a
|
# SWA mode uses google/gemma-4-E2B-it, whose forward does a
|
||||||
# ``positions += 1`` in-place. canary's WRITE/VERIFY require
|
# ``positions += 1`` in-place. canary's WRITE/VERIFY require
|
||||||
# forward_batch.positions to stay 0-indexed, so flip the gemma
|
# forward_batch.positions to stay 0-indexed, so flip the gemma
|
||||||
@@ -134,6 +139,70 @@ class CanaryE2EBase(CapturedServerE2EBase):
|
|||||||
self.assertEqual(result.get("status_code"), 200, result)
|
self.assertEqual(result.get("status_code"), 200, result)
|
||||||
return results
|
return results
|
||||||
|
|
||||||
|
def maybe_assert_swa_divergence_observed(self) -> None:
|
||||||
|
if self.model_mode == "swa":
|
||||||
|
self.assert_swa_divergence_observed()
|
||||||
|
|
||||||
|
def assert_swa_divergence_observed(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
min_swa_out_of_window_tokens: int = 1,
|
||||||
|
min_swa_full_idx_divergence: int = 1,
|
||||||
|
require_verify_lag: bool = True,
|
||||||
|
flush_wait_seconds: float = 3.0,
|
||||||
|
max_retries: int = 10,
|
||||||
|
) -> None:
|
||||||
|
"""Assert that the SWA path was genuinely exercised.
|
||||||
|
|
||||||
|
Three signals must all hold:
|
||||||
|
- ``swa_out_of_window_tokens >= 1``: at least one prefix token has been clipped
|
||||||
|
out of the sliding window (its SWA mapping is 0). Any prompt longer than the
|
||||||
|
SWA window produces this — proves the SWA window slide actually ran.
|
||||||
|
- ``swa_full_idx_divergence >= 1``: SWA pool has actually remapped at least one
|
||||||
|
slot to a non-identity index (i.e. real slot reuse / eviction occurred). The
|
||||||
|
workload must drive SWA pool pressure for this to fire — required because the
|
||||||
|
"pool reuse" path is the one production hits under sustained long-context
|
||||||
|
traffic, and we must keep it covered.
|
||||||
|
- ``verify_swa < verify_full``: SWA verify kernel processed fewer tokens than
|
||||||
|
FULL — proves both kernel groups ran and the window short-circuited SWA.
|
||||||
|
"""
|
||||||
|
last_parsed = None
|
||||||
|
last_line: str = ""
|
||||||
|
for _ in range(max_retries):
|
||||||
|
time.sleep(flush_wait_seconds)
|
||||||
|
log_text = self._captured_log_text()
|
||||||
|
found = SwaDivergenceLog.find_last(log_text)
|
||||||
|
if found is not None:
|
||||||
|
last_parsed, last_line = found
|
||||||
|
break
|
||||||
|
|
||||||
|
if last_parsed is None:
|
||||||
|
raise AssertionError(
|
||||||
|
"No kv_canary swa_divergence line found in server log after "
|
||||||
|
f"{max_retries} retries (wait={flush_wait_seconds}s each). "
|
||||||
|
f"Log tail:\n{self._captured_log_text()[-2000:]}"
|
||||||
|
)
|
||||||
|
|
||||||
|
if last_parsed.swa_out_of_window_tokens < min_swa_out_of_window_tokens:
|
||||||
|
raise AssertionError(
|
||||||
|
f"SWA path not exercised: swa_out_of_window_tokens={last_parsed.swa_out_of_window_tokens} "
|
||||||
|
f"< min={min_swa_out_of_window_tokens}. Line: {last_line}"
|
||||||
|
)
|
||||||
|
if last_parsed.swa_full_idx_divergence < min_swa_full_idx_divergence:
|
||||||
|
raise AssertionError(
|
||||||
|
f"SWA pool reuse not exercised: swa_full_idx_divergence={last_parsed.swa_full_idx_divergence} "
|
||||||
|
f"< min={min_swa_full_idx_divergence}. The workload did not drive enough SWA pool pressure "
|
||||||
|
f"to force slot remap. Line: {last_line}"
|
||||||
|
)
|
||||||
|
if require_verify_lag and not (
|
||||||
|
last_parsed.verify_swa < last_parsed.verify_full
|
||||||
|
):
|
||||||
|
raise AssertionError(
|
||||||
|
f"SWA path not exercised: verify_swa={last_parsed.verify_swa} "
|
||||||
|
f"not strictly less than verify_full={last_parsed.verify_full}. "
|
||||||
|
f"Line: {last_line}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _make_unique_prompts(n: int) -> list[str]:
|
def _make_unique_prompts(n: int) -> list[str]:
|
||||||
if n > len(_UNIQUE_PROMPT_FIRST_CHARS):
|
if n > len(_UNIQUE_PROMPT_FIRST_CHARS):
|
||||||
|
|||||||
@@ -28,6 +28,7 @@ class _BaselineBase(CanaryE2EBase):
|
|||||||
for _ in range(self.workload_n_batches):
|
for _ in range(self.workload_n_batches):
|
||||||
self.send_parallel_requests()
|
self.send_parallel_requests()
|
||||||
self.assert_no_violation(wait_seconds=2.0)
|
self.assert_no_violation(wait_seconds=2.0)
|
||||||
|
self.maybe_assert_swa_divergence_observed()
|
||||||
|
|
||||||
|
|
||||||
class TestBaselineMha(_BaselineBase):
|
class TestBaselineMha(_BaselineBase):
|
||||||
|
|||||||
@@ -51,6 +51,7 @@ class _PerturbRealKvUnusedCacheBase(CanaryE2EBase):
|
|||||||
target_group=self.target_group,
|
target_group=self.target_group,
|
||||||
flush_wait_seconds=5.0,
|
flush_wait_seconds=5.0,
|
||||||
)
|
)
|
||||||
|
self.maybe_assert_swa_divergence_observed()
|
||||||
|
|
||||||
|
|
||||||
class TestPerturbRealKvUnusedCacheMhaFull(_PerturbRealKvUnusedCacheBase):
|
class TestPerturbRealKvUnusedCacheMhaFull(_PerturbRealKvUnusedCacheBase):
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ class _PerturbRealKvUsedBase(CanaryE2EBase):
|
|||||||
fail_reason="verify_real_kv_hash",
|
fail_reason="verify_real_kv_hash",
|
||||||
target_group=self.target_group,
|
target_group=self.target_group,
|
||||||
)
|
)
|
||||||
|
self.maybe_assert_swa_divergence_observed()
|
||||||
|
|
||||||
|
|
||||||
class TestPerturbRealKvUsedMhaFull(_PerturbRealKvUsedBase):
|
class TestPerturbRealKvUsedMhaFull(_PerturbRealKvUsedBase):
|
||||||
|
|||||||
@@ -34,6 +34,7 @@ class _PerturbReqToTokenBase(CanaryE2EBase):
|
|||||||
for _ in range(self.workload_n_batches):
|
for _ in range(self.workload_n_batches):
|
||||||
self.send_parallel_requests()
|
self.send_parallel_requests()
|
||||||
self.assert_per_forward_violation_reported(fail_reason="verify_chain_hash")
|
self.assert_per_forward_violation_reported(fail_reason="verify_chain_hash")
|
||||||
|
self.maybe_assert_swa_divergence_observed()
|
||||||
|
|
||||||
|
|
||||||
class TestPerturbReqToTokenMha(_PerturbReqToTokenBase):
|
class TestPerturbReqToTokenMha(_PerturbReqToTokenBase):
|
||||||
|
|||||||
@@ -0,0 +1,163 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from sglang.srt.kv_canary.runner.swa_divergence import SwaDivergenceLog
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.kv_canary.e2e_base import CanaryE2EBase
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-b-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
_GOOD_LINE: str = SwaDivergenceLog(
|
||||||
|
forward_ct=120,
|
||||||
|
verify_full=10000,
|
||||||
|
verify_swa=4200,
|
||||||
|
swa_full_idx_divergence=512,
|
||||||
|
swa_out_of_window_tokens=8192,
|
||||||
|
).format()
|
||||||
|
_LATER_LINE: str = SwaDivergenceLog(
|
||||||
|
forward_ct=240,
|
||||||
|
verify_full=20000,
|
||||||
|
verify_swa=8400,
|
||||||
|
swa_full_idx_divergence=1024,
|
||||||
|
swa_out_of_window_tokens=16384,
|
||||||
|
).format()
|
||||||
|
|
||||||
|
|
||||||
|
class _DummyHarness(CanaryE2EBase):
|
||||||
|
model_mode = "swa"
|
||||||
|
kv_canary_mode = "log"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def setUpClass(cls) -> None:
|
||||||
|
return
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def tearDownClass(cls) -> None:
|
||||||
|
return
|
||||||
|
|
||||||
|
|
||||||
|
class TestAssertSwaDivergenceObserved(CustomTestCase):
|
||||||
|
def _make_harness(
|
||||||
|
self, log_text_or_sequence
|
||||||
|
) -> tuple[_DummyHarness, "patch._patch[None]"]:
|
||||||
|
harness = _DummyHarness()
|
||||||
|
harness._stderr_buf = None
|
||||||
|
harness._stdout_buf = None
|
||||||
|
if isinstance(log_text_or_sequence, list):
|
||||||
|
patcher = patch.object(
|
||||||
|
_DummyHarness, "_captured_log_text", side_effect=log_text_or_sequence
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
patcher = patch.object(
|
||||||
|
_DummyHarness,
|
||||||
|
"_captured_log_text",
|
||||||
|
return_value=log_text_or_sequence,
|
||||||
|
)
|
||||||
|
return harness, patcher
|
||||||
|
|
||||||
|
def test_assert_swa_divergence_observed_passes_when_above_threshold(self) -> None:
|
||||||
|
harness, patcher = self._make_harness(_LATER_LINE + "\n" + _GOOD_LINE + "\n")
|
||||||
|
with patcher:
|
||||||
|
harness.assert_swa_divergence_observed(
|
||||||
|
min_swa_full_idx_divergence=100,
|
||||||
|
require_verify_lag=True,
|
||||||
|
flush_wait_seconds=0.0,
|
||||||
|
max_retries=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_assert_swa_divergence_observed_uses_latest_line(self) -> None:
|
||||||
|
log = _GOOD_LINE + "\n" + _LATER_LINE + "\n"
|
||||||
|
harness, patcher = self._make_harness(log)
|
||||||
|
with patcher:
|
||||||
|
harness.assert_swa_divergence_observed(
|
||||||
|
min_swa_full_idx_divergence=1000,
|
||||||
|
require_verify_lag=True,
|
||||||
|
flush_wait_seconds=0.0,
|
||||||
|
max_retries=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_assert_swa_divergence_observed_raises_when_below_threshold(self) -> None:
|
||||||
|
zero_mapping_line = SwaDivergenceLog(
|
||||||
|
forward_ct=100,
|
||||||
|
verify_full=5000,
|
||||||
|
verify_swa=2000,
|
||||||
|
swa_full_idx_divergence=0,
|
||||||
|
swa_out_of_window_tokens=8192,
|
||||||
|
).format()
|
||||||
|
harness, patcher = self._make_harness(zero_mapping_line + "\n")
|
||||||
|
with patcher:
|
||||||
|
with self.assertRaisesRegex(AssertionError, "swa_full_idx_divergence=0"):
|
||||||
|
harness.assert_swa_divergence_observed(
|
||||||
|
min_swa_full_idx_divergence=1,
|
||||||
|
require_verify_lag=False,
|
||||||
|
flush_wait_seconds=0.0,
|
||||||
|
max_retries=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_assert_swa_divergence_observed_raises_when_no_verify_lag(self) -> None:
|
||||||
|
equal_verify_line = SwaDivergenceLog(
|
||||||
|
forward_ct=100,
|
||||||
|
verify_full=5000,
|
||||||
|
verify_swa=5000,
|
||||||
|
swa_full_idx_divergence=200,
|
||||||
|
swa_out_of_window_tokens=8192,
|
||||||
|
).format()
|
||||||
|
harness, patcher = self._make_harness(equal_verify_line + "\n")
|
||||||
|
with patcher:
|
||||||
|
with self.assertRaisesRegex(AssertionError, "verify_swa=5000"):
|
||||||
|
harness.assert_swa_divergence_observed(
|
||||||
|
min_swa_full_idx_divergence=1,
|
||||||
|
require_verify_lag=True,
|
||||||
|
flush_wait_seconds=0.0,
|
||||||
|
max_retries=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_assert_swa_divergence_observed_retries_until_stats_emitted(self) -> None:
|
||||||
|
sequence = ["", "", "", _GOOD_LINE + "\n", _GOOD_LINE + "\n"]
|
||||||
|
harness, patcher = self._make_harness(sequence)
|
||||||
|
with patcher:
|
||||||
|
harness.assert_swa_divergence_observed(
|
||||||
|
min_swa_full_idx_divergence=1,
|
||||||
|
require_verify_lag=True,
|
||||||
|
flush_wait_seconds=0.0,
|
||||||
|
max_retries=5,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_assert_swa_divergence_observed_raises_when_no_stats_emitted(self) -> None:
|
||||||
|
harness, patcher = self._make_harness("nothing here\n")
|
||||||
|
with patcher:
|
||||||
|
with self.assertRaisesRegex(AssertionError, "No kv_canary swa_divergence"):
|
||||||
|
harness.assert_swa_divergence_observed(
|
||||||
|
min_swa_full_idx_divergence=1,
|
||||||
|
require_verify_lag=True,
|
||||||
|
flush_wait_seconds=0.0,
|
||||||
|
max_retries=2,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_assert_swa_divergence_observed_catches_zero_swa_full_idx_divergence(
|
||||||
|
self,
|
||||||
|
) -> None:
|
||||||
|
zero_divergence_line = SwaDivergenceLog(
|
||||||
|
forward_ct=200,
|
||||||
|
verify_full=10000,
|
||||||
|
verify_swa=2000,
|
||||||
|
swa_full_idx_divergence=0,
|
||||||
|
swa_out_of_window_tokens=8192,
|
||||||
|
).format()
|
||||||
|
harness, patcher = self._make_harness(zero_divergence_line + "\n")
|
||||||
|
with patcher:
|
||||||
|
with self.assertRaisesRegex(AssertionError, "swa_full_idx_divergence=0"):
|
||||||
|
harness.assert_swa_divergence_observed(
|
||||||
|
min_swa_full_idx_divergence=1,
|
||||||
|
require_verify_lag=True,
|
||||||
|
flush_wait_seconds=0.0,
|
||||||
|
max_retries=1,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,431 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.jit_kernel.kv_canary.verify import VerifyPlan
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.kv_canary.buffer_group import PoolKind
|
||||||
|
from sglang.srt.kv_canary.runner import swa_divergence as swa_div_module
|
||||||
|
from sglang.srt.kv_canary.runner.swa_divergence import (
|
||||||
|
SwaDivergenceLog,
|
||||||
|
SwaDivergenceReporter,
|
||||||
|
compute_swa_full_idx_divergence,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
from sglang.test.kv_canary.fixtures import make_buffer_group
|
||||||
|
from sglang.test.kv_canary.runner_test_base import CanaryManagerTestCase, make_manager
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=45, stage="extra-a", runner_config="1-gpu-small")
|
||||||
|
|
||||||
|
_DEVICE = torch.device("cuda")
|
||||||
|
|
||||||
|
_EMPTY_FORWARD_BATCH = SimpleNamespace(
|
||||||
|
req_pool_indices=torch.empty(0, dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.empty(0, dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_verify_plan(value: int) -> VerifyPlan:
|
||||||
|
plan = VerifyPlan.allocate(verify_capacity=4, device=_DEVICE)
|
||||||
|
plan.verify_num_valid.copy_(torch.tensor([value], dtype=torch.int32))
|
||||||
|
return plan
|
||||||
|
|
||||||
|
|
||||||
|
def _make_allocator_stub(mapping: torch.Tensor) -> SimpleNamespace:
|
||||||
|
return SimpleNamespace(full_to_swa_index_mapping=mapping)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_req_to_token_pool_stub(req_to_token: torch.Tensor) -> SimpleNamespace:
|
||||||
|
return SimpleNamespace(req_to_token=req_to_token)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_identity_mapping(size: int) -> torch.Tensor:
|
||||||
|
return torch.arange(size, dtype=torch.int64, device=_DEVICE)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_identity_req_to_token(num_reqs: int, max_seq_len: int) -> torch.Tensor:
|
||||||
|
base = torch.arange(num_reqs * max_seq_len, dtype=torch.int64, device=_DEVICE)
|
||||||
|
return base.view(num_reqs, max_seq_len)
|
||||||
|
|
||||||
|
|
||||||
|
def _make_forward_batch(
|
||||||
|
*, req_pool_indices: torch.Tensor, seq_lens: torch.Tensor
|
||||||
|
) -> SimpleNamespace:
|
||||||
|
return SimpleNamespace(req_pool_indices=req_pool_indices, seq_lens=seq_lens)
|
||||||
|
|
||||||
|
|
||||||
|
def _parse_swa_divergence_line(line: str) -> SwaDivergenceLog:
|
||||||
|
parsed = SwaDivergenceLog.parse(line)
|
||||||
|
if parsed is None:
|
||||||
|
raise AssertionError(f"line does not match swa_divergence format: {line!r}")
|
||||||
|
return parsed
|
||||||
|
|
||||||
|
|
||||||
|
def _run_compute(
|
||||||
|
*,
|
||||||
|
swa_allocator: SimpleNamespace,
|
||||||
|
req_to_token_pool: SimpleNamespace,
|
||||||
|
forward_batch: SimpleNamespace,
|
||||||
|
) -> int:
|
||||||
|
count = compute_swa_full_idx_divergence(
|
||||||
|
swa_allocator=swa_allocator,
|
||||||
|
req_to_token_pool=req_to_token_pool,
|
||||||
|
maybe_inaccurate_forward_batch=forward_batch,
|
||||||
|
)
|
||||||
|
return int(count.item())
|
||||||
|
|
||||||
|
|
||||||
|
class TestSwaDivergenceReporter(CustomTestCase):
|
||||||
|
def test_swa_divergence_log_emitted(self) -> None:
|
||||||
|
d2h_stream = torch.cuda.Stream(device=_DEVICE)
|
||||||
|
stats = SwaDivergenceReporter(
|
||||||
|
device=_DEVICE,
|
||||||
|
d2h_stream=d2h_stream,
|
||||||
|
interval=10,
|
||||||
|
swa_allocator=None,
|
||||||
|
req_to_token_pool=None,
|
||||||
|
)
|
||||||
|
# First 3 forwards stay below the interval trigger (1, 2, 3 % 10 != 0) so
|
||||||
|
# step() just bumps forward_ct and stages nothing.
|
||||||
|
for forward_idx in range(3):
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.FULL, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(10),
|
||||||
|
)
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.SWA, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(3),
|
||||||
|
)
|
||||||
|
stats.step(
|
||||||
|
outer_step_counter=forward_idx + 1,
|
||||||
|
maybe_inaccurate_forward_batch=_EMPTY_FORWARD_BATCH,
|
||||||
|
)
|
||||||
|
# 4th forward lands on outer_step_counter=10 = interval, so compute_on_device
|
||||||
|
# snapshots {forward_ct:4, verify_full:40, verify_swa:12} into the dict and
|
||||||
|
# the staged future hangs onto it. forward_ct is now 4.
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.FULL, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(10),
|
||||||
|
)
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.SWA, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(3),
|
||||||
|
)
|
||||||
|
stats.step(
|
||||||
|
outer_step_counter=10, maybe_inaccurate_forward_batch=_EMPTY_FORWARD_BATCH
|
||||||
|
)
|
||||||
|
|
||||||
|
# 5th step drains the previous stage and emits the log; forward_ct is now 5
|
||||||
|
# but the staged dict still carries the snapshot forward_ct=4 from step 4.
|
||||||
|
with self.assertLogs(
|
||||||
|
swa_div_module.logger.name, level=logging.INFO
|
||||||
|
) as captured:
|
||||||
|
stats.step(
|
||||||
|
outer_step_counter=11,
|
||||||
|
maybe_inaccurate_forward_batch=_EMPTY_FORWARD_BATCH,
|
||||||
|
)
|
||||||
|
|
||||||
|
lines = [
|
||||||
|
line for line in captured.output if SwaDivergenceLog.parse(line) is not None
|
||||||
|
]
|
||||||
|
self.assertEqual(len(lines), 1, lines)
|
||||||
|
fields = _parse_swa_divergence_line(lines[0])
|
||||||
|
self.assertEqual(fields.forward_ct, 4)
|
||||||
|
self.assertEqual(fields.verify_full, 40)
|
||||||
|
self.assertEqual(fields.verify_swa, 12)
|
||||||
|
self.assertEqual(fields.swa_full_idx_divergence, 0)
|
||||||
|
|
||||||
|
def test_swa_divergence_counts_monotonic_increasing(self) -> None:
|
||||||
|
d2h_stream = torch.cuda.Stream(device=_DEVICE)
|
||||||
|
stats = SwaDivergenceReporter(
|
||||||
|
device=_DEVICE,
|
||||||
|
d2h_stream=d2h_stream,
|
||||||
|
interval=10,
|
||||||
|
swa_allocator=None,
|
||||||
|
req_to_token_pool=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
snapshots: list[SwaDivergenceLog] = []
|
||||||
|
|
||||||
|
def _take_snapshot(stage_step: int, drain_step: int) -> None:
|
||||||
|
# Stage the dict at the interval-aligned step (no log emitted yet,
|
||||||
|
# DelayedDeviceHostHandler still has nothing to drain), then call
|
||||||
|
# step() again at the next counter to drain and emit the log.
|
||||||
|
stats.step(
|
||||||
|
outer_step_counter=stage_step,
|
||||||
|
maybe_inaccurate_forward_batch=_EMPTY_FORWARD_BATCH,
|
||||||
|
)
|
||||||
|
with self.assertLogs(
|
||||||
|
swa_div_module.logger.name, level=logging.INFO
|
||||||
|
) as captured:
|
||||||
|
stats.step(
|
||||||
|
outer_step_counter=drain_step,
|
||||||
|
maybe_inaccurate_forward_batch=_EMPTY_FORWARD_BATCH,
|
||||||
|
)
|
||||||
|
matching = [
|
||||||
|
line
|
||||||
|
for line in captured.output
|
||||||
|
if SwaDivergenceLog.parse(line) is not None
|
||||||
|
]
|
||||||
|
self.assertTrue(matching, captured.output)
|
||||||
|
snapshots.append(_parse_swa_divergence_line(matching[-1]))
|
||||||
|
|
||||||
|
for batch in range(3):
|
||||||
|
for _ in range(5):
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.FULL, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(7),
|
||||||
|
)
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.SWA, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(2),
|
||||||
|
)
|
||||||
|
stage_step = 10 + 20 * batch
|
||||||
|
_take_snapshot(stage_step=stage_step, drain_step=stage_step + 1)
|
||||||
|
|
||||||
|
for idx in range(1, len(snapshots)):
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
snapshots[idx].verify_full, snapshots[idx - 1].verify_full
|
||||||
|
)
|
||||||
|
self.assertGreaterEqual(
|
||||||
|
snapshots[idx].verify_swa, snapshots[idx - 1].verify_swa
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSwaFullIdxDivergenceCompute(CustomTestCase):
|
||||||
|
def test_compute_returns_zero_when_empty_batch(self) -> None:
|
||||||
|
mapping = _make_identity_mapping(size=64)
|
||||||
|
req_to_token = _make_identity_req_to_token(num_reqs=4, max_seq_len=16)
|
||||||
|
|
||||||
|
forward_batch = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.empty(0, dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.empty(0, dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_run_compute(
|
||||||
|
swa_allocator=_make_allocator_stub(mapping),
|
||||||
|
req_to_token_pool=_make_req_to_token_pool_stub(req_to_token),
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
),
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_compute_returns_zero_when_all_identity(self) -> None:
|
||||||
|
mapping = _make_identity_mapping(size=64)
|
||||||
|
req_to_token = _make_identity_req_to_token(num_reqs=4, max_seq_len=16)
|
||||||
|
|
||||||
|
forward_batch = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.tensor([0, 2], dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.tensor([8, 5], dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_run_compute(
|
||||||
|
swa_allocator=_make_allocator_stub(mapping),
|
||||||
|
req_to_token_pool=_make_req_to_token_pool_stub(req_to_token),
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
),
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_compute_counts_swa_full_idx_divergence_in_live_range(self) -> None:
|
||||||
|
mapping = _make_identity_mapping(size=64)
|
||||||
|
req_to_token = _make_identity_req_to_token(num_reqs=4, max_seq_len=16)
|
||||||
|
|
||||||
|
mapping[0] = 50
|
||||||
|
mapping[1] = 51
|
||||||
|
mapping[17] = 60
|
||||||
|
|
||||||
|
forward_batch = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.tensor([0, 1], dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.tensor([8, 8], dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_run_compute(
|
||||||
|
swa_allocator=_make_allocator_stub(mapping),
|
||||||
|
req_to_token_pool=_make_req_to_token_pool_stub(req_to_token),
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
),
|
||||||
|
3,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_compute_ignores_swa_mapping_zero(self) -> None:
|
||||||
|
# SWATokenToKVPoolAllocator writes 0 into full_to_swa_index_mapping for
|
||||||
|
# FULL pool slots beyond the sliding window. Those entries are expected,
|
||||||
|
# not real divergence, so the count must skip them.
|
||||||
|
mapping = _make_identity_mapping(size=64)
|
||||||
|
req_to_token = _make_identity_req_to_token(num_reqs=4, max_seq_len=16)
|
||||||
|
|
||||||
|
mapping[3] = 0
|
||||||
|
mapping[5] = 0
|
||||||
|
mapping[7] = 42
|
||||||
|
|
||||||
|
forward_batch = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.tensor([0], dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.tensor([8], dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_run_compute(
|
||||||
|
swa_allocator=_make_allocator_stub(mapping),
|
||||||
|
req_to_token_pool=_make_req_to_token_pool_stub(req_to_token),
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
),
|
||||||
|
1,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_compute_ignores_writes_outside_seq_lens(self) -> None:
|
||||||
|
mapping = _make_identity_mapping(size=128)
|
||||||
|
req_to_token = _make_identity_req_to_token(num_reqs=4, max_seq_len=32)
|
||||||
|
|
||||||
|
mapping[20] = 99
|
||||||
|
mapping[28] = 77
|
||||||
|
|
||||||
|
forward_batch = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.tensor([0], dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.tensor([10], dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_run_compute(
|
||||||
|
swa_allocator=_make_allocator_stub(mapping),
|
||||||
|
req_to_token_pool=_make_req_to_token_pool_stub(req_to_token),
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
),
|
||||||
|
0,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_compute_reflects_current_forward_batch(self) -> None:
|
||||||
|
mapping = _make_identity_mapping(size=64)
|
||||||
|
req_to_token = _make_identity_req_to_token(num_reqs=4, max_seq_len=16)
|
||||||
|
|
||||||
|
mapping[0] = 41
|
||||||
|
mapping[1] = 42
|
||||||
|
mapping[32] = 99
|
||||||
|
mapping[33] = 100
|
||||||
|
|
||||||
|
fb_req0 = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.tensor([0], dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.tensor([4], dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
fb_req2 = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.tensor([2], dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.tensor([4], dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
_run_compute(
|
||||||
|
swa_allocator=_make_allocator_stub(mapping),
|
||||||
|
req_to_token_pool=_make_req_to_token_pool_stub(req_to_token),
|
||||||
|
forward_batch=fb_req0,
|
||||||
|
),
|
||||||
|
2,
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
_run_compute(
|
||||||
|
swa_allocator=_make_allocator_stub(mapping),
|
||||||
|
req_to_token_pool=_make_req_to_token_pool_stub(req_to_token),
|
||||||
|
forward_batch=fb_req2,
|
||||||
|
),
|
||||||
|
2,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestSwaDivergenceReporterWithCompute(CustomTestCase):
|
||||||
|
def test_swa_divergence_report_emits_swa_full_idx_divergence_from_compute(
|
||||||
|
self,
|
||||||
|
) -> None:
|
||||||
|
mapping = _make_identity_mapping(size=64)
|
||||||
|
req_to_token = _make_identity_req_to_token(num_reqs=4, max_seq_len=16)
|
||||||
|
|
||||||
|
mapping[0] = 50
|
||||||
|
mapping[1] = 51
|
||||||
|
mapping[2] = 52
|
||||||
|
|
||||||
|
forward_batch = _make_forward_batch(
|
||||||
|
req_pool_indices=torch.tensor([0], dtype=torch.int64, device=_DEVICE),
|
||||||
|
seq_lens=torch.tensor([8], dtype=torch.int64, device=_DEVICE),
|
||||||
|
)
|
||||||
|
|
||||||
|
swa_allocator = _make_allocator_stub(mapping)
|
||||||
|
req_to_token_pool = _make_req_to_token_pool_stub(req_to_token)
|
||||||
|
d2h_stream = torch.cuda.Stream(device=_DEVICE)
|
||||||
|
stats = SwaDivergenceReporter(
|
||||||
|
device=_DEVICE,
|
||||||
|
d2h_stream=d2h_stream,
|
||||||
|
interval=10,
|
||||||
|
swa_allocator=swa_allocator,
|
||||||
|
req_to_token_pool=req_to_token_pool,
|
||||||
|
)
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.FULL, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(11),
|
||||||
|
)
|
||||||
|
stats.observe_after_invoke_plan(
|
||||||
|
group=make_buffer_group(
|
||||||
|
device=_DEVICE, kind=PoolKind.SWA, has_v=False, num_slots=1
|
||||||
|
),
|
||||||
|
verify_plan=_make_verify_plan(3),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Stage at the interval-aligned step, then drain on the next step so the
|
||||||
|
# DelayedDeviceHostHandler has a pending future to postprocess.
|
||||||
|
stats.step(outer_step_counter=10, maybe_inaccurate_forward_batch=forward_batch)
|
||||||
|
with self.assertLogs(
|
||||||
|
swa_div_module.logger.name, level=logging.INFO
|
||||||
|
) as captured:
|
||||||
|
stats.step(
|
||||||
|
outer_step_counter=11, maybe_inaccurate_forward_batch=forward_batch
|
||||||
|
)
|
||||||
|
|
||||||
|
matching = [
|
||||||
|
line for line in captured.output if SwaDivergenceLog.parse(line) is not None
|
||||||
|
]
|
||||||
|
self.assertEqual(len(matching), 1, matching)
|
||||||
|
parsed = SwaDivergenceLog.parse(matching[0])
|
||||||
|
assert parsed is not None
|
||||||
|
self.assertEqual(parsed.swa_full_idx_divergence, 3)
|
||||||
|
self.assertEqual(parsed.verify_full, 11)
|
||||||
|
self.assertEqual(parsed.verify_swa, 3)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCanaryManagerSwaDivergenceWiring(CanaryManagerTestCase):
|
||||||
|
def test_swa_divergence_report_is_none_when_env_disabled(self) -> None:
|
||||||
|
with envs.SGLANG_KV_CANARY_SWA_DIVERGENCE_STATS_INTERVAL.override(
|
||||||
|
0
|
||||||
|
), envs.SGLANG_KV_CANARY_PERTURB_TARGET_GROUP.override("full"):
|
||||||
|
manager = make_manager(device=self.device)
|
||||||
|
self.assertIsNone(manager._swa_divergence_report)
|
||||||
|
|
||||||
|
def test_swa_divergence_report_present_when_env_enabled(self) -> None:
|
||||||
|
with envs.SGLANG_KV_CANARY_SWA_DIVERGENCE_STATS_INTERVAL.override(
|
||||||
|
20
|
||||||
|
), envs.SGLANG_KV_CANARY_PERTURB_TARGET_GROUP.override("full"):
|
||||||
|
manager = make_manager(device=self.device)
|
||||||
|
self.assertIsNotNone(manager._swa_divergence_report)
|
||||||
|
self.assertIsInstance(manager._swa_divergence_report, SwaDivergenceReporter)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user