[Model] Complete dots.note.omni support with native encoders, video preprocessing, and MTP decoding (#33829)
Co-authored-by: miraclezqc <dysania@pku.edu.cn>
This commit is contained in:
co-authored by
miraclezqc
parent
c35683fda0
commit
af39ad9349
@@ -8,11 +8,19 @@ cuda-graph buffer plumbing is covered by the backend SWA integration tests.
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||
|
||||
maybe_stub_sgl_kernel()
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.layers.attention.dots_hybrid_backend import DotsSWAMLAAttnBackend
|
||||
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionBackend
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc, MLATokenToKVPool
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
@@ -78,6 +86,78 @@ class TestSWAKVPoolSetKVBuffer(CustomTestCase):
|
||||
)
|
||||
self.assertIs(recorded["full_loc"], loc)
|
||||
|
||||
def test_composed_mla_pools_route_local_layer_ids(self):
|
||||
pool = object.__new__(SWAKVPool)
|
||||
pool.layers_mapping = {7: (1, False), 8: (2, True)}
|
||||
recorded = {}
|
||||
|
||||
def make_mla_pool(name):
|
||||
mla_pool = object.__new__(MLATokenToKVPool)
|
||||
|
||||
def set_kv(layer, loc, k, v, layer_id_override=None):
|
||||
recorded[f"{name}_kv"] = (layer, loc, layer_id_override)
|
||||
|
||||
def set_mla(layer, loc, k_nope, k_rope, layer_id_override=None):
|
||||
recorded[f"{name}_mla"] = (layer, loc, layer_id_override)
|
||||
|
||||
mla_pool.set_kv_buffer = set_kv
|
||||
mla_pool.set_mla_kv_buffer = set_mla
|
||||
return mla_pool
|
||||
|
||||
pool.full_kv_pool = make_mla_pool("full")
|
||||
pool.swa_kv_pool = make_mla_pool("swa")
|
||||
full_loc = torch.tensor([3, 4])
|
||||
swa_loc = torch.tensor([7, 8])
|
||||
|
||||
pool.set_kv_buffer(
|
||||
SimpleNamespace(layer_id=7), KVWriteLoc(full_loc, swa_loc), None, None
|
||||
)
|
||||
pool.set_mla_kv_buffer(
|
||||
SimpleNamespace(layer_id=8), KVWriteLoc(full_loc, swa_loc), None, None
|
||||
)
|
||||
|
||||
full_layer, recorded_full_loc, full_layer_id = recorded["full_kv"]
|
||||
swa_layer, recorded_swa_loc, swa_layer_id = recorded["swa_mla"]
|
||||
self.assertIsNone(full_layer)
|
||||
self.assertIs(recorded_full_loc, full_loc)
|
||||
self.assertEqual(full_layer_id, 1)
|
||||
self.assertIsNone(swa_layer)
|
||||
self.assertIs(recorded_swa_loc, swa_loc)
|
||||
self.assertEqual(swa_layer_id, 2)
|
||||
|
||||
|
||||
class TestDotsDraftSWAOutCacheLoc(CustomTestCase):
|
||||
def test_metadata_sees_only_current_step_and_forward_batch_is_restored(self):
|
||||
backend = object.__new__(FlashAttentionBackend)
|
||||
backend.topk = 2
|
||||
backend.speculative_num_steps = 3
|
||||
backend.speculative_step_id = 1
|
||||
|
||||
seen = []
|
||||
backend.init_forward_metadata_out_graph = MagicMock(
|
||||
side_effect=lambda forward_batch, in_capture=False: seen.append(
|
||||
forward_batch.out_cache_loc.clone()
|
||||
)
|
||||
)
|
||||
|
||||
wrapper = object.__new__(DotsSWAMLAAttnBackend)
|
||||
wrapper.backend = backend
|
||||
wrapper._active_backend = backend
|
||||
wrapper._prefill_metadata = None
|
||||
|
||||
original = torch.arange(12)
|
||||
forward_batch = SimpleNamespace(
|
||||
batch_size=2,
|
||||
forward_mode=ForwardMode.DECODE,
|
||||
out_cache_loc=original,
|
||||
spec_info=object(),
|
||||
)
|
||||
|
||||
wrapper.init_forward_metadata_out_graph(forward_batch)
|
||||
|
||||
torch.testing.assert_close(seen[0], torch.tensor([1, 4, 7, 10]))
|
||||
self.assertIs(forward_batch.out_cache_loc, original)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
import json
|
||||
import unittest
|
||||
|
||||
from sglang.srt.entrypoints.openai.protocol import Function, Tool
|
||||
from sglang.srt.function_call.dots_detector import DotsToolDetector
|
||||
from sglang.srt.function_call.function_call_parser import FunctionCallParser
|
||||
from sglang.srt.parser.reasoning_parser import Qwen3Detector, ReasoningParser
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _tool(name: str, properties: dict) -> Tool:
|
||||
return Tool(
|
||||
type="function",
|
||||
function=Function(
|
||||
name=name,
|
||||
description="test tool",
|
||||
parameters={"type": "object", "properties": properties},
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class TestDotsToolDetector(unittest.TestCase):
|
||||
def test_dots_parsers_are_registered(self):
|
||||
self.assertIs(ReasoningParser.DetectorMap["dots"], Qwen3Detector)
|
||||
self.assertIs(FunctionCallParser.ToolCallParserEnum["dots"], DotsToolDetector)
|
||||
|
||||
def test_dots_reasoning_uses_qwen3_format(self):
|
||||
parser = ReasoningParser("dots", stream_reasoning=False, force_reasoning=True)
|
||||
reasoning, content = parser.parse_non_stream(
|
||||
"Need to inspect inputs.</think>Final answer"
|
||||
)
|
||||
self.assertEqual(reasoning, "Need to inspect inputs.")
|
||||
self.assertEqual(content, "Final answer")
|
||||
|
||||
def test_non_stream_xml_converts_schema_types_and_resolves_ref(self):
|
||||
tool = Tool(
|
||||
type="function",
|
||||
function=Function(
|
||||
name="set_location",
|
||||
description="Set location",
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {"$ref": "#/$defs/Location"},
|
||||
"days": {"type": "integer"},
|
||||
"include_weather": {"type": "boolean"},
|
||||
},
|
||||
"$defs": {
|
||||
"Location": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
}
|
||||
},
|
||||
},
|
||||
),
|
||||
)
|
||||
parser = FunctionCallParser([tool], "dots")
|
||||
text = (
|
||||
"ok<dots_function_call>"
|
||||
'<invoke name="set_location">'
|
||||
'<parameter name="location">{"city": "Shanghai"}</parameter>'
|
||||
'<parameter name="days">3</parameter>'
|
||||
'<parameter name="include_weather">true</parameter>'
|
||||
"</invoke>"
|
||||
"</dots_function_call>"
|
||||
)
|
||||
|
||||
normal_text, calls = parser.parse_non_stream(text)
|
||||
|
||||
self.assertEqual(normal_text, "ok")
|
||||
self.assertEqual(len(calls), 1)
|
||||
self.assertEqual(calls[0].name, "set_location")
|
||||
self.assertEqual(
|
||||
json.loads(calls[0].parameters),
|
||||
{
|
||||
"location": {"city": "Shanghai"},
|
||||
"days": 3,
|
||||
"include_weather": True,
|
||||
},
|
||||
)
|
||||
|
||||
def test_non_stream_supports_multiple_invokes_and_json_fallback(self):
|
||||
tools = [
|
||||
_tool("search", {"query": {"type": "string"}}),
|
||||
_tool("open", {"id": {"type": "integer"}}),
|
||||
]
|
||||
parser = FunctionCallParser(tools, "dots")
|
||||
text = (
|
||||
"<dots_function_call>"
|
||||
'<invoke name="search"><parameter name="query">chairs</parameter></invoke>'
|
||||
'<invoke name="open"><parameter name="id">7</parameter></invoke>'
|
||||
"</dots_function_call>"
|
||||
'<dots_function_call>{"name":"search","arguments":{"query":"tables"}}'
|
||||
"</dots_function_call>"
|
||||
)
|
||||
|
||||
_, calls = parser.parse_non_stream(text)
|
||||
|
||||
self.assertEqual([call.name for call in calls], ["search", "open", "search"])
|
||||
self.assertEqual(
|
||||
[json.loads(call.parameters) for call in calls],
|
||||
[
|
||||
{"query": "chairs"},
|
||||
{"id": 7},
|
||||
{"query": "tables"},
|
||||
],
|
||||
)
|
||||
|
||||
def test_streaming_buffers_partial_marker_and_emits_all_complete_calls(self):
|
||||
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||
detector = DotsToolDetector()
|
||||
chunks = [
|
||||
"visible<dots_func",
|
||||
(
|
||||
"tion_call>"
|
||||
'<invoke name="search"><parameter name="query">chairs</parameter></invoke>'
|
||||
"</dots_function_call>"
|
||||
"<dots_function_call>"
|
||||
'<invoke name="search"><parameter name="query">tables</parameter></invoke>'
|
||||
"</dots_function_call>"
|
||||
),
|
||||
]
|
||||
|
||||
results = [detector.parse_streaming_increment(chunk, tools) for chunk in chunks]
|
||||
|
||||
self.assertEqual("".join(result.normal_text for result in results), "visible")
|
||||
calls = [call for result in results for call in result.calls]
|
||||
self.assertEqual([call.tool_index for call in calls], [0, 1])
|
||||
self.assertEqual(
|
||||
[json.loads(call.parameters) for call in calls],
|
||||
[
|
||||
{"query": "chairs"},
|
||||
{"query": "tables"},
|
||||
],
|
||||
)
|
||||
|
||||
def test_streaming_filters_unknown_tools_and_surfaces_the_content(self):
|
||||
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||
detector = DotsToolDetector()
|
||||
text = (
|
||||
"<dots_function_call>"
|
||||
'<invoke name="ghost"><parameter name="query">chairs</parameter></invoke>'
|
||||
"</dots_function_call>"
|
||||
)
|
||||
|
||||
result = detector.parse_streaming_increment(text, tools)
|
||||
|
||||
self.assertEqual(result.calls, [])
|
||||
self.assertIn("ghost", result.normal_text)
|
||||
self.assertEqual(detector._buffer, "")
|
||||
|
||||
def test_streaming_malformed_block_does_not_block_a_later_valid_call(self):
|
||||
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||
detector = DotsToolDetector()
|
||||
|
||||
malformed = detector.parse_streaming_increment(
|
||||
"<dots_function_call>garbage</dots_function_call>", tools
|
||||
)
|
||||
valid = detector.parse_streaming_increment(
|
||||
"<dots_function_call>"
|
||||
'<invoke name="search"><parameter name="query">chairs</parameter></invoke>'
|
||||
"</dots_function_call>",
|
||||
tools,
|
||||
)
|
||||
|
||||
self.assertEqual(malformed.calls, [])
|
||||
self.assertEqual(malformed.normal_text, "garbage")
|
||||
self.assertEqual([call.name for call in valid.calls], ["search"])
|
||||
|
||||
def test_streaming_strips_stray_end_marker_from_normal_text(self):
|
||||
detector = DotsToolDetector()
|
||||
|
||||
result = detector.parse_streaming_increment(
|
||||
"some text </dots_function_call>", []
|
||||
)
|
||||
|
||||
self.assertEqual(result.calls, [])
|
||||
self.assertEqual(result.normal_text, "some text ")
|
||||
|
||||
def test_streaming_flushes_partial_opening_marker_at_eof(self):
|
||||
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||
detector = DotsToolDetector()
|
||||
|
||||
result = detector.parse_streaming_increment("answer <dots_func", tools)
|
||||
|
||||
self.assertEqual(result.calls, [])
|
||||
self.assertEqual(result.normal_text, "answer ")
|
||||
self.assertEqual(detector.flush_pending_normal_text(), "<dots_func")
|
||||
self.assertEqual(detector.flush_pending_normal_text(), "")
|
||||
|
||||
def test_streaming_emits_complete_json_body_before_end_marker_without_duplication(
|
||||
self,
|
||||
):
|
||||
tools = [_tool("search", {"query": {"type": "string"}})]
|
||||
detector = DotsToolDetector()
|
||||
|
||||
opening = detector.parse_streaming_increment("<dots_function_call>", tools)
|
||||
body = detector.parse_streaming_increment(
|
||||
'{"name":"search","arguments":{"query":"chairs"}}', tools
|
||||
)
|
||||
closing = detector.parse_streaming_increment("</dots_function_call>", tools)
|
||||
|
||||
self.assertEqual(opening.calls, [])
|
||||
self.assertEqual([call.name for call in body.calls], ["search", None])
|
||||
self.assertEqual(
|
||||
"".join(call.parameters for call in body.calls), '{"query": "chairs"}'
|
||||
)
|
||||
self.assertEqual(closing.calls, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,135 @@
|
||||
import sys
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||
|
||||
maybe_stub_sgl_kernel()
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.attention.dots_hybrid_backend import (
|
||||
DotsHybridAttnBackend,
|
||||
DotsSWAMLAAttnBackend,
|
||||
_metadata_mismatches_dp_padded_batch,
|
||||
_normalize_cache_seqlens_rows,
|
||||
)
|
||||
from sglang.srt.layers.attention.flashattention_backend import FlashAttentionMetadata
|
||||
from sglang.srt.layers.attention.swa_mla_fallback.ops import (
|
||||
gather_page64_kv_latent,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
|
||||
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
|
||||
|
||||
|
||||
def _batch(*, bs: int, num_tokens: int, original_bs: int | None = None):
|
||||
return SimpleNamespace(
|
||||
batch_size=bs,
|
||||
out_cache_loc=torch.zeros(num_tokens, dtype=torch.int64),
|
||||
forward_mode=SimpleNamespace(),
|
||||
_original_batch_size=original_bs,
|
||||
)
|
||||
|
||||
|
||||
def _fa_metadata(*, bs: int, num_tokens: int):
|
||||
return FlashAttentionMetadata(
|
||||
page_table=torch.zeros((bs, 4), dtype=torch.int32),
|
||||
swa_page_table=torch.zeros((bs, 4), dtype=torch.int32),
|
||||
cache_seqlens_int32=torch.ones(bs, dtype=torch.int32),
|
||||
swa_out_cache_loc=torch.zeros(num_tokens, dtype=torch.int64),
|
||||
)
|
||||
|
||||
|
||||
def test_mismatch_detects_short_page_table_and_swa_loc():
|
||||
metadata = _fa_metadata(bs=1, num_tokens=1)
|
||||
assert _metadata_mismatches_dp_padded_batch(metadata, _batch(bs=2, num_tokens=2))
|
||||
|
||||
metadata = _fa_metadata(bs=2, num_tokens=2)
|
||||
assert not _metadata_mismatches_dp_padded_batch(
|
||||
metadata, _batch(bs=2, num_tokens=2)
|
||||
)
|
||||
|
||||
|
||||
def test_swa_backend_rebuilds_when_dp_padding_changes_rows():
|
||||
inner = SimpleNamespace(
|
||||
forward_metadata=_fa_metadata(bs=1, num_tokens=1),
|
||||
init_forward_metadata=MagicMock(),
|
||||
)
|
||||
backend = object.__new__(DotsSWAMLAAttnBackend)
|
||||
backend.backend = inner
|
||||
backend._active_backend = inner
|
||||
backend._prefill_metadata = None
|
||||
backend._dp_rebuilt_batch_id = None
|
||||
backend.init_forward_metadata = MagicMock()
|
||||
|
||||
backend.maybe_rebuild_metadata_after_dp_padding(_batch(bs=2, num_tokens=2))
|
||||
backend.init_forward_metadata.assert_called_once()
|
||||
|
||||
|
||||
def test_hybrid_rebuilds_when_dp_padding_changes_batch_size():
|
||||
matching = _fa_metadata(bs=2, num_tokens=2)
|
||||
hybrid = object.__new__(DotsHybridAttnBackend)
|
||||
hybrid.dsa_backend = SimpleNamespace(forward_metadata=matching)
|
||||
hybrid.swa_backend = SimpleNamespace(forward_metadata=matching)
|
||||
hybrid._dp_rebuilt_batch_id = None
|
||||
hybrid.init_forward_metadata = MagicMock()
|
||||
|
||||
hybrid.maybe_rebuild_metadata_after_dp_padding(
|
||||
_batch(bs=2, num_tokens=2, original_bs=1)
|
||||
)
|
||||
hybrid.init_forward_metadata.assert_called_once()
|
||||
|
||||
|
||||
def test_normalize_cache_seqlens_preserves_planned_rows_and_pads_dummy_rows():
|
||||
cache_seqlens = torch.tensor([17, 23], dtype=torch.int32)
|
||||
seq_lens = torch.tensor([100, 200, 300, 400], dtype=torch.int64)
|
||||
|
||||
normalized = _normalize_cache_seqlens_rows(cache_seqlens, seq_lens, 4)
|
||||
|
||||
assert torch.equal(normalized, torch.tensor([17, 23, 300, 400], dtype=torch.int32))
|
||||
|
||||
|
||||
def test_normalize_cache_seqlens_truncates_extra_rows():
|
||||
cache_seqlens = torch.tensor([17, 23, 29], dtype=torch.int32)
|
||||
seq_lens = torch.tensor([100, 200], dtype=torch.int64)
|
||||
|
||||
normalized = _normalize_cache_seqlens_rows(cache_seqlens, seq_lens, 2)
|
||||
|
||||
assert torch.equal(normalized, torch.tensor([17, 23], dtype=torch.int32))
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA")
|
||||
def test_page64_gather_masks_out_of_range_page_table_entries():
|
||||
kv_cache_dim = 128
|
||||
k_cache = torch.arange(64 * kv_cache_dim, dtype=torch.float32, device="cuda").view(
|
||||
64, 1, kv_cache_dim
|
||||
)
|
||||
# Row 0 has a sequence longer than its one-page table. Row 1 points past
|
||||
# the physical KV pool. Both can occur transiently when DP padding changes
|
||||
# the live batch after speculative metadata was planned.
|
||||
block_table = torch.tensor([[0], [9]], dtype=torch.int32, device="cuda")
|
||||
cache_seqlens = torch.tensor([130, 64], dtype=torch.int32, device="cuda")
|
||||
|
||||
gathered, valid = gather_page64_kv_latent(
|
||||
k_cache=k_cache,
|
||||
block_table=block_table,
|
||||
cache_seqlens=cache_seqlens,
|
||||
window_size=128,
|
||||
s_q=1,
|
||||
kv_cache_dim=kv_cache_dim,
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
assert valid[0, :62].all()
|
||||
assert not valid[0, 62:].any()
|
||||
assert not valid[1].any()
|
||||
torch.testing.assert_close(gathered[0, :62], k_cache[2:64, 0])
|
||||
assert not gathered[0, 62:].any()
|
||||
assert not gathered[1].any()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__]))
|
||||
@@ -170,6 +170,38 @@ class TestMlpSyncPadUnpad(CustomTestCase):
|
||||
# row count must match the real request count.
|
||||
self.assertEqual((fb.seq_lens - 1).shape[0], fb.batch_size)
|
||||
|
||||
def test_draft_extend_dummy_request_pads_cpu_and_gpu_lens(self):
|
||||
spec_info = MagicMock()
|
||||
spec_info.num_tokens_per_req = 4
|
||||
spec_info.is_draft_input.return_value = False
|
||||
fb = ForwardBatch(
|
||||
forward_mode=ForwardMode.DRAFT_EXTEND_V2,
|
||||
batch_size=1,
|
||||
input_ids=torch.empty(0, dtype=torch.int64),
|
||||
req_pool_indices=torch.empty(0, dtype=torch.int64),
|
||||
seq_lens=torch.empty(0, dtype=torch.int64),
|
||||
seq_lens_sum=0,
|
||||
out_cache_loc=torch.empty(0, dtype=torch.int64),
|
||||
positions=torch.empty(0, dtype=torch.int64),
|
||||
seq_lens_cpu=torch.empty(0, dtype=torch.int64),
|
||||
extend_seq_lens=torch.empty(0, dtype=torch.int32),
|
||||
extend_prefix_lens=torch.empty(0, dtype=torch.int64),
|
||||
extend_seq_lens_cpu=[],
|
||||
extend_prefix_lens_cpu=[],
|
||||
extend_logprob_start_lens_cpu=[],
|
||||
spec_info=spec_info,
|
||||
)
|
||||
|
||||
fb._pad_inputs_to_size(_mock_model_runner(), num_tokens=4, bs=1)
|
||||
|
||||
torch.testing.assert_close(
|
||||
fb.extend_seq_lens, torch.tensor([4], dtype=torch.int32)
|
||||
)
|
||||
torch.testing.assert_close(fb.extend_prefix_lens, torch.tensor([0]))
|
||||
self.assertEqual(fb.extend_seq_lens_cpu, [4])
|
||||
self.assertEqual(fb.extend_prefix_lens_cpu, [0])
|
||||
self.assertEqual(fb.extend_logprob_start_lens_cpu, [0])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -10,6 +10,7 @@ import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from sglang.srt.configs.model_config import AttentionArch
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.runtime_context import get_memory, get_parallel, get_server_args
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
@@ -79,6 +80,8 @@ def _make_model_runner(
|
||||
disaggregation_decode_extra_slots=0,
|
||||
kv_lora_rank=512,
|
||||
qk_rope_head_dim=64,
|
||||
swa_kv_lora_rank=128,
|
||||
swa_qk_rope_head_dim=32,
|
||||
):
|
||||
"""Create a mock ModelRunner with the fields configurators need."""
|
||||
mr = MagicMock()
|
||||
@@ -99,6 +102,9 @@ def _make_model_runner(
|
||||
mc.v_head_dim = v_head_dim
|
||||
mc.kv_lora_rank = kv_lora_rank
|
||||
mc.qk_rope_head_dim = qk_rope_head_dim
|
||||
mc.swa_kv_lora_rank = swa_kv_lora_rank
|
||||
mc.swa_qk_rope_head_dim = swa_qk_rope_head_dim
|
||||
mc.attention_arch = AttentionArch.MLA if use_mla_backend else AttentionArch.MHA
|
||||
mc.is_hybrid_swa = is_hybrid_swa
|
||||
mc.full_attention_layer_ids = (
|
||||
full_attention_layer_ids
|
||||
@@ -158,7 +164,9 @@ def _make_model_runner(
|
||||
mr.ps = ParallelState.trivial()
|
||||
mr.pp_group = SimpleNamespace(rank_in_group=0)
|
||||
mr.spec_aux_config = SimpleNamespace(
|
||||
eagle_draft_num_layers=None, dflash_draft_num_layers=None
|
||||
eagle_draft_num_layers=None,
|
||||
eagle_draft_swa_num_layers=None,
|
||||
dflash_draft_num_layers=None,
|
||||
)
|
||||
|
||||
return mr
|
||||
@@ -314,6 +322,58 @@ class TestHybridSWAConfigurator(CustomTestCase):
|
||||
self.assertLessEqual(used, available)
|
||||
self.assertGreater(used, available * 0.99)
|
||||
|
||||
@patch(
|
||||
"sglang.srt.mem_cache.kv_cache_configurator.calculate_mla_kv_cache_dim",
|
||||
return_value=576,
|
||||
)
|
||||
def test_mla_uses_full_and_swa_latent_geometry(
|
||||
self,
|
||||
mock_calculate_mla_kv_cache_dim,
|
||||
):
|
||||
"""Hybrid MLA pools must not be sized from MHA head geometry."""
|
||||
available = 10_000_000
|
||||
full_layers = 2
|
||||
swa_layers = 3
|
||||
swa_kv_lora_rank = 128
|
||||
swa_qk_rope_head_dim = 32
|
||||
mr = _make_model_runner(
|
||||
self,
|
||||
num_kv_heads=32,
|
||||
head_dim=256,
|
||||
v_head_dim=256,
|
||||
use_mla_backend=True,
|
||||
is_hybrid_swa=True,
|
||||
full_attention_layer_ids=list(range(full_layers)),
|
||||
swa_attention_layer_ids=list(range(full_layers, full_layers + swa_layers)),
|
||||
swa_num_kv_heads=16,
|
||||
swa_head_dim=128,
|
||||
swa_v_head_dim=128,
|
||||
swa_kv_lora_rank=swa_kv_lora_rank,
|
||||
swa_qk_rope_head_dim=swa_qk_rope_head_dim,
|
||||
swa_full_tokens_ratio=0.5,
|
||||
)
|
||||
|
||||
with mock_cpu_env(kv_size=2):
|
||||
from sglang.srt.model_executor.pool_configurator import (
|
||||
create_memory_pool_configurator,
|
||||
)
|
||||
|
||||
cfg = create_memory_pool_configurator(mr)
|
||||
config = cfg.calculate_pool_sizes(available, page_size=1)
|
||||
|
||||
expected_full_per_token = 576 * 2
|
||||
expected_swa_per_token = (swa_kv_lora_rank + swa_qk_rope_head_dim) * 2
|
||||
self.assertEqual(cfg._full_per_token, expected_full_per_token)
|
||||
self.assertEqual(cfg._swa_per_token, expected_swa_per_token)
|
||||
mock_calculate_mla_kv_cache_dim.assert_called_once()
|
||||
|
||||
used = (
|
||||
config.full_max_total_num_tokens * expected_full_per_token * full_layers
|
||||
+ config.swa_max_total_num_tokens * expected_swa_per_token * swa_layers
|
||||
)
|
||||
self.assertLessEqual(used, available)
|
||||
self.assertGreater(used, available * 0.99)
|
||||
|
||||
def test_ratio_respected(self):
|
||||
"""swa_tokens ~= full_tokens * ratio (within page alignment)"""
|
||||
available = 10_000_000
|
||||
@@ -443,6 +503,40 @@ class TestHybridSWAConfigurator(CustomTestCase):
|
||||
self.assertEqual(config.swa_max_total_num_tokens, 91)
|
||||
self.assertLessEqual(_actual_memory_used(mr, config), available)
|
||||
|
||||
def test_chunk_cache_cap_accounts_for_draft_swa_layers(self):
|
||||
"""Draft SWA tensors consume the same fixed-capacity pool as target SWA."""
|
||||
available = 1_000_000
|
||||
mr = _make_model_runner(
|
||||
self,
|
||||
is_hybrid_swa=True,
|
||||
full_attention_layer_ids=[0],
|
||||
swa_attention_layer_ids=[1],
|
||||
swa_num_kv_heads=4,
|
||||
disable_radix_cache=True,
|
||||
chunked_prefill_size=4,
|
||||
sliding_window_size=8,
|
||||
page_size=1,
|
||||
max_running_requests=2,
|
||||
)
|
||||
mr.spec_algorithm.is_eagle.return_value = True
|
||||
mr.spec_algorithm.is_none.return_value = False
|
||||
mr.spec_aux_config.eagle_draft_num_layers = 1
|
||||
mr.spec_aux_config.eagle_draft_swa_num_layers = 1
|
||||
|
||||
with mock_cpu_env():
|
||||
from sglang.srt.model_executor.pool_configurator import (
|
||||
create_memory_pool_configurator,
|
||||
)
|
||||
|
||||
cfg = create_memory_pool_configurator(mr)
|
||||
config = cfg.calculate_pool_sizes(available, page_size=1)
|
||||
|
||||
full_tokens = config.full_max_total_num_tokens
|
||||
swa_tokens = config.swa_max_total_num_tokens
|
||||
used = full_tokens * _full_per_token(mr) + swa_tokens * _swa_per_token(mr) * 2
|
||||
self.assertLessEqual(used, available)
|
||||
self.assertGreater(used, available * 0.99)
|
||||
|
||||
def test_chunk_cache_cap_drops_prefill_for_disagg_decode(self):
|
||||
available = 1_000_000
|
||||
mr = _make_model_runner(
|
||||
@@ -658,6 +752,45 @@ class TestEagleConfigurator(CustomTestCase):
|
||||
available,
|
||||
)
|
||||
|
||||
def test_hybrid_swa_draft_uses_swa_geometry_and_capacity(self):
|
||||
"""SWA draft layers use SWA KV geometry and capacity."""
|
||||
available = 10_000_000
|
||||
ratio = 0.25
|
||||
mr = _make_model_runner(
|
||||
self,
|
||||
num_kv_heads=8,
|
||||
head_dim=64,
|
||||
v_head_dim=64,
|
||||
num_layers=4,
|
||||
is_hybrid_swa=True,
|
||||
full_attention_layer_ids=[0, 1],
|
||||
swa_attention_layer_ids=[2, 3],
|
||||
swa_num_kv_heads=2,
|
||||
swa_head_dim=32,
|
||||
swa_v_head_dim=32,
|
||||
swa_full_tokens_ratio=ratio,
|
||||
)
|
||||
mr.spec_algorithm.is_eagle.return_value = True
|
||||
mr.spec_algorithm.is_none.return_value = False
|
||||
mr.spec_aux_config.eagle_draft_num_layers = 1
|
||||
mr.spec_aux_config.eagle_draft_swa_num_layers = 1
|
||||
|
||||
with mock_cpu_env():
|
||||
from sglang.srt.model_executor.pool_configurator import (
|
||||
create_memory_pool_configurator,
|
||||
)
|
||||
|
||||
cfg = create_memory_pool_configurator(mr)
|
||||
config = cfg.calculate_pool_sizes(available, page_size=1)
|
||||
|
||||
full_tokens = config.full_max_total_num_tokens
|
||||
swa_tokens = config.swa_max_total_num_tokens
|
||||
full_pt = _full_per_token(mr)
|
||||
swa_pt = _swa_per_token(mr)
|
||||
used = full_tokens * full_pt * 2 + swa_tokens * swa_pt * 3
|
||||
self.assertLessEqual(used, available)
|
||||
self.assertGreater(used, available * 0.99)
|
||||
|
||||
|
||||
class TestDSAIndexerAllocationPolicy(CustomTestCase):
|
||||
@patch(
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import re
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from sglang.test.test_utils import maybe_stub_sgl_kernel
|
||||
|
||||
maybe_stub_sgl_kernel()
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.managers.schedule_batch import Modality
|
||||
from sglang.srt.multimodal.processors.dots_note_omni import DotsNoteOmniProcessor
|
||||
from sglang.test.ci.ci_register import register_cpu_ci
|
||||
from sglang.test.test_utils import CustomTestCase
|
||||
|
||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||
|
||||
_IM_TOKEN_ID = 100
|
||||
_AUDIO_TOKEN_ID = 200
|
||||
# The dots.note chat template renders a video content part as this single token.
|
||||
_VIDEO_PLACEHOLDER = "<|video_pad|>"
|
||||
|
||||
|
||||
class _FakeMultimodalTokens:
|
||||
_pattern = re.compile(r"(<image>|<audio>)")
|
||||
|
||||
def get_combined_regex(self):
|
||||
return self._pattern
|
||||
|
||||
def get_modality_of_token(self, token):
|
||||
return {
|
||||
"<image>": Modality.IMAGE,
|
||||
"<audio>": Modality.AUDIO,
|
||||
}.get(token)
|
||||
|
||||
|
||||
class _FakeTokenizer:
|
||||
"""Maps each media marker to one pad id and other text to per-char ids."""
|
||||
|
||||
def encode(self, text, add_special_tokens=False):
|
||||
ids = []
|
||||
for part in _FakeMultimodalTokens._pattern.split(text):
|
||||
if part == "<image>":
|
||||
ids.append(_IM_TOKEN_ID)
|
||||
elif part == "<audio>":
|
||||
ids.append(_AUDIO_TOKEN_ID)
|
||||
else:
|
||||
ids.extend(ord(char) for char in part)
|
||||
return ids
|
||||
|
||||
|
||||
def _fake_preprocess_dots_video(raw_video, question, **kwargs):
|
||||
"""Each video flattens into one frame, one audio segment and the question."""
|
||||
return [
|
||||
{"type": "image_url", "image_url": {"url": f"{raw_video}-frame"}},
|
||||
{"type": "audio_url", "audio_url": {"url": f"{raw_video}-audio"}},
|
||||
{"type": "text", "text": question},
|
||||
]
|
||||
|
||||
|
||||
class TestDotsNoteOmniVideoMixing(CustomTestCase):
|
||||
def setUp(self):
|
||||
self.processor = DotsNoteOmniProcessor.__new__(DotsNoteOmniProcessor)
|
||||
self.processor.image_start_token = ""
|
||||
self.processor.image_token = "<image>"
|
||||
self.processor.image_end_token = ""
|
||||
self.processor.audio_start_token = ""
|
||||
self.processor.audio_token = "<audio>"
|
||||
self.processor.audio_end_token = ""
|
||||
self.processor.mm_tokens = _FakeMultimodalTokens()
|
||||
self.processor.video_placeholder_regex = re.compile(
|
||||
re.escape(_VIDEO_PLACEHOLDER)
|
||||
)
|
||||
|
||||
def test_multiple_videos_and_native_media_keep_prompt_order(self):
|
||||
prompt = f"{_VIDEO_PLACEHOLDER}<image>between{_VIDEO_PLACEHOLDER}question"
|
||||
all_video_media = {}
|
||||
contents = [
|
||||
[
|
||||
{"type": "image_url", "image_url": {"url": "video-0-frame"}},
|
||||
{"type": "text", "text": "question"},
|
||||
],
|
||||
[
|
||||
{"type": "audio_url", "audio_url": {"url": "video-1-audio"}},
|
||||
{"type": "text", "text": "question"},
|
||||
],
|
||||
]
|
||||
|
||||
for index, content in enumerate(contents):
|
||||
prompt, video_media = self.processor._render_video_content(
|
||||
prompt, "question", index, content
|
||||
)
|
||||
all_video_media.update(video_media)
|
||||
|
||||
prompt, images, audios = self.processor._merge_video_media(
|
||||
prompt,
|
||||
image_data=["native-image"],
|
||||
audio_data=None,
|
||||
video_media=all_video_media,
|
||||
)
|
||||
|
||||
self.assertEqual(prompt, "<image><image>between<audio>question")
|
||||
self.assertEqual(images, ["video-0-frame", "native-image"])
|
||||
self.assertEqual(audios, ["video-1-audio"])
|
||||
|
||||
def test_template_without_video_placeholders_inserts_each_video_once(self):
|
||||
prompt = "<|user|>question"
|
||||
all_video_media = {}
|
||||
|
||||
for index in range(2):
|
||||
prompt, video_media = self.processor._render_video_content(
|
||||
prompt,
|
||||
"question",
|
||||
index,
|
||||
[
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"video-{index}-frame"},
|
||||
},
|
||||
{"type": "text", "text": "question"},
|
||||
],
|
||||
)
|
||||
all_video_media.update(video_media)
|
||||
|
||||
prompt, images, audios = self.processor._merge_video_media(
|
||||
prompt, None, None, all_video_media
|
||||
)
|
||||
|
||||
self.assertEqual(prompt, "<|user|><image><image>question")
|
||||
self.assertEqual(images, ["video-0-frame", "video-1-frame"])
|
||||
self.assertEqual(audios, [])
|
||||
|
||||
|
||||
class TestDotsNoteOmniProcessMmDataAsync(CustomTestCase):
|
||||
"""Drive the request entry point that used to reject mixed video inputs."""
|
||||
|
||||
def setUp(self):
|
||||
self.processor = DotsNoteOmniProcessor.__new__(DotsNoteOmniProcessor)
|
||||
self.processor.image_start_token = ""
|
||||
self.processor.image_token = "<image>"
|
||||
self.processor.image_end_token = ""
|
||||
self.processor.audio_start_token = ""
|
||||
self.processor.audio_token = "<audio>"
|
||||
self.processor.audio_end_token = ""
|
||||
self.processor.mm_tokens = _FakeMultimodalTokens()
|
||||
self.processor.video_placeholder_regex = re.compile(
|
||||
re.escape(_VIDEO_PLACEHOLDER)
|
||||
)
|
||||
self.processor.mm_token_ids = {
|
||||
"im_start_id": 98,
|
||||
"im_token_id": _IM_TOKEN_ID,
|
||||
"im_end_id": 99,
|
||||
"audio_start_id": 198,
|
||||
"audio_token_id": _AUDIO_TOKEN_ID,
|
||||
"audio_end_id": 199,
|
||||
}
|
||||
self.processor._tokenizer = _FakeTokenizer()
|
||||
self.processor.audio_processor_config = types.SimpleNamespace(
|
||||
sampling_rate=16000
|
||||
)
|
||||
self.processor.image_preprocessor = types.SimpleNamespace(
|
||||
process_images=lambda images: (
|
||||
[torch.tensor([float(index)]) for index in range(len(images))],
|
||||
[torch.tensor([1, 1, 4]) for _ in images],
|
||||
["<image>" for _ in images],
|
||||
)
|
||||
)
|
||||
self.processor.io_executor = concurrent.futures.ThreadPoolExecutor(
|
||||
max_workers=2
|
||||
)
|
||||
self.addCleanup(self.processor.io_executor.shutdown)
|
||||
|
||||
self.loaded = {}
|
||||
|
||||
async def fake_load_mm_data(prompt, image_data=None, audio_data=None, **kwargs):
|
||||
self.loaded["prompt"] = prompt
|
||||
self.loaded["image_data"] = list(image_data or [])
|
||||
self.loaded["audio_data"] = list(audio_data or [])
|
||||
return types.SimpleNamespace(
|
||||
input_text=prompt,
|
||||
images=list(image_data or []),
|
||||
audios=[torch.zeros(4) for _ in audio_data or []],
|
||||
)
|
||||
|
||||
self.processor.load_mm_data = fake_load_mm_data
|
||||
|
||||
def _run(self, request_obj, **kwargs):
|
||||
with (
|
||||
patch(
|
||||
"sglang.srt.multimodal.processors.dots_note_omni.preprocess_dots_video",
|
||||
_fake_preprocess_dots_video,
|
||||
),
|
||||
patch(
|
||||
"sglang.srt.multimodal.processors.dots_note_omni.get_audio_token_string",
|
||||
lambda *args, **kwargs: "<audio>",
|
||||
),
|
||||
):
|
||||
return asyncio.run(
|
||||
self.processor.process_mm_data_async(
|
||||
request_obj.text,
|
||||
request_obj,
|
||||
max_req_input_len=4096,
|
||||
**kwargs,
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _request(text, video_data):
|
||||
return types.SimpleNamespace(
|
||||
text=text,
|
||||
video_data=video_data,
|
||||
video_config={
|
||||
"_question": "question",
|
||||
"seq": 131072,
|
||||
"audio_cap": 1.0,
|
||||
"audio_sr": 16000,
|
||||
"k_mode": "eval_ek",
|
||||
},
|
||||
sampling_params={"max_new_tokens": 16},
|
||||
rid="test-rid",
|
||||
)
|
||||
|
||||
def test_video_mixed_with_image_and_audio(self):
|
||||
request_obj = self._request(
|
||||
f"<image>{_VIDEO_PLACEHOLDER}middle<audio>question", ["video-0"]
|
||||
)
|
||||
|
||||
output = self._run(
|
||||
request_obj, image_data=["native-image"], audio_data=["native-audio"]
|
||||
)
|
||||
|
||||
self.assertEqual(
|
||||
self.loaded["prompt"], "<image><image><audio>middle<audio>question"
|
||||
)
|
||||
self.assertEqual(self.loaded["image_data"], ["native-image", "video-0-frame"])
|
||||
self.assertEqual(self.loaded["audio_data"], ["video-0-audio", "native-audio"])
|
||||
self.assertEqual(
|
||||
[item.modality for item in output.mm_items],
|
||||
[
|
||||
Modality.IMAGE,
|
||||
Modality.IMAGE,
|
||||
Modality.AUDIO,
|
||||
Modality.AUDIO,
|
||||
],
|
||||
)
|
||||
|
||||
def test_multiple_videos_mixed_with_image(self):
|
||||
request_obj = self._request(
|
||||
f"{_VIDEO_PLACEHOLDER}<image>middle{_VIDEO_PLACEHOLDER}question",
|
||||
["video-0", "video-1"],
|
||||
)
|
||||
|
||||
output = self._run(request_obj, image_data=["native-image"])
|
||||
|
||||
self.assertEqual(
|
||||
self.loaded["prompt"],
|
||||
"<image><audio><image>middle<image><audio>question",
|
||||
)
|
||||
self.assertEqual(
|
||||
self.loaded["image_data"],
|
||||
["video-0-frame", "native-image", "video-1-frame"],
|
||||
)
|
||||
self.assertEqual(self.loaded["audio_data"], ["video-0-audio", "video-1-audio"])
|
||||
self.assertEqual(len(output.mm_items), 5)
|
||||
|
||||
def test_extra_native_image_without_placeholder_is_rejected(self):
|
||||
request_obj = self._request(f"{_VIDEO_PLACEHOLDER}question", ["video-0"])
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "Image placeholder count"):
|
||||
self._run(request_obj, image_data=["native-image"])
|
||||
|
||||
def test_unconsumed_video_placeholder_is_rejected(self):
|
||||
request_obj = self._request(
|
||||
f"{_VIDEO_PLACEHOLDER}{_VIDEO_PLACEHOLDER}question", ["video-0"]
|
||||
)
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "Video placeholder count"):
|
||||
self._run(request_obj)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -6,6 +6,7 @@ tests check that the pre-allocated `parent_list` / `top_scores_index` match the
|
||||
slow path (`organize_draft_results`) for num_steps in {1, 2, 3, 4}.
|
||||
"""
|
||||
|
||||
import contextlib
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
@@ -134,6 +135,31 @@ class TestEagleWorkerV2Topk1FastPath(CustomTestCase):
|
||||
with self.assertRaises(AssertionError):
|
||||
worker._rebuild_topk1_chain_buffers()
|
||||
|
||||
def test_idle_draft_runs_each_eager_forward_without_tree_layout(self):
|
||||
worker = object.__new__(EagleDraftWorker)
|
||||
worker.speculative_num_steps = 3
|
||||
worker.draft_attn_backend = SimpleNamespace(attn_backends=[object(), object()])
|
||||
worker.draft_runner = SimpleNamespace(
|
||||
canary_manager=None,
|
||||
forward=MagicMock(),
|
||||
)
|
||||
spec_info = SimpleNamespace(hidden_states=torch.empty((0, 8), device=DEVICE))
|
||||
forward_batch = SimpleNamespace(
|
||||
forward_mode=ForwardMode.IDLE,
|
||||
input_ids=torch.empty((0,), dtype=torch.long, device=DEVICE),
|
||||
out_cache_loc=torch.empty((0,), dtype=torch.long, device=DEVICE),
|
||||
spec_info=spec_info,
|
||||
)
|
||||
|
||||
with patch(
|
||||
"sglang.srt.speculative.eagle_worker_v2.forward_context",
|
||||
side_effect=lambda *_args, **_kwargs: contextlib.nullcontext(),
|
||||
):
|
||||
result = worker.draft_forward(forward_batch)
|
||||
|
||||
self.assertEqual(result, (None, None, None, None))
|
||||
self.assertEqual(worker.draft_runner.forward.call_count, 2)
|
||||
|
||||
|
||||
class TestEagleWorkerV2BackendFallback(CustomTestCase):
|
||||
def setUp(self):
|
||||
|
||||
Reference in New Issue
Block a user