Introduce NgramEmbeddingManager component (#31154)
This commit is contained in:
@@ -199,6 +199,8 @@ class MlxModelRunnerStub(ModelRunner):
|
||||
self.graph_mem_usage = 0
|
||||
self.attn_backend = None
|
||||
|
||||
self.init_ngram_embedding_manager()
|
||||
|
||||
logger.info(
|
||||
f"MLX stub: initialized minimal pools "
|
||||
f"(max_total_num_tokens={self.max_total_num_tokens}, "
|
||||
|
||||
@@ -38,7 +38,6 @@ import torch.distributed
|
||||
from torch.cuda import Stream as CudaStream
|
||||
from torch.distributed import barrier
|
||||
|
||||
from sglang.jit_kernel.ngram_embedding import update_token_table
|
||||
from sglang.srt.configs.model_config import ModelConfig, ModelImpl, is_minimax_sparse
|
||||
from sglang.srt.constrained.grammar_manager import GrammarManager
|
||||
from sglang.srt.debug_utils.pr_fix_toggle import maybe_revert_pr_fix
|
||||
@@ -225,7 +224,7 @@ from sglang.srt.managers.utils import (
|
||||
)
|
||||
from sglang.srt.mem_cache import kv_cache_builder
|
||||
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode, PPProxyTensors
|
||||
from sglang.srt.model_executor.forward_batch_info import PPProxyTensors
|
||||
from sglang.srt.model_loader.utils import get_resolved_model_impl
|
||||
from sglang.srt.multiplex.multiplexing_mixin import SchedulerMultiplexMixin
|
||||
from sglang.srt.observability.metrics_collector import SchedulerMetricsCollector
|
||||
@@ -1303,71 +1302,16 @@ class Scheduler(
|
||||
self.batch_record_ct = 0
|
||||
|
||||
def maybe_init_ngram_embedding(self):
|
||||
self.ngram_embedding_manager = (
|
||||
self.tp_worker.model_runner.ngram_embedding_manager
|
||||
)
|
||||
self.use_ngram_embedding = self.tp_worker.model_config.use_ngram_embedding
|
||||
if self.use_ngram_embedding:
|
||||
self.token_table = self.tp_worker.model_runner.token_table
|
||||
self.token_table = self.tp_worker.model_runner.ngram_embedding_manager.table
|
||||
hf_config = self.tp_worker.model_config.hf_config
|
||||
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
||||
self.ngram_embedding_k = hf_config.ngram_embedding_k
|
||||
|
||||
def _maybe_prepare_ngram_embedding(
|
||||
self, batch: Optional[ScheduleBatch]
|
||||
) -> Optional[ScheduleBatch]:
|
||||
"""Fill the token table for ngram embedding before a forward pass."""
|
||||
if batch is None or not self.use_ngram_embedding:
|
||||
return batch
|
||||
batch.ne_token_table = self.token_table
|
||||
if batch.forward_mode == ForwardMode.EXTEND:
|
||||
all_tokens = []
|
||||
column_starts = []
|
||||
request_lengths = []
|
||||
for req in batch.reqs:
|
||||
start = len(req.prefix_indices)
|
||||
end = start + req.extend_range.length
|
||||
fill_ids = req.origin_input_ids + req.output_ids
|
||||
if start == 0:
|
||||
tokens = fill_ids[start:end]
|
||||
column_starts.append(0)
|
||||
elif start < self.ngram_embedding_n:
|
||||
tokens = fill_ids[0:end]
|
||||
column_starts.append(0)
|
||||
else:
|
||||
# Prepend n-1 tokens before prefix_len for n-gram context
|
||||
tokens = fill_ids[start - self.ngram_embedding_n + 1 : end]
|
||||
column_starts.append(start - self.ngram_embedding_n + 1)
|
||||
all_tokens.extend(tokens)
|
||||
request_lengths.append(len(tokens))
|
||||
dtype = self.token_table.dtype
|
||||
device = self.token_table.device
|
||||
update_token_table(
|
||||
ne_token_table=self.token_table,
|
||||
tokens=torch.tensor(all_tokens, dtype=dtype, device=device),
|
||||
row_indices=batch.req_pool_indices,
|
||||
column_starts=torch.tensor(
|
||||
column_starts, dtype=torch.int32, device=device
|
||||
),
|
||||
req_lens=torch.tensor(
|
||||
request_lengths, dtype=torch.int32, device=device
|
||||
),
|
||||
ignore_tokens=None,
|
||||
)
|
||||
# Mark the chunked (not-yet-finished) prefill request so sample()
|
||||
# skips writing its pseudo next-token into the ngram token table.
|
||||
# Use self.chunked_req identity (not req.is_chunked) to avoid
|
||||
# overlap-scheduling timing issues.
|
||||
if self.chunked_req is not None:
|
||||
skip_token_table_update = [
|
||||
req is self.chunked_req for req in batch.reqs
|
||||
]
|
||||
batch.ne_skip_token_table_update = (
|
||||
torch.tensor(
|
||||
skip_token_table_update, dtype=torch.bool, device=device
|
||||
)
|
||||
if any(skip_token_table_update)
|
||||
else None
|
||||
)
|
||||
return batch
|
||||
|
||||
def init_deterministic_inference_config(self):
|
||||
"""Initialize deterministic inference configuration for different attention backends."""
|
||||
if not self.server_args.enable_deterministic_inference:
|
||||
@@ -2806,7 +2750,9 @@ class Scheduler(
|
||||
)
|
||||
|
||||
# Handle ngram embedding
|
||||
ret = self._maybe_prepare_ngram_embedding(ret)
|
||||
ret = self.ngram_embedding_manager.prepare_for_forward(
|
||||
ret, chunked_req=self.chunked_req
|
||||
)
|
||||
|
||||
if ret:
|
||||
set_schedule_time_batch(ret)
|
||||
|
||||
@@ -848,7 +848,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
ret.positions = positions
|
||||
ret.extend_logprob_start_lens_cpu = extend_logprob_start_lens
|
||||
|
||||
if model_runner.use_ngram_embedding:
|
||||
if model_runner.ngram_embedding_manager.enabled:
|
||||
ret._init_ngram_embedding_info(batch, device)
|
||||
|
||||
if model_runner.model_config.model_is_mrope:
|
||||
|
||||
@@ -129,6 +129,9 @@ from sglang.srt.model_executor.forward_context import (
|
||||
)
|
||||
from sglang.srt.model_executor.graph_shared_output import GraphSharedOutput
|
||||
from sglang.srt.model_executor.hook_manager import register_forward_hooks
|
||||
from sglang.srt.model_executor.model_runner_components.ngram_embedding_manager import (
|
||||
NgramEmbeddingManager,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.remote_instance_weight_transporter import (
|
||||
RemoteInstanceWeightTransporter,
|
||||
)
|
||||
@@ -141,9 +144,6 @@ from sglang.srt.model_executor.model_runner_components.weight_updater import (
|
||||
from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
|
||||
ModelRunnerKVCacheMixin,
|
||||
)
|
||||
from sglang.srt.model_executor.ngram_token_table import (
|
||||
update_ngram_token_table_after_sampling,
|
||||
)
|
||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||
from sglang.srt.model_executor.runner import (
|
||||
EagerRunner,
|
||||
@@ -558,6 +558,16 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
gpu_id=self.gpu_id,
|
||||
)
|
||||
|
||||
def init_ngram_embedding_manager(self):
|
||||
self.ngram_embedding_manager = NgramEmbeddingManager.from_model(
|
||||
model=self.model,
|
||||
model_config=self.model_config,
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
server_args=self.server_args,
|
||||
max_running_requests=self.max_running_requests,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def init_msprobe(self):
|
||||
# Init the msprobe
|
||||
try:
|
||||
@@ -773,7 +783,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
)
|
||||
|
||||
# Init ngram embedding token table
|
||||
self.maybe_init_ngram_embedding()
|
||||
self.init_ngram_embedding_manager()
|
||||
|
||||
if self.enable_hisparse:
|
||||
from sglang.srt.managers.hisparse_coordinator import HiSparseCoordinator
|
||||
@@ -1627,45 +1637,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
full_attention_backend = ATTENTION_BACKENDS[backend_str](self)
|
||||
return attn_backend_wrapper(self, full_attention_backend)
|
||||
|
||||
def maybe_init_ngram_embedding(self):
|
||||
self.use_ngram_embedding = self.model_config.use_ngram_embedding
|
||||
if self.use_ngram_embedding:
|
||||
from sglang.srt.layers.n_gram_embedding import NgramEmbedding
|
||||
|
||||
# Sized to mirror req_to_token (indexed by req_pool_idx).
|
||||
self.token_table = torch.empty(
|
||||
self.req_to_token_pool.req_to_token.shape[0],
|
||||
self.model_config.context_len,
|
||||
dtype=torch.int32,
|
||||
device=self.device,
|
||||
)
|
||||
chunked_prefill_size = self.server_args.chunked_prefill_size
|
||||
assert (
|
||||
chunked_prefill_size is not None and chunked_prefill_size > 0
|
||||
), "Ngram embedding requires chunked prefill to be enabled (chunked_prefill_size > 0)"
|
||||
for module in self.model.modules():
|
||||
if isinstance(module, NgramEmbedding):
|
||||
module.init_buffers(
|
||||
self.max_running_requests, chunked_prefill_size, self.device
|
||||
)
|
||||
|
||||
def maybe_update_ngram_token_table(
|
||||
self,
|
||||
next_token_ids: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
"""Update the ngram embedding token table after sampling."""
|
||||
ngram_embedding_info = forward_batch.ngram_embedding_info
|
||||
if ngram_embedding_info is None:
|
||||
return
|
||||
update_ngram_token_table_after_sampling(
|
||||
ngram_embedding_info=ngram_embedding_info,
|
||||
next_token_ids=next_token_ids,
|
||||
req_pool_indices=forward_batch.req_pool_indices,
|
||||
seq_lens=forward_batch.seq_lens,
|
||||
batch_size=forward_batch.batch_size,
|
||||
)
|
||||
|
||||
def init_decode_cuda_graph(self):
|
||||
"""Capture device graphs."""
|
||||
self.decode_cuda_graph_runner = None
|
||||
@@ -2316,7 +2287,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
else forward_batch.seq_lens - 1
|
||||
),
|
||||
)
|
||||
self.maybe_update_ngram_token_table(next_token_ids, forward_batch)
|
||||
self.ngram_embedding_manager.update_after_decode(
|
||||
next_token_ids=next_token_ids,
|
||||
forward_batch=forward_batch,
|
||||
)
|
||||
return next_token_ids
|
||||
|
||||
def compute_logprobs_only(
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
"""Utilities for updating LongCat ngram embedding token tables."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.ngram_embedding import update_token_table
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.managers.schedule_batch import ForwardMode
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||
class NgramEmbeddingManager:
|
||||
enabled: bool
|
||||
table: Optional[torch.Tensor]
|
||||
n: int
|
||||
k: int
|
||||
|
||||
@classmethod
|
||||
def from_model(
|
||||
cls,
|
||||
*,
|
||||
model: torch.nn.Module,
|
||||
model_config: ModelConfig,
|
||||
req_to_token_pool: ReqToTokenPool,
|
||||
server_args: ServerArgs,
|
||||
max_running_requests: int,
|
||||
device: str,
|
||||
):
|
||||
token_table = None
|
||||
ngram_embedding_n = 0
|
||||
ngram_embedding_k = 0
|
||||
use_ngram_embedding = model_config.use_ngram_embedding
|
||||
if use_ngram_embedding:
|
||||
from sglang.srt.layers.n_gram_embedding import NgramEmbedding
|
||||
|
||||
# Sized to mirror req_to_token (indexed by req_pool_idx).
|
||||
token_table = torch.empty(
|
||||
req_to_token_pool.req_to_token.shape[0],
|
||||
model_config.context_len,
|
||||
dtype=torch.int32,
|
||||
device=device,
|
||||
)
|
||||
chunked_prefill_size = server_args.chunked_prefill_size
|
||||
assert (
|
||||
chunked_prefill_size is not None and chunked_prefill_size > 0
|
||||
), "Ngram embedding requires chunked prefill to be enabled (chunked_prefill_size > 0)"
|
||||
for module in model.modules():
|
||||
if isinstance(module, NgramEmbedding):
|
||||
module.init_buffers(
|
||||
max_running_requests, chunked_prefill_size, device
|
||||
)
|
||||
hf_config = model_config.hf_config
|
||||
ngram_embedding_n = hf_config.ngram_embedding_n
|
||||
ngram_embedding_k = hf_config.ngram_embedding_k
|
||||
return cls(
|
||||
enabled=use_ngram_embedding,
|
||||
table=token_table,
|
||||
n=ngram_embedding_n,
|
||||
k=ngram_embedding_k,
|
||||
)
|
||||
|
||||
def update_after_decode(
|
||||
self,
|
||||
next_token_ids: torch.Tensor,
|
||||
forward_batch: ForwardBatch,
|
||||
):
|
||||
"""Update the ngram embedding token table after sampling."""
|
||||
ngram_embedding_info = forward_batch.ngram_embedding_info
|
||||
if ngram_embedding_info is None:
|
||||
return
|
||||
update_ngram_token_table_after_sampling(
|
||||
ngram_embedding_info=ngram_embedding_info,
|
||||
next_token_ids=next_token_ids,
|
||||
req_pool_indices=forward_batch.req_pool_indices,
|
||||
seq_lens=forward_batch.seq_lens,
|
||||
batch_size=forward_batch.batch_size,
|
||||
)
|
||||
|
||||
def prepare_for_forward(
|
||||
self,
|
||||
batch: Optional[ScheduleBatch],
|
||||
*,
|
||||
chunked_req: Optional[Req],
|
||||
) -> Optional[ScheduleBatch]:
|
||||
"""Fill the token table for ngram embedding before a forward pass."""
|
||||
if batch is None or not self.enabled:
|
||||
return batch
|
||||
batch.ne_token_table = self.table
|
||||
if batch.forward_mode == ForwardMode.EXTEND:
|
||||
all_tokens = []
|
||||
column_starts = []
|
||||
request_lengths = []
|
||||
for req in batch.reqs:
|
||||
start = len(req.prefix_indices)
|
||||
end = start + req.extend_range.length
|
||||
fill_ids = req.origin_input_ids + req.output_ids
|
||||
if start == 0:
|
||||
tokens = fill_ids[start:end]
|
||||
column_starts.append(0)
|
||||
elif start < self.n:
|
||||
tokens = fill_ids[0:end]
|
||||
column_starts.append(0)
|
||||
else:
|
||||
# Prepend n-1 tokens before prefix_len for n-gram context
|
||||
tokens = fill_ids[start - self.n + 1 : end]
|
||||
column_starts.append(start - self.n + 1)
|
||||
all_tokens.extend(tokens)
|
||||
request_lengths.append(len(tokens))
|
||||
dtype = self.table.dtype
|
||||
device = self.table.device
|
||||
update_token_table(
|
||||
ne_token_table=self.table,
|
||||
tokens=torch.tensor(all_tokens, dtype=dtype, device=device),
|
||||
row_indices=batch.req_pool_indices,
|
||||
column_starts=torch.tensor(
|
||||
column_starts, dtype=torch.int32, device=device
|
||||
),
|
||||
req_lens=torch.tensor(
|
||||
request_lengths, dtype=torch.int32, device=device
|
||||
),
|
||||
ignore_tokens=None,
|
||||
)
|
||||
# Mark the chunked (not-yet-finished) prefill request so sample()
|
||||
# skips writing its pseudo next-token into the ngram token table.
|
||||
# Use self.chunked_req identity (not req.is_chunked) to avoid
|
||||
# overlap-scheduling timing issues.
|
||||
if chunked_req is not None:
|
||||
skip_token_table_update = [req is chunked_req for req in batch.reqs]
|
||||
batch.ne_skip_token_table_update = (
|
||||
torch.tensor(
|
||||
skip_token_table_update, dtype=torch.bool, device=device
|
||||
)
|
||||
if any(skip_token_table_update)
|
||||
else None
|
||||
)
|
||||
return batch
|
||||
|
||||
|
||||
def update_ngram_token_table_after_sampling(
|
||||
*,
|
||||
ngram_embedding_info,
|
||||
next_token_ids: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
batch_size: int,
|
||||
) -> bool:
|
||||
"""Update the ngram token table with sampled tokens.
|
||||
|
||||
Returns whether the token table was updated.
|
||||
"""
|
||||
skip_token_table_update = ngram_embedding_info.skip_token_table_update
|
||||
if skip_token_table_update is not None:
|
||||
# Skip chunked (not-yet-finished) prefill requests: their sampled token
|
||||
# is a pseudo prediction and must not pollute the token table.
|
||||
indices = (~skip_token_table_update).nonzero(as_tuple=True)[0]
|
||||
if indices.numel() == 0:
|
||||
return False
|
||||
update_token_table(
|
||||
ne_token_table=ngram_embedding_info.token_table,
|
||||
tokens=next_token_ids[indices].to(torch.int32),
|
||||
row_indices=req_pool_indices[indices],
|
||||
column_starts=seq_lens[indices].to(torch.int32),
|
||||
req_lens=torch.ones(
|
||||
indices.numel(), dtype=torch.int32, device=next_token_ids.device
|
||||
),
|
||||
ignore_tokens=None,
|
||||
)
|
||||
return True
|
||||
|
||||
ngram_embedding_info.out_column_starts[:batch_size] = seq_lens
|
||||
ngram_embedding_info.out_req_lens[:batch_size] = 1
|
||||
update_token_table(
|
||||
ne_token_table=ngram_embedding_info.token_table,
|
||||
tokens=next_token_ids.to(torch.int32),
|
||||
row_indices=req_pool_indices,
|
||||
column_starts=ngram_embedding_info.out_column_starts,
|
||||
req_lens=ngram_embedding_info.out_req_lens,
|
||||
ignore_tokens=None,
|
||||
)
|
||||
return True
|
||||
@@ -1,51 +0,0 @@
|
||||
"""Utilities for updating LongCat ngram embedding token tables."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.jit_kernel.ngram_embedding import update_token_table
|
||||
|
||||
|
||||
def update_ngram_token_table_after_sampling(
|
||||
*,
|
||||
ngram_embedding_info,
|
||||
next_token_ids: torch.Tensor,
|
||||
req_pool_indices: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
batch_size: int,
|
||||
) -> bool:
|
||||
"""Update the ngram token table with sampled tokens.
|
||||
|
||||
Returns whether the token table was updated.
|
||||
"""
|
||||
skip_token_table_update = ngram_embedding_info.skip_token_table_update
|
||||
if skip_token_table_update is not None:
|
||||
# Skip chunked (not-yet-finished) prefill requests: their sampled token
|
||||
# is a pseudo prediction and must not pollute the token table.
|
||||
indices = (~skip_token_table_update).nonzero(as_tuple=True)[0]
|
||||
if indices.numel() == 0:
|
||||
return False
|
||||
update_token_table(
|
||||
ne_token_table=ngram_embedding_info.token_table,
|
||||
tokens=next_token_ids[indices].to(torch.int32),
|
||||
row_indices=req_pool_indices[indices],
|
||||
column_starts=seq_lens[indices].to(torch.int32),
|
||||
req_lens=torch.ones(
|
||||
indices.numel(), dtype=torch.int32, device=next_token_ids.device
|
||||
),
|
||||
ignore_tokens=None,
|
||||
)
|
||||
return True
|
||||
|
||||
ngram_embedding_info.out_column_starts[:batch_size] = seq_lens
|
||||
ngram_embedding_info.out_req_lens[:batch_size] = 1
|
||||
update_token_table(
|
||||
ne_token_table=ngram_embedding_info.token_table,
|
||||
tokens=next_token_ids.to(torch.int32),
|
||||
row_indices=req_pool_indices,
|
||||
column_starts=ngram_embedding_info.out_column_starts,
|
||||
req_lens=ngram_embedding_info.out_req_lens,
|
||||
ignore_tokens=None,
|
||||
)
|
||||
return True
|
||||
@@ -312,7 +312,11 @@ class BaseRunner(ABC):
|
||||
num_tokens_per_req=num_tokens_per_req,
|
||||
cache_loc_dtype=torch.int64,
|
||||
enable_mamba_track=False,
|
||||
ne_token_table=mr.token_table if mr.use_ngram_embedding else None,
|
||||
ne_token_table=(
|
||||
mr.ngram_embedding_manager.table
|
||||
if mr.ngram_embedding_manager.enabled
|
||||
else None
|
||||
),
|
||||
hc_hidden_size=getattr(mr.model_config, "hc_hidden_size", None),
|
||||
pp_proxy_topk_size=mr.get_pp_proxy_topk_size(),
|
||||
)
|
||||
|
||||
@@ -213,7 +213,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.enable_two_batch_overlap = (
|
||||
model_runner.server_args.enable_two_batch_overlap
|
||||
)
|
||||
self.use_ngram_embedding = model_runner.use_ngram_embedding
|
||||
self.use_ngram_embedding = model_runner.ngram_embedding_manager.enabled
|
||||
if self.use_ngram_embedding:
|
||||
hf_config = model_runner.model_config.hf_config
|
||||
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
||||
@@ -362,7 +362,9 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
cache_loc_dtype=self._cache_loc_dtype(),
|
||||
enable_mamba_track=enable_mamba_track,
|
||||
ne_token_table=(
|
||||
model_runner.token_table if self.use_ngram_embedding else None
|
||||
model_runner.ngram_embedding_manager.table
|
||||
if self.use_ngram_embedding
|
||||
else None
|
||||
),
|
||||
hc_hidden_size=getattr(
|
||||
self.model_runner.model_config, "hc_hidden_size", None
|
||||
|
||||
Reference in New Issue
Block a user