Introduce NgramEmbeddingManager component (#31154)

This commit is contained in:
fzyzcjy
2026-07-14 15:58:08 +08:00
committed by GitHub
parent 0f20f52e5e
commit d15f6a9ac3
11 changed files with 242 additions and 166 deletions
@@ -37,6 +37,7 @@ def make_runner(
return SimpleNamespace(
server_args=args,
model_config=SimpleNamespace(),
hybrid_gdn_config=SimpleNamespace(
linear_key_head_dim=key_dim,
linear_value_head_dim=value_dim,
@@ -60,6 +61,11 @@ class TestFlashInferGDNPrefillBackendPolicy(unittest.TestCase):
flashinfer_available=True,
):
with (
patch.object(
gdn_backend,
"hybrid_gdn_config",
return_value=runner.hybrid_gdn_config,
),
patch.object(gdn_backend, "is_cuda", return_value=cuda),
patch.object(torch.cuda, "get_device_capability", return_value=capability),
patch.object(torch.version, "cuda", cuda_version),
@@ -97,7 +97,10 @@ def _scheduler_for_get_next_batch(*, tree_cache, chunked_req) -> Scheduler:
s.dp_attn_adapter.maybe_prepare_mlp_sync_batch = MagicMock(
side_effect=lambda batch, **_: batch
)
s._maybe_prepare_ngram_embedding = MagicMock(side_effect=lambda batch: batch)
s.ngram_embedding_manager = MagicMock()
s.ngram_embedding_manager.prepare_for_forward = MagicMock(
side_effect=lambda batch, **_: batch
)
s.update_running_batch = MagicMock(side_effect=lambda batch: batch)
s.tree_cache = tree_cache
s.chunked_req = chunked_req
@@ -11,7 +11,7 @@ from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.model_executor.ngram_token_table import ( # noqa: E402
from sglang.srt.model_executor.model_runner_components.ngram_embedding_manager import ( # noqa: E402
update_ngram_token_table_after_sampling,
)
@@ -37,7 +37,7 @@ class TestNgramTokenTableUpdate(CustomTestCase):
seq_lens = torch.tensor([11, 22, 33, 44], dtype=torch.int64)
with patch(
"sglang.srt.model_executor.ngram_token_table.update_token_table"
"sglang.srt.model_executor.model_runner_components.ngram_embedding_manager.update_token_table"
) as update_mock:
updated = update_ngram_token_table_after_sampling(
ngram_embedding_info=info,
@@ -71,7 +71,7 @@ class TestNgramTokenTableUpdate(CustomTestCase):
info = _make_ngram_info(2, skip_token_table_update=torch.tensor([True, True]))
with patch(
"sglang.srt.model_executor.ngram_token_table.update_token_table"
"sglang.srt.model_executor.model_runner_components.ngram_embedding_manager.update_token_table"
) as update_mock:
updated = update_ngram_token_table_after_sampling(
ngram_embedding_info=info,
@@ -91,7 +91,7 @@ class TestNgramTokenTableUpdate(CustomTestCase):
seq_lens = torch.tensor([11, 22, 33], dtype=torch.int64)
with patch(
"sglang.srt.model_executor.ngram_token_table.update_token_table"
"sglang.srt.model_executor.model_runner_components.ngram_embedding_manager.update_token_table"
) as update_mock:
updated = update_ngram_token_table_after_sampling(
ngram_embedding_info=info,