[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:
Jianfei Wang
2026-08-22 14:19:14 +08:00
committed by GitHub
co-authored by miraclezqc
parent c35683fda0
commit af39ad9349
55 changed files with 9638 additions and 154 deletions
@@ -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):