Add rids/bootstrap-room int-hash plumbing for deterministic per-request identification (#26814)
This commit is contained in:
@@ -753,6 +753,7 @@ class Envs:
|
|||||||
# KV-Canary (testing-only)
|
# KV-Canary (testing-only)
|
||||||
# ===================================================================
|
# ===================================================================
|
||||||
SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024)
|
SGLANG_KV_CANARY_RING_CAPACITY = EnvInt(1024)
|
||||||
|
SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE = EnvBool(False)
|
||||||
SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False)
|
SGLANG_KV_CANARY_ENABLE_MHA_V = EnvBool(False)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -149,6 +149,8 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
encoder_lens: Optional[torch.Tensor]
|
encoder_lens: Optional[torch.Tensor]
|
||||||
pp_proxy_tensors: Optional[Dict[str, torch.Tensor]]
|
pp_proxy_tensors: Optional[Dict[str, torch.Tensor]]
|
||||||
ngram_embedding_info: Optional["NgramEmbeddingInfo"]
|
ngram_embedding_info: Optional["NgramEmbeddingInfo"]
|
||||||
|
rids_int: Optional[torch.Tensor]
|
||||||
|
bootstrap_room_ids_int: Optional[torch.Tensor]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def create(
|
def create(
|
||||||
@@ -240,6 +242,13 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
|
||||||
|
rids_int = torch.zeros((max_bs,), dtype=torch.int64)
|
||||||
|
bootstrap_room_ids_int = torch.full((max_bs,), -1, dtype=torch.int64)
|
||||||
|
else:
|
||||||
|
rids_int = None
|
||||||
|
bootstrap_room_ids_int = None
|
||||||
|
|
||||||
# Keep seq_lens_cpu as a true CPU tensor, like the old implementation.
|
# Keep seq_lens_cpu as a true CPU tensor, like the old implementation.
|
||||||
seq_lens_cpu = torch.full(
|
seq_lens_cpu = torch.full(
|
||||||
(max_bs,),
|
(max_bs,),
|
||||||
@@ -267,6 +276,8 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||||
pp_proxy_tensors=pp_proxy_tensors,
|
pp_proxy_tensors=pp_proxy_tensors,
|
||||||
ngram_embedding_info=ngram_embedding_info,
|
ngram_embedding_info=ngram_embedding_info,
|
||||||
|
rids_int=rids_int,
|
||||||
|
bootstrap_room_ids_int=bootstrap_room_ids_int,
|
||||||
)
|
)
|
||||||
|
|
||||||
def populate_from_forward_batch(
|
def populate_from_forward_batch(
|
||||||
@@ -342,6 +353,16 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
dsts.append(self.mrope_positions[:, :raw_num_token])
|
dsts.append(self.mrope_positions[:, :raw_num_token])
|
||||||
srcs.append(forward_batch.mrope_positions)
|
srcs.append(forward_batch.mrope_positions)
|
||||||
|
|
||||||
|
if self.rids_int is not None and forward_batch.rids_int is not None:
|
||||||
|
dsts.append(self.rids_int[:raw_bs])
|
||||||
|
srcs.append(forward_batch.rids_int)
|
||||||
|
if (
|
||||||
|
self.bootstrap_room_ids_int is not None
|
||||||
|
and forward_batch.bootstrap_room_ids_int is not None
|
||||||
|
):
|
||||||
|
dsts.append(self.bootstrap_room_ids_int[:raw_bs])
|
||||||
|
srcs.append(forward_batch.bootstrap_room_ids_int)
|
||||||
|
|
||||||
if require_gathered_buffer:
|
if require_gathered_buffer:
|
||||||
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs)
|
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs)
|
||||||
self.global_num_tokens_for_logprob_gpu.fill_(bs * num_tokens_per_bs)
|
self.global_num_tokens_for_logprob_gpu.fill_(bs * num_tokens_per_bs)
|
||||||
@@ -935,6 +956,12 @@ class CudaGraphRunner:
|
|||||||
encoder_lens = None
|
encoder_lens = None
|
||||||
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
||||||
next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens]
|
next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens]
|
||||||
|
rids_int = buffers.rids_int[:bs] if buffers.rids_int is not None else None
|
||||||
|
bootstrap_room_ids_int = (
|
||||||
|
buffers.bootstrap_room_ids_int[:bs]
|
||||||
|
if buffers.bootstrap_room_ids_int is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
|
||||||
# Adjust for attention TP if needed (matching replay path in
|
# Adjust for attention TP if needed (matching replay path in
|
||||||
# populate_from_forward_batch).
|
# populate_from_forward_batch).
|
||||||
@@ -1050,6 +1077,8 @@ class CudaGraphRunner:
|
|||||||
num_token_non_padded=buffers.num_token_non_padded,
|
num_token_non_padded=buffers.num_token_non_padded,
|
||||||
global_forward_mode=self.capture_forward_mode,
|
global_forward_mode=self.capture_forward_mode,
|
||||||
lora_ids=lora_ids,
|
lora_ids=lora_ids,
|
||||||
|
rids_int=rids_int,
|
||||||
|
bootstrap_room_ids_int=bootstrap_room_ids_int,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Trip the coordinator so the hisparse code path is captured into the
|
# Trip the coordinator so the hisparse code path is captured into the
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ ScheduleBatch -> ForwardBatch
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import hashlib
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from enum import IntEnum, auto
|
from enum import IntEnum, auto
|
||||||
from functools import total_ordering
|
from functools import total_ordering
|
||||||
@@ -449,6 +450,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# For ngram embedding
|
# For ngram embedding
|
||||||
ngram_embedding_info: Optional[NgramEmbeddingInfo] = None
|
ngram_embedding_info: Optional[NgramEmbeddingInfo] = None
|
||||||
|
|
||||||
|
# For dumper: int-hashed request / bootstrap-room IDs (derived from rids)
|
||||||
|
rids_int: Optional[torch.Tensor] = None
|
||||||
|
bootstrap_room_ids_int: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# kv-canary token-id validator snapshot
|
# kv-canary token-id validator snapshot
|
||||||
req_all_ids_flat: Optional[torch.Tensor] = None
|
req_all_ids_flat: Optional[torch.Tensor] = None
|
||||||
req_all_ids_lens: Optional[torch.Tensor] = None
|
req_all_ids_lens: Optional[torch.Tensor] = None
|
||||||
@@ -568,6 +573,20 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
|
|
||||||
device = model_runner.device
|
device = model_runner.device
|
||||||
|
|
||||||
|
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
|
||||||
|
hashed = _hash_rids_to_tensor(
|
||||||
|
rids=[req.rid for req in batch.reqs],
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
bootstrap_room_ids = _bootstrap_rooms_to_tensor(
|
||||||
|
bootstrap_rooms=[req.bootstrap_room for req in batch.reqs],
|
||||||
|
device=device,
|
||||||
|
)
|
||||||
|
batch.sampling_info.rids_int = hashed
|
||||||
|
batch.sampling_info.bootstrap_room_ids_int = bootstrap_room_ids
|
||||||
|
ret.rids_int = hashed
|
||||||
|
ret.bootstrap_room_ids_int = bootstrap_room_ids
|
||||||
|
|
||||||
if batch.extend_input_logprob_token_ids is not None:
|
if batch.extend_input_logprob_token_ids is not None:
|
||||||
ret.extend_input_logprob_token_ids_gpu = (
|
ret.extend_input_logprob_token_ids_gpu = (
|
||||||
batch.extend_input_logprob_token_ids.to(device, non_blocking=True)
|
batch.extend_input_logprob_token_ids.to(device, non_blocking=True)
|
||||||
@@ -1086,6 +1105,17 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
if self.extend_seq_lens is not None:
|
if self.extend_seq_lens is not None:
|
||||||
self.extend_seq_lens = self._pad_tensor_to_size(self.extend_seq_lens, bs)
|
self.extend_seq_lens = self._pad_tensor_to_size(self.extend_seq_lens, bs)
|
||||||
|
|
||||||
|
if self.rids_int is not None:
|
||||||
|
self.rids_int = self._pad_tensor_to_size(self.rids_int, bs)
|
||||||
|
if self.sampling_info is not None:
|
||||||
|
self.sampling_info.rids_int = self.rids_int
|
||||||
|
if self.bootstrap_room_ids_int is not None:
|
||||||
|
self.bootstrap_room_ids_int = self._pad_tensor_to_size(
|
||||||
|
self.bootstrap_room_ids_int, bs, value=-1
|
||||||
|
)
|
||||||
|
if self.sampling_info is not None:
|
||||||
|
self.sampling_info.bootstrap_room_ids_int = self.bootstrap_room_ids_int
|
||||||
|
|
||||||
if self.spec_info is not None and self.spec_info.is_draft_input():
|
if self.spec_info is not None and self.spec_info.is_draft_input():
|
||||||
spec_info = self.spec_info
|
spec_info = self.spec_info
|
||||||
self.output_cache_loc_backup = self.out_cache_loc
|
self.output_cache_loc_backup = self.out_cache_loc
|
||||||
@@ -1321,3 +1351,18 @@ else:
|
|||||||
clamp_position = _clamp_position_native
|
clamp_position = _clamp_position_native
|
||||||
|
|
||||||
|
|
||||||
|
def _hash_rids_to_tensor(*, rids: List[str], device: torch.device) -> torch.Tensor:
|
||||||
|
values: List[int] = [_stable_hash_str_to_i64(rid) for rid in rids]
|
||||||
|
return torch.tensor(values, dtype=torch.int64, device=device)
|
||||||
|
|
||||||
|
|
||||||
|
def _bootstrap_rooms_to_tensor(
|
||||||
|
*, bootstrap_rooms: List[Optional[int]], device: torch.device
|
||||||
|
) -> torch.Tensor:
|
||||||
|
values: List[int] = [room if room is not None else -1 for room in bootstrap_rooms]
|
||||||
|
return torch.tensor(values, dtype=torch.int64, device=device)
|
||||||
|
|
||||||
|
|
||||||
|
def _stable_hash_str_to_i64(rid: str) -> int:
|
||||||
|
digest = hashlib.blake2b(rid.encode("utf-8"), digest_size=8).digest()
|
||||||
|
return int.from_bytes(digest, "little", signed=True)
|
||||||
|
|||||||
@@ -42,6 +42,8 @@ class SamplingBatchInfo:
|
|||||||
# Masking tensors for grammar-guided structured outputs
|
# Masking tensors for grammar-guided structured outputs
|
||||||
vocab_size: int
|
vocab_size: int
|
||||||
grammars: Optional[List] = None
|
grammars: Optional[List] = None
|
||||||
|
rids_int: Optional[torch.Tensor] = None
|
||||||
|
bootstrap_room_ids_int: Optional[torch.Tensor] = None
|
||||||
vocab_mask: Optional[torch.Tensor] = None
|
vocab_mask: Optional[torch.Tensor] = None
|
||||||
apply_mask_func: Optional[Callable[[torch.Tensor, torch.Tensor], None]] = None
|
apply_mask_func: Optional[Callable[[torch.Tensor, torch.Tensor], None]] = None
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from typing import TYPE_CHECKING, Callable, Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import DpPaddingMode, set_dp_buffer_len
|
from sglang.srt.layers.dp_attention import DpPaddingMode, set_dp_buffer_len
|
||||||
from sglang.srt.model_executor.cuda_graph_runner import (
|
from sglang.srt.model_executor.cuda_graph_runner import (
|
||||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||||
@@ -46,6 +47,8 @@ class EagleDraftInputBuffers(ForwardInputBuffers):
|
|||||||
out_cache_loc: torch.Tensor
|
out_cache_loc: torch.Tensor
|
||||||
positions: torch.Tensor
|
positions: torch.Tensor
|
||||||
mrope_positions: torch.Tensor
|
mrope_positions: torch.Tensor
|
||||||
|
rids_int: Optional[torch.Tensor]
|
||||||
|
bootstrap_room_ids_int: Optional[torch.Tensor]
|
||||||
seq_lens: torch.Tensor
|
seq_lens: torch.Tensor
|
||||||
seq_lens_cpu: torch.Tensor
|
seq_lens_cpu: torch.Tensor
|
||||||
extend_seq_lens: torch.Tensor
|
extend_seq_lens: torch.Tensor
|
||||||
@@ -124,6 +127,16 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
)
|
)
|
||||||
positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
positions = torch.zeros((self.max_num_token,), dtype=torch.int64)
|
||||||
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
mrope_positions = torch.zeros((3, self.max_num_token), dtype=torch.int64)
|
||||||
|
rids_int = (
|
||||||
|
torch.zeros((self.max_bs,), dtype=torch.int64)
|
||||||
|
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get()
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
bootstrap_room_ids_int = (
|
||||||
|
torch.full((self.max_bs,), -1, dtype=torch.int64)
|
||||||
|
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get()
|
||||||
|
else None
|
||||||
|
)
|
||||||
seq_lens = torch.full(
|
seq_lens = torch.full(
|
||||||
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
|
(self.max_bs,), self.seq_len_fill_value, dtype=torch.int32
|
||||||
)
|
)
|
||||||
@@ -164,6 +177,8 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
out_cache_loc=out_cache_loc,
|
out_cache_loc=out_cache_loc,
|
||||||
positions=positions,
|
positions=positions,
|
||||||
mrope_positions=mrope_positions,
|
mrope_positions=mrope_positions,
|
||||||
|
rids_int=rids_int,
|
||||||
|
bootstrap_room_ids_int=bootstrap_room_ids_int,
|
||||||
seq_lens=seq_lens,
|
seq_lens=seq_lens,
|
||||||
seq_lens_cpu=seq_lens_cpu,
|
seq_lens_cpu=seq_lens_cpu,
|
||||||
extend_seq_lens=extend_seq_lens,
|
extend_seq_lens=extend_seq_lens,
|
||||||
@@ -259,6 +274,12 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
out_cache_loc = buffers.out_cache_loc[: num_tokens * self.speculative_num_steps]
|
out_cache_loc = buffers.out_cache_loc[: num_tokens * self.speculative_num_steps]
|
||||||
positions = buffers.positions[:num_tokens]
|
positions = buffers.positions[:num_tokens]
|
||||||
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
mrope_positions = buffers.mrope_positions[:, :num_tokens]
|
||||||
|
rids_int = buffers.rids_int[:num_seqs] if buffers.rids_int is not None else None
|
||||||
|
bootstrap_room_ids_int = (
|
||||||
|
buffers.bootstrap_room_ids_int[:num_seqs]
|
||||||
|
if buffers.bootstrap_room_ids_int is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
hidden_states = (
|
hidden_states = (
|
||||||
buffers.hidden_states[:num_seqs]
|
buffers.hidden_states[:num_seqs]
|
||||||
if buffers.hidden_states is not None
|
if buffers.hidden_states is not None
|
||||||
@@ -341,6 +362,8 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
global_dp_buffer_len=global_dp_buffer_len,
|
global_dp_buffer_len=global_dp_buffer_len,
|
||||||
spec_algorithm=self.model_runner.spec_algorithm,
|
spec_algorithm=self.model_runner.spec_algorithm,
|
||||||
spec_info=spec_info,
|
spec_info=spec_info,
|
||||||
|
rids_int=rids_int,
|
||||||
|
bootstrap_room_ids_int=bootstrap_room_ids_int,
|
||||||
capture_hidden_mode=(
|
capture_hidden_mode=(
|
||||||
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
|
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
|
||||||
),
|
),
|
||||||
@@ -412,6 +435,10 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
buffers.seq_lens.fill_(self.seq_len_fill_value)
|
||||||
buffers.out_cache_loc.zero_()
|
buffers.out_cache_loc.zero_()
|
||||||
buffers.positions.zero_()
|
buffers.positions.zero_()
|
||||||
|
if buffers.rids_int is not None:
|
||||||
|
buffers.rids_int.zero_()
|
||||||
|
if buffers.bootstrap_room_ids_int is not None:
|
||||||
|
buffers.bootstrap_room_ids_int.fill_(-1)
|
||||||
buffers.topk_p.zero_()
|
buffers.topk_p.zero_()
|
||||||
buffers.topk_index.zero_()
|
buffers.topk_index.zero_()
|
||||||
if buffers.hidden_states is not None:
|
if buffers.hidden_states is not None:
|
||||||
@@ -426,6 +453,15 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
forward_batch.out_cache_loc
|
forward_batch.out_cache_loc
|
||||||
)
|
)
|
||||||
buffers.positions[:raw_num_token].copy_(forward_batch.positions)
|
buffers.positions[:raw_num_token].copy_(forward_batch.positions)
|
||||||
|
if buffers.rids_int is not None and forward_batch.rids_int is not None:
|
||||||
|
buffers.rids_int[:raw_bs].copy_(forward_batch.rids_int)
|
||||||
|
if (
|
||||||
|
buffers.bootstrap_room_ids_int is not None
|
||||||
|
and forward_batch.bootstrap_room_ids_int is not None
|
||||||
|
):
|
||||||
|
buffers.bootstrap_room_ids_int[:raw_bs].copy_(
|
||||||
|
forward_batch.bootstrap_room_ids_int
|
||||||
|
)
|
||||||
maybe_detect_nan(
|
maybe_detect_nan(
|
||||||
forward_batch.spec_info.topk_p,
|
forward_batch.spec_info.topk_p,
|
||||||
"EagleDraftCudaGraphRunner.replay: topk_p",
|
"EagleDraftCudaGraphRunner.replay: topk_p",
|
||||||
@@ -457,6 +493,15 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
forward_batch.seq_lens = buffers.seq_lens[:bs]
|
forward_batch.seq_lens = buffers.seq_lens[:bs]
|
||||||
forward_batch.req_pool_indices = buffers.req_pool_indices[:bs]
|
forward_batch.req_pool_indices = buffers.req_pool_indices[:bs]
|
||||||
forward_batch.positions = buffers.positions[:num_tokens]
|
forward_batch.positions = buffers.positions[:num_tokens]
|
||||||
|
if buffers.rids_int is not None and forward_batch.rids_int is not None:
|
||||||
|
forward_batch.rids_int = buffers.rids_int[:bs]
|
||||||
|
if (
|
||||||
|
buffers.bootstrap_room_ids_int is not None
|
||||||
|
and forward_batch.bootstrap_room_ids_int is not None
|
||||||
|
):
|
||||||
|
forward_batch.bootstrap_room_ids_int = buffers.bootstrap_room_ids_int[
|
||||||
|
:bs
|
||||||
|
]
|
||||||
|
|
||||||
if forward_batch.seq_lens_cpu is not None:
|
if forward_batch.seq_lens_cpu is not None:
|
||||||
if bs != raw_bs:
|
if bs != raw_bs:
|
||||||
@@ -481,6 +526,15 @@ class EAGLEDraftCudaGraphRunner:
|
|||||||
forward_batch.positions = buffers.positions[:raw_num_token]
|
forward_batch.positions = buffers.positions[:raw_num_token]
|
||||||
forward_batch.seq_lens = buffers.seq_lens[:raw_bs]
|
forward_batch.seq_lens = buffers.seq_lens[:raw_bs]
|
||||||
forward_batch.req_pool_indices = buffers.req_pool_indices[:raw_bs]
|
forward_batch.req_pool_indices = buffers.req_pool_indices[:raw_bs]
|
||||||
|
if buffers.rids_int is not None and forward_batch.rids_int is not None:
|
||||||
|
forward_batch.rids_int = buffers.rids_int[:raw_bs]
|
||||||
|
if (
|
||||||
|
buffers.bootstrap_room_ids_int is not None
|
||||||
|
and forward_batch.bootstrap_room_ids_int is not None
|
||||||
|
):
|
||||||
|
forward_batch.bootstrap_room_ids_int = buffers.bootstrap_room_ids_int[
|
||||||
|
:raw_bs
|
||||||
|
]
|
||||||
if forward_batch.seq_lens_cpu is not None:
|
if forward_batch.seq_lens_cpu is not None:
|
||||||
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:raw_bs]
|
forward_batch.seq_lens_cpu = buffers.seq_lens_cpu[:raw_bs]
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user