Introduce NgramEmbeddingManager component (#31154)
This commit is contained in:
@@ -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
|
||||
|
||||
+4
-4
@@ -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,
|
||||
Reference in New Issue
Block a user