Add a deterministic token oracle and production write-input assertion (#26815)
This commit is contained in:
@@ -750,9 +750,12 @@ class Envs:
|
||||
SGLANG_PLUGINS = EnvStr("")
|
||||
|
||||
# ===================================================================
|
||||
# KV-Canary (testing-only)
|
||||
# KV-Canary / Token-Oracle (testing-only)
|
||||
# ===================================================================
|
||||
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_NEXT_TOKEN_SWAP_PROB = EnvFloat(0.0)
|
||||
SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE = EnvBool(False)
|
||||
SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False)
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ from sglang.srt.kv_canary.runner.canary_manager import CanaryManager
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
@@ -23,6 +24,7 @@ def install_canary(
|
||||
*,
|
||||
server_args: "ServerArgs",
|
||||
model_runner: "ModelRunner",
|
||||
token_oracle_manager: Optional["TokenOracleManager"] = None,
|
||||
) -> Optional[CanaryManager]:
|
||||
config = CanaryConfig.from_env(server_args)
|
||||
if config.mode is CanaryMode.NONE:
|
||||
@@ -59,7 +61,9 @@ def install_canary(
|
||||
req_to_token_pool=model_runner.req_to_token_pool,
|
||||
launch_capacities=launch_capacities,
|
||||
swa_window_size=swa_window_size,
|
||||
token_oracle_manager=token_oracle_manager,
|
||||
speculative_num_steps=speculative_num_steps,
|
||||
is_eagle_draft_decode=model_runner.is_draft_worker,
|
||||
)
|
||||
|
||||
_patch_model_forward(model_runner=model_runner, manager=manager)
|
||||
|
||||
@@ -20,7 +20,7 @@ class CanaryMode(str, Enum):
|
||||
class CanaryConfig:
|
||||
"""Top-level canary configuration. All knobs live here; nothing reads env vars deeper in the stack.
|
||||
|
||||
Constructed once inside install_canary(server_args, model_runner) via
|
||||
Constructed once inside install_canary(server_args, model_runner, token_oracle_manager) via
|
||||
CanaryConfig.from_env(server_args), then frozen and threaded through the canary stack.
|
||||
Subsequent runtime never mutates it.
|
||||
|
||||
@@ -33,11 +33,17 @@ class CanaryConfig:
|
||||
sweep_interval: 0 disables sweep entirely; positive N means every N-th forward step the runner
|
||||
additionally walks all radix-tree-held slots (overlap with per-forward HEAD/TAIL is harmless
|
||||
redundancy) and verifies them.
|
||||
enable_write_input_assert: bool. True = launch_canary_write_kernel additionally compares
|
||||
forward_batch.input_ids[i] / positions[i] against caller-supplied expected_input_tokens[i] /
|
||||
expected_input_positions[i]; mismatch records a violation. Only useful when something else
|
||||
(e.g. token_oracle.oracle_manager.fill_expected_inputs) is feeding the expected_* placeholders
|
||||
per forward — canary itself knows no oracle.
|
||||
"""
|
||||
|
||||
mode: CanaryMode
|
||||
ring_capacity: int
|
||||
sweep_interval: int
|
||||
enable_write_input_assert: bool
|
||||
|
||||
@classmethod
|
||||
def from_env(cls, server_args: "ServerArgs") -> "CanaryConfig":
|
||||
@@ -51,4 +57,5 @@ class CanaryConfig:
|
||||
mode=CanaryMode(mode_raw),
|
||||
ring_capacity=envs.SGLANG_KV_CANARY_RING_CAPACITY.get(),
|
||||
sweep_interval=server_args.kv_canary_sweep_interval,
|
||||
enable_write_input_assert=envs.SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT.get(),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
"""Swap two requests' sampled next tokens at the sampler exit.
|
||||
|
||||
KV path is untouched, so kv_canary KV-side fail_reasons stay silent. The
|
||||
token-oracle input check downstream MUST report fail_reason=write_token — this
|
||||
validates that the input-check link is genuinely active.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||
class NextTokenSwapConfig:
|
||||
prob: float
|
||||
warmup_steps: int
|
||||
|
||||
@classmethod
|
||||
def from_env(cls) -> "NextTokenSwapConfig":
|
||||
return cls(
|
||||
prob=envs.SGLANG_KV_CANARY_PERTURB_NEXT_TOKEN_SWAP_PROB.get(),
|
||||
warmup_steps=envs.SGLANG_KV_CANARY_PERTURB_WARMUP_STEPS.get(),
|
||||
)
|
||||
|
||||
|
||||
_config: Optional[NextTokenSwapConfig] = None
|
||||
_step_counter: int = 0
|
||||
|
||||
|
||||
def _get_config() -> NextTokenSwapConfig:
|
||||
global _config
|
||||
if _config is None:
|
||||
_config = NextTokenSwapConfig.from_env()
|
||||
return _config
|
||||
|
||||
|
||||
def maybe_perturb_swap_next_tokens(
|
||||
batch_next_token_ids: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
global _step_counter
|
||||
|
||||
config = _get_config()
|
||||
step = _step_counter
|
||||
_step_counter += 1
|
||||
|
||||
if config.prob <= 0.0:
|
||||
return batch_next_token_ids
|
||||
if step < config.warmup_steps:
|
||||
return batch_next_token_ids
|
||||
if batch_next_token_ids.shape[0] < 2:
|
||||
return batch_next_token_ids
|
||||
|
||||
if random.random() >= config.prob:
|
||||
return batch_next_token_ids
|
||||
|
||||
batch_size = batch_next_token_ids.shape[0]
|
||||
i = random.randrange(batch_size)
|
||||
j = random.randrange(batch_size)
|
||||
while j == i:
|
||||
j = random.randrange(batch_size)
|
||||
|
||||
swapped = batch_next_token_ids.clone()
|
||||
swapped[i], swapped[j] = (
|
||||
batch_next_token_ids[j].clone(),
|
||||
batch_next_token_ids[i].clone(),
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"kv_canary perturb next_token_swap: swapped i=%d j=%d step=%d",
|
||||
i,
|
||||
j,
|
||||
step,
|
||||
)
|
||||
return swapped
|
||||
@@ -23,6 +23,7 @@ from sglang.srt.kv_canary.single_forward_manager.manager import (
|
||||
_PreOpsMaybeInsideGraphOutput,
|
||||
)
|
||||
from sglang.srt.kv_canary.state import CanaryDeviceState
|
||||
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
@@ -42,7 +43,9 @@ class CanaryManager:
|
||||
req_to_token_pool: "ReqToTokenPool",
|
||||
launch_capacities: CanaryLaunchCapacities,
|
||||
swa_window_size: int = 0,
|
||||
token_oracle_manager: Optional[TokenOracleManager] = None,
|
||||
speculative_num_steps: int = 1,
|
||||
is_eagle_draft_decode: bool = False,
|
||||
) -> None:
|
||||
self.config = config
|
||||
self._req_to_token_pool = req_to_token_pool
|
||||
@@ -107,6 +110,8 @@ class CanaryManager:
|
||||
per_forward_write_req_capacity=launch_capacities.per_forward_write_req_capacity,
|
||||
per_forward_write_entry_capacity=launch_capacities.per_forward_write_entry_capacity,
|
||||
d2h_stream=self._d2h_stream,
|
||||
token_oracle_manager=token_oracle_manager,
|
||||
is_eagle_draft_decode=is_eagle_draft_decode,
|
||||
)
|
||||
for _ in range(num_sfms)
|
||||
)
|
||||
|
||||
@@ -64,8 +64,8 @@ def launch_endpoints_per_forward(
|
||||
forward_batch: "ForwardBatch",
|
||||
expected_inputs: ExpectedInputs,
|
||||
violation_log: ViolationLog,
|
||||
enable_write_input_assert: bool = False,
|
||||
enable_verify_token_assert: bool = False,
|
||||
enable_write_input_assert: bool,
|
||||
enable_verify_token_assert: bool,
|
||||
) -> None:
|
||||
positions = _canonicalize_boundary_int64(forward_batch.positions, _POSITIONS)
|
||||
out_cache_loc = _canonicalize_boundary_int64(forward_batch.out_cache_loc, _OUT_LOC)
|
||||
|
||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -22,6 +22,7 @@ from sglang.srt.kv_canary.single_forward_manager.data import (
|
||||
PostOpsInsideGraphOutputBuffer,
|
||||
)
|
||||
from sglang.srt.kv_canary.state import CanaryDeviceState
|
||||
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
|
||||
from sglang.srt.utils.phase_checker import SimplePhaseChecker
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -67,6 +68,8 @@ class SingleForwardManager:
|
||||
per_forward_write_req_capacity: int,
|
||||
per_forward_write_entry_capacity: int,
|
||||
d2h_stream: torch.cuda.Stream,
|
||||
token_oracle_manager: Optional[TokenOracleManager],
|
||||
is_eagle_draft_decode: bool,
|
||||
) -> None:
|
||||
self._config = config
|
||||
self._device = device
|
||||
@@ -76,6 +79,8 @@ class SingleForwardManager:
|
||||
self._req_to_token_pool = req_to_token_pool
|
||||
self._swa_window_size = swa_window_size
|
||||
self._d2h_stream = d2h_stream
|
||||
self._token_oracle_manager: Optional[TokenOracleManager] = token_oracle_manager
|
||||
self._is_eagle_draft_decode: bool = is_eagle_draft_decode
|
||||
|
||||
self._write_req_capacity = per_forward_write_req_capacity
|
||||
self._write_entry_capacity = per_forward_write_entry_capacity
|
||||
@@ -153,6 +158,22 @@ class SingleForwardManager:
|
||||
bs_capacity=self._write_req_capacity, device=self._device
|
||||
)
|
||||
|
||||
enable_write_input_assert = self._should_enable_write_input_assert_for_launch(
|
||||
forward_batch
|
||||
)
|
||||
if enable_write_input_assert:
|
||||
manager = self._token_oracle_manager
|
||||
if manager is None:
|
||||
raise RuntimeError(
|
||||
"kv-canary: enable_write_input_assert=True requires a TokenOracleManager; pass "
|
||||
"token_oracle_manager=install_oracle_sampler(oracle=...) into "
|
||||
"install_canary(...)"
|
||||
)
|
||||
manager.fill_expected_inputs(
|
||||
forward_batch=forward_batch,
|
||||
expected_inputs_out=expected_inputs,
|
||||
)
|
||||
|
||||
plan_input.fill_from_forward_batch(forward_batch=forward_batch)
|
||||
|
||||
violation_log = self._device_state.violation_log
|
||||
@@ -180,6 +201,8 @@ class SingleForwardManager:
|
||||
forward_batch=forward_batch,
|
||||
expected_inputs=expected_inputs_slice,
|
||||
violation_log=violation_log,
|
||||
enable_write_input_assert=enable_write_input_assert,
|
||||
enable_verify_token_assert=False,
|
||||
)
|
||||
|
||||
return _PreOpsMaybeInsideGraphOutput(
|
||||
@@ -202,6 +225,9 @@ class SingleForwardManager:
|
||||
violation_log = self._device_state.violation_log
|
||||
num_tokens = int(forward_batch.positions.shape[0])
|
||||
expected_inputs_slice = pre_ops_output.expected_inputs.slice(num_tokens)
|
||||
enable_write_input_assert = self._should_enable_write_input_assert_for_launch(
|
||||
forward_batch
|
||||
)
|
||||
for group_idx, group in enumerate(self._buffer_groups):
|
||||
launch_endpoints_per_forward(
|
||||
endpoints=self._endpoints,
|
||||
@@ -212,6 +238,8 @@ class SingleForwardManager:
|
||||
forward_batch=forward_batch,
|
||||
expected_inputs=expected_inputs_slice,
|
||||
violation_log=violation_log,
|
||||
enable_write_input_assert=enable_write_input_assert,
|
||||
enable_verify_token_assert=False,
|
||||
)
|
||||
|
||||
verify_plan_enable_combined = _torch_reduce_minimum(
|
||||
@@ -233,6 +261,20 @@ class SingleForwardManager:
|
||||
|
||||
self._enable_warner.tick(self._output_buffer.verify_plan_enable)
|
||||
|
||||
def _should_enable_write_input_assert_for_launch(
|
||||
self, forward_batch: "ForwardBatch"
|
||||
) -> bool:
|
||||
if not self._config.enable_write_input_assert:
|
||||
return False
|
||||
forward_mode = forward_batch.forward_mode
|
||||
if (
|
||||
self._is_eagle_draft_decode
|
||||
and forward_mode is not None
|
||||
and forward_mode.is_decode()
|
||||
):
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _is_head_tag(tag: CanaryLaunchTag) -> bool:
|
||||
return tag in (
|
||||
|
||||
@@ -0,0 +1,22 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from sglang.srt.kv_canary.token_oracle.oracle import HashOracle
|
||||
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
|
||||
from sglang.srt.kv_canary.token_oracle.sampler import install_oracle_sampler
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
|
||||
def install_token_oracle_from_env(
|
||||
*, server_args: "ServerArgs", vocab_size: int
|
||||
) -> Optional[TokenOracleManager]:
|
||||
# Must be called before create_sampler() so the factory is present when the
|
||||
# Sampler is first constructed.
|
||||
if server_args.sampling_backend != "token_oracle":
|
||||
return None
|
||||
|
||||
oracle = HashOracle(vocab_size=vocab_size)
|
||||
return install_oracle_sampler(oracle=oracle)
|
||||
@@ -0,0 +1,50 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Protocol
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class TokenOracle(Protocol):
|
||||
"""Deterministic (generalized_req_id, position) -> token_id mapping."""
|
||||
|
||||
def expected_tokens(
|
||||
self, *, generalized_req_ids: torch.Tensor, positions: torch.Tensor
|
||||
) -> torch.Tensor: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||
class HashOracle:
|
||||
"""token_id = splitmix64(generalized_req_id XOR position) % vocab_size."""
|
||||
|
||||
vocab_size: int
|
||||
|
||||
def expected_tokens(
|
||||
self, *, generalized_req_ids: torch.Tensor, positions: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
x = generalized_req_ids.to(torch.int64) ^ positions.to(torch.int64)
|
||||
x = _splitmix64_tensor(x)
|
||||
return _uint64_mod(x, self.vocab_size).to(torch.int32)
|
||||
|
||||
|
||||
_C1: int = -4658895280553007687 # 0xBF58476D1CE4E5B9 as signed int64
|
||||
_C2: int = -7723592293110705685 # 0x94D049BB133111EB as signed int64
|
||||
|
||||
|
||||
def _splitmix64_tensor(x: torch.Tensor) -> torch.Tensor:
|
||||
x = (x ^ _logical_shr(x, 30)) * _C1
|
||||
x = (x ^ _logical_shr(x, 27)) * _C2
|
||||
x = x ^ _logical_shr(x, 31)
|
||||
return x
|
||||
|
||||
|
||||
def _logical_shr(x: torch.Tensor, n: int) -> torch.Tensor:
|
||||
return (x >> n) & ((1 << (64 - n)) - 1)
|
||||
|
||||
|
||||
def _uint64_mod(x: torch.Tensor, mod: int) -> torch.Tensor:
|
||||
offset = (1 << 64) % mod
|
||||
base = x % mod
|
||||
correction = (x < 0).to(x.dtype) * offset
|
||||
return (base + correction) % mod
|
||||
@@ -0,0 +1,110 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.kv_canary.expected_inputs import ExpectedInputs
|
||||
from sglang.srt.kv_canary.token_oracle.oracle import TokenOracle
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
|
||||
class TokenOracleManager:
|
||||
def __init__(self, *, oracle: TokenOracle) -> None:
|
||||
self.oracle = oracle
|
||||
|
||||
def fill_expected_inputs(
|
||||
self,
|
||||
*,
|
||||
forward_batch: "ForwardBatch",
|
||||
expected_inputs_out: ExpectedInputs,
|
||||
) -> None:
|
||||
positions = forward_batch.positions
|
||||
input_ids = forward_batch.input_ids
|
||||
num_tokens = int(input_ids.shape[0])
|
||||
|
||||
if num_tokens == 0:
|
||||
return
|
||||
|
||||
generalized_req_ids = _build_generalized_req_id_per_token(
|
||||
forward_batch=forward_batch,
|
||||
num_tokens=num_tokens,
|
||||
generalized_req_ids_per_row=select_generalized_req_ids(
|
||||
vanilla_req_ids=forward_batch.rids_int,
|
||||
bootstrap_room_ids_int=forward_batch.bootstrap_room_ids_int,
|
||||
),
|
||||
)
|
||||
if forward_batch.forward_mode.is_extend():
|
||||
expected_tokens = input_ids
|
||||
else:
|
||||
expected_tokens = self.oracle.expected_tokens(
|
||||
generalized_req_ids=generalized_req_ids,
|
||||
positions=positions.to(torch.int64),
|
||||
)
|
||||
expected_inputs_out.tokens[:num_tokens].copy_(expected_tokens.to(torch.int64))
|
||||
expected_inputs_out.positions[:num_tokens].copy_(positions.to(torch.int64))
|
||||
|
||||
def sample_next_tokens(
|
||||
self, *, generalized_req_ids: torch.Tensor, logits_positions: torch.Tensor
|
||||
) -> torch.Tensor:
|
||||
return self.oracle.expected_tokens(
|
||||
generalized_req_ids=generalized_req_ids,
|
||||
positions=logits_positions.to(torch.int64) + 1,
|
||||
)
|
||||
|
||||
|
||||
def _build_generalized_req_id_per_token(
|
||||
*,
|
||||
forward_batch: "ForwardBatch",
|
||||
num_tokens: int,
|
||||
generalized_req_ids_per_row: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
forward_mode = forward_batch.forward_mode
|
||||
if forward_mode.is_target_verify():
|
||||
per_req = int(forward_batch.spec_info.draft_token_num)
|
||||
result = _expand_uniform(generalized_req_ids_per_row, per_req)
|
||||
elif forward_mode.is_draft_extend(include_v2=True):
|
||||
per_req = int(forward_batch.spec_info.num_tokens_per_req)
|
||||
result = _expand_uniform(generalized_req_ids_per_row, per_req)
|
||||
elif forward_mode.is_extend():
|
||||
extend_seq_lens = forward_batch.extend_seq_lens
|
||||
if extend_seq_lens is None:
|
||||
raise RuntimeError(
|
||||
"_build_generalized_req_id_per_token: extend_seq_lens is None in extend mode"
|
||||
)
|
||||
lens = extend_seq_lens.to(torch.int64)
|
||||
result = torch.repeat_interleave(generalized_req_ids_per_row, lens)
|
||||
else:
|
||||
result = generalized_req_ids_per_row
|
||||
|
||||
if int(result.shape[0]) != num_tokens:
|
||||
raise RuntimeError(
|
||||
f"fill_expected_inputs: sum(lens)={int(result.shape[0])} != num_tokens={num_tokens}"
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
def _expand_uniform(values: torch.Tensor, per_row: int) -> torch.Tensor:
|
||||
bs = int(values.shape[0])
|
||||
return values.unsqueeze(1).expand(bs, per_row).reshape(bs * per_row)
|
||||
|
||||
|
||||
def select_generalized_req_ids(
|
||||
*,
|
||||
vanilla_req_ids: torch.Tensor,
|
||||
bootstrap_room_ids_int: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
if bootstrap_room_ids_int is None:
|
||||
return vanilla_req_ids
|
||||
|
||||
bootstrap_room_ids_int = bootstrap_room_ids_int.to(
|
||||
device=vanilla_req_ids.device,
|
||||
dtype=torch.int64,
|
||||
)
|
||||
return torch.where(
|
||||
bootstrap_room_ids_int >= 0,
|
||||
bootstrap_room_ids_int,
|
||||
vanilla_req_ids.to(torch.int64),
|
||||
)
|
||||
@@ -0,0 +1,59 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.kv_canary.perturb.next_token_swap import maybe_perturb_swap_next_tokens
|
||||
from sglang.srt.kv_canary.token_oracle.oracle import TokenOracle
|
||||
from sglang.srt.kv_canary.token_oracle.oracle_manager import (
|
||||
TokenOracleManager,
|
||||
select_generalized_req_ids,
|
||||
)
|
||||
from sglang.srt.layers.sampler import Sampler, register_sampler_backend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
|
||||
|
||||
def install_oracle_sampler(*, oracle: TokenOracle) -> TokenOracleManager:
|
||||
manager = TokenOracleManager(oracle=oracle)
|
||||
register_sampler_backend(
|
||||
"token_oracle",
|
||||
lambda: _OracleSampler(token_oracle_manager=manager),
|
||||
)
|
||||
return manager
|
||||
|
||||
|
||||
class _OracleSampler(Sampler):
|
||||
def __init__(self, *, token_oracle_manager: TokenOracleManager) -> None:
|
||||
super().__init__()
|
||||
self._token_oracle_manager = token_oracle_manager
|
||||
|
||||
def forward(
|
||||
self,
|
||||
logits_output: "LogitsProcessorOutput",
|
||||
sampling_info: "SamplingBatchInfo",
|
||||
return_logprob: bool,
|
||||
top_logprobs_nums: List[int],
|
||||
token_ids_logprobs: List[List[int]],
|
||||
positions: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
vanilla_req_ids = sampling_info.rids_int
|
||||
if vanilla_req_ids is None:
|
||||
raise RuntimeError(
|
||||
"_OracleSampler.forward: generalized_req_id source tensor is None; "
|
||||
"token oracle requires a per-forward generalized_req_id source tensor "
|
||||
"(set in ForwardBatch.init_new when SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE=1)"
|
||||
)
|
||||
batch_next_token_ids = self._token_oracle_manager.sample_next_tokens(
|
||||
generalized_req_ids=select_generalized_req_ids(
|
||||
vanilla_req_ids=vanilla_req_ids,
|
||||
bootstrap_room_ids_int=sampling_info.bootstrap_room_ids_int,
|
||||
),
|
||||
logits_positions=positions,
|
||||
)
|
||||
|
||||
batch_next_token_ids = maybe_perturb_swap_next_tokens(batch_next_token_ids)
|
||||
return batch_next_token_ids
|
||||
@@ -107,6 +107,7 @@ from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
|
||||
from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner
|
||||
from sglang.srt.kv_canary.api import install_canary
|
||||
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
||||
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.attention.attention_registry import (
|
||||
ATTENTION_BACKENDS,
|
||||
@@ -650,6 +651,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
|
||||
if self.server_args.elastic_ep_backend:
|
||||
ElasticEPStateManager.init(self.server_args)
|
||||
self._token_oracle_manager = install_token_oracle_from_env(
|
||||
server_args=server_args,
|
||||
vocab_size=self.model_config.vocab_size,
|
||||
)
|
||||
# Load the model
|
||||
self.sampler = create_sampler()
|
||||
self.load_model()
|
||||
@@ -759,6 +764,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
self.canary_manager = install_canary(
|
||||
server_args=server_args,
|
||||
model_runner=self,
|
||||
token_oracle_manager=self._token_oracle_manager,
|
||||
)
|
||||
|
||||
# Init ngram embedding token table
|
||||
|
||||
@@ -100,6 +100,8 @@ LLAMA4_MODEL_ARCHS = (
|
||||
)
|
||||
|
||||
SAMPLING_BACKEND_CHOICES = {"flashinfer", "pytorch", "ascend"}
|
||||
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
|
||||
SAMPLING_BACKEND_CHOICES.add("token_oracle")
|
||||
|
||||
LOAD_FORMAT_CHOICES = [
|
||||
"auto",
|
||||
|
||||
@@ -136,6 +136,7 @@ def make_base_config() -> CanaryConfig:
|
||||
mode=CanaryMode.RAISE,
|
||||
ring_capacity=1024,
|
||||
sweep_interval=0,
|
||||
enable_write_input_assert=False,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -25,11 +25,13 @@ def make_config(
|
||||
mode: CanaryMode = CanaryMode.RAISE,
|
||||
ring_capacity: int = 1024,
|
||||
sweep_interval: int = 0,
|
||||
enable_write_input_assert: bool = False,
|
||||
) -> CanaryConfig:
|
||||
return CanaryConfig(
|
||||
mode=mode,
|
||||
ring_capacity=ring_capacity,
|
||||
sweep_interval=sweep_interval,
|
||||
enable_write_input_assert=enable_write_input_assert,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ _MOCK_MODEL_SERVER_ARGS_NO_CANARY: list[str] = [
|
||||
"--load-format",
|
||||
"dummy",
|
||||
"--sampling-backend",
|
||||
"pytorch",
|
||||
"token_oracle",
|
||||
"--disable-piecewise-cuda-graph",
|
||||
]
|
||||
|
||||
@@ -51,7 +51,12 @@ def mock_model_server_args(*extra_args: str, canary_mode: str = "raise") -> list
|
||||
|
||||
def mock_model_server_env(*, input_check_enabled: bool = True) -> dict[str, str]:
|
||||
"""Return env overrides for popen_launch_server in mock-model + canary mode."""
|
||||
return {}
|
||||
return {
|
||||
"SGLANG_KV_CANARY_ENABLE_WRITE_INPUT_ASSERT": (
|
||||
"1" if input_check_enabled else "0"
|
||||
),
|
||||
"SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE": "1",
|
||||
}
|
||||
|
||||
|
||||
def run_mock_model_bench_serving(
|
||||
|
||||
Reference in New Issue
Block a user