From bfeb7cd9b2a2954c637421dfb83f19bf9b19b015 Mon Sep 17 00:00:00 2001 From: Xinyuan Tong Date: Fri, 18 Sep 2026 18:31:47 +0000 Subject: [PATCH] model: support DeepSeek V4.1 vision with interleave prefill CP The CP runner bypassed the vision merge and used bare text embeddings. Merge image features before sharding so request-global offsets stay valid. Canonicalize model IDs separately to preserve scheduler hash IDs. Keep unsupported combinations guarded and isolate embedding overrides from multimodal prefills without starving queued FCFS requests. --- .../sglang/srt/arg_groups/deepseek_v4_hook.py | 45 ++- python/sglang/srt/managers/schedule_policy.py | 9 + python/sglang/srt/managers/scheduler.py | 2 + .../sglang/srt/managers/tokenizer_manager.py | 10 + .../sglang/srt/model_executor/model_runner.py | 1 + .../model_runner_components/misc_utils.py | 22 ++ .../srt/model_executor/runner/eager_runner.py | 13 +- python/sglang/srt/models/deepseek_v4.py | 69 +++-- python/sglang/srt/models/deepseek_v4_nextn.py | 1 + .../unit/managers/test_embed_overrides.py | 89 ++++++ .../unit/managers/test_prefill_adder.py | 141 ++++++++- .../test_deepseek_v41_vision_cp_inputs.py | 275 ++++++++++++++++++ .../unit/server_args/test_server_args.py | 91 ++++++ 13 files changed, 731 insertions(+), 37 deletions(-) create mode 100644 test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py diff --git a/python/sglang/srt/arg_groups/deepseek_v4_hook.py b/python/sglang/srt/arg_groups/deepseek_v4_hook.py index 500dc7240..d64fee256 100644 --- a/python/sglang/srt/arg_groups/deepseek_v4_hook.py +++ b/python/sglang/srt/arg_groups/deepseek_v4_hook.py @@ -181,12 +181,15 @@ def validate_deepseek_v41_features(server_args: ServerArgs) -> None: ) cfg = resolving_view(server_args) - if model_config_of(server_args).hf_config.model_type != "deepseek_v41": + hf_config = model_config_of(server_args).hf_config + if hf_config.model_type != "deepseek_v41": if cfg.enable_encoder_swa_bounded_replay: raise ValueError( "--enable-encoder-swa-bounded-replay requires DeepSeek-V4.1" ) return + if hf_config.vision_n_layers > 0 and cfg.enable_prefill_cp: + _validate_deepseek_v41_vision_prefill_cp(server_args) if cfg.enable_encoder_swa_bounded_replay: from sglang.srt.model_executor.cuda_graph_config import Backend @@ -197,7 +200,8 @@ def validate_deepseek_v41_features(server_args: ServerArgs) -> None: cfg.cuda_graph_config.prefill.backend != Backend.DISABLED, ), ("DP attention", cfg.enable_dp_attention), - ("context parallelism", cfg.attn_cp_size > 1), + # Prefill CP declares attn_cp_size and DP attention only later. + ("context parallelism", cfg.attn_cp_size > 1 or cfg.enable_prefill_cp), ("external cache linker", cfg.enable_unified_cache_external_linker), ("unified memory", cfg.enable_unified_memory), ("PD disaggregation", cfg.disaggregation_mode != "null"), @@ -306,3 +310,40 @@ def validate_deepseek_v41_features(server_args: ServerArgs) -> None: "--enable-decoder-swa-bounded-replay cannot be combined with " f"{feature} yet; disable one of them." ) + + +def _validate_deepseek_v41_vision_prefill_cp(server_args: ServerArgs) -> None: + from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase + + cfg = resolving_view(server_args) + if cfg.cp_strategy != "interleave": + raise ValueError( + "DeepSeek-V4.1 vision with prefill CP requires --cp-strategy " + f"interleave; got {cfg.cp_strategy!r}." + ) + if cfg.cuda_graph_config.prefill.backend != Backend.DISABLED: + # The CP runner merges image features eagerly; no capture path replays it. + locked = getattr(server_args, "_cuda_graph_config_locked", set()) + if (Phase.PREFILL, "backend") in locked: + raise ValueError( + "DeepSeek-V4.1 vision with prefill CP runs eager prefill; remove " + "the explicit prefill CUDA graph backend." + ) + declare_resolution( + server_args, + "validate_deepseek_v41_features", + cuda_graph_config=with_phase( + cfg.cuda_graph_config, Phase.PREFILL, backend=Backend.DISABLED + ), + ) + logger.warning( + "Disabling the prefill CUDA graph for DeepSeek-V4.1 vision with prefill CP." + ) + if ( + str(cfg.speculative_algorithm).upper() == "DSPARK" + and cfg.enable_decoder_swa_bounded_replay + ): + raise ValueError( + "DeepSeek-V4.1 vision with prefill CP does not support DSpark together " + "with --enable-decoder-swa-bounded-replay yet." + ) diff --git a/python/sglang/srt/managers/schedule_policy.py b/python/sglang/srt/managers/schedule_policy.py index f83b18178..5191db406 100644 --- a/python/sglang/srt/managers/schedule_policy.py +++ b/python/sglang/srt/managers/schedule_policy.py @@ -1019,6 +1019,15 @@ class PrefillAdder: else AddReqResult.CONTINUE ) + def can_share_extend_batch(self, req: Req) -> bool: + # Token embedding overrides embed the batch's raw input_ids before the + # model runs, and that lookup cannot index multimodal placeholder hash IDs. + if req.positional_embed_overrides is not None: + return all(r.multimodal_inputs is None for r in self.can_run_list) + if req.multimodal_inputs is not None: + return all(r.positional_embed_overrides is None for r in self.can_run_list) + return True + def add_chunked_req(self, req: Req): if self.dllm_config is not None: _rem_tokens = self._get_dllm_remain_tokens() diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index b694ba635..6c99d9e5c 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -3940,6 +3940,8 @@ class Scheduler( for req in self.waiting_queue: if self.enable_lora and not self.can_schedule_lora_req(req, running_loras): continue + if not adder.can_share_extend_batch(req): + break running_bs = len(running_batch.reqs) candidate_beam_width = ( diff --git a/python/sglang/srt/managers/tokenizer_manager.py b/python/sglang/srt/managers/tokenizer_manager.py index ef151ffba..528662e42 100644 --- a/python/sglang/srt/managers/tokenizer_manager.py +++ b/python/sglang/srt/managers/tokenizer_manager.py @@ -1279,6 +1279,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin): raise ValueError( "encoder SWA replay cannot return cached prompt logprobs" ) + requests_embed_overrides = obj.positional_embed_overrides is not None or ( + isinstance(obj, EmbeddingReqInput) + and obj.embed_overrides is not None + and obj.embed_override_token_id is not None + ) + if requests_embed_overrides and obj.contains_mm_input(): + raise ValueError( + "embedding overrides cannot be combined with image, video, or audio " + "inputs" + ) _max_req_len = self.context_len input_token_num = len(input_ids) if input_ids is not None else 0 input_token_num += self.num_reserved_tokens diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 9f7a69c7b..667d26e56 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -1646,6 +1646,7 @@ class ModelRunner: forward_batch.replace_embeds is not None and forward_batch.replace_positions is not None ): + misc_utils.validate_replace_embeds_batch(forward_batch) # Token embedding overrides: get base embeddings, scatter replacements if "input_embeds" not in kwargs: embed_layer = self.model.get_input_embeddings() diff --git a/python/sglang/srt/model_executor/model_runner_components/misc_utils.py b/python/sglang/srt/model_executor/model_runner_components/misc_utils.py index 026ca0434..4d00407f8 100644 --- a/python/sglang/srt/model_executor/model_runner_components/misc_utils.py +++ b/python/sglang/srt/model_executor/model_runner_components/misc_utils.py @@ -18,6 +18,7 @@ from sglang.srt.server_args import CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACK if TYPE_CHECKING: from sglang.srt.configs.model_config import ModelConfig + from sglang.srt.model_executor.forward_batch_info import ForwardBatch logger = logging.getLogger(__name__) @@ -105,3 +106,24 @@ def resolve_pp_proxy_dspark_hidden_size( if isinstance(model, _SupportsDSparkPPProxy): return model.get_pp_proxy_dspark_hidden_size() return 0 + + +def validate_replace_embeds_batch(forward_batch: ForwardBatch) -> None: + if forward_batch.mm_inputs is None: + return + for mm_inputs, prefix_len, extend_len in zip( + forward_batch.mm_inputs, + forward_batch.extend_prefix_lens_cpu, + forward_batch.extend_seq_lens_cpu, + ): + if mm_inputs is None: + continue + chunk_end = prefix_len + extend_len + for item in mm_inputs.mm_items: + for start, end in item.offsets or (): + if start < chunk_end and end >= prefix_len: + # Placeholder rows carry hash IDs the base embedding lookup cannot index. + raise ValueError( + "Token embedding overrides cannot share an extend batch with " + "multimodal placeholders" + ) diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py index 5714d01fb..28e5f8e72 100644 --- a/python/sglang/srt/model_executor/runner/eager_runner.py +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -387,12 +387,13 @@ class EagerRunner(BaseRunner): input_ids = forward_batch.input_ids input_embeds = kwargs.get("input_embeds") - # Multimodal spans must be embedded in global token order, before CP - # slicing. The model may also normalize image hash IDs for its router. - prepare_inputs = getattr(model, "prepare_language_model_inputs", None) - if prepare_inputs is not None: - input_ids, input_embeds = prepare_inputs( - input_ids, forward_batch, input_embeds + if hasattr(model, "prepare_model_inputs"): + # Multimodal offsets are request-global, so the merge and the + # placeholder-ID remap must see the full extend layout first. + input_ids, input_embeds = model.prepare_model_inputs( + input_ids=input_ids, + forward_batch=forward_batch, + input_embeds=input_embeds, ) if input_embeds is None: input_embeds = model.get_input_embeddings()(input_ids) diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index 3a352c484..5e8f4bcc7 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -74,6 +74,7 @@ from sglang.srt.layers.communicator_dsa_cp import ( dsa_cp_gather_hidden_states, dsa_cp_reduce_scatter_hidden_states, ) +from sglang.srt.layers.cp.base import is_zigzag from sglang.srt.layers.cp.cp_decode_attn_tp import get_cp_decode_attn_tp_ctx from sglang.srt.layers.cp.utils import ( cp_gather_full_sequence_states, @@ -4897,14 +4898,18 @@ class DeepseekV4ForCausalLM(nn.Module): and not getattr(config, "language_model_only", False) ): if ( - get_parallel().attn_cp_size != 1 - or get_pp_group().world_size != 1 + get_pp_group().world_size != 1 or not _v41_vision_a2a_supported() ): raise ValueError( - "V4.1 vision supports TP/EP/DP without CP or PP; " + "V4.1 vision supports TP/EP/DP without PP; " "MoE A2A is supported only with MegaMoE on a PD decode node" ) + if get_parallel().attn_cp_size != 1 and (_is_npu or is_zigzag()): + raise ValueError( + "V4.1 vision context parallelism requires the CUDA interleave " + "strategy; NPU and zigzag CP are not supported yet" + ) args = SimpleNamespace(**vars(config), dim=config.hidden_size) self.vision = ViT(args) @@ -5065,6 +5070,34 @@ class DeepseekV4ForCausalLM(nn.Module): def get_input_embeddings(self) -> nn.Module: return self.model.get_input_embeddings() + def prepare_model_inputs( + self, + input_ids: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor], + ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + if self.vision is None: + return input_ids, input_embeds + if ( + not forward_batch.forward_mode.is_decode() + and not forward_batch.forward_mode.is_target_verify() + and forward_batch.mm_inputs is not None + and any(x is not None for x in forward_batch.mm_inputs) + ): + if input_embeds is not None: + raise ValueError("Cannot combine input_embeds and image inputs") + input_embeds = self._prepare_mm_embeddings(input_ids, forward_batch) + if not ( + forward_batch.forward_mode.is_decode_or_idle() + or forward_batch.forward_mode.is_target_verify() + ): + # Decode/verify IDs are already vocabulary IDs; remap prompt image + # hashes for Engram and routing. + input_ids = input_ids.masked_fill( + input_ids >= MM_PAD_SHIFT_VALUE, self.config.image_token_id + ) + return input_ids, input_embeds + def set_dspark_layers_to_capture(self, layer_ids: List[int]) -> None: if not self.pp_group.is_last_rank: return @@ -5115,31 +5148,11 @@ class DeepseekV4ForCausalLM(nn.Module): input_ids: torch.Tensor, forward_batch: ForwardBatch, input_embeds: Optional[torch.Tensor] = None, - ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: - """Prepare full-sequence image embeddings and model IDs before CP splits. - - Scheduler hash IDs stay intact for multimodal cache keys; the language - model uses image_token_id for Engram masking and visual MoE routing. - """ - if ( - getattr(self, "vision", None) is not None - and not forward_batch.forward_mode.is_decode() - and not forward_batch.forward_mode.is_target_verify() - and forward_batch.mm_inputs is not None - and any(x is not None for x in forward_batch.mm_inputs) - ): - if input_embeds is not None: - raise ValueError("Cannot combine input_embeds and image inputs") - input_embeds = self._prepare_mm_embeddings(input_ids, forward_batch) - if getattr(self, "vision", None) is not None and not ( - forward_batch.forward_mode.is_decode_or_idle() - or forward_batch.forward_mode.is_target_verify() - ): - # Decode/verify IDs are already vocabulary IDs; remap prompt image - # hashes for Engram and routing. - input_ids = input_ids.masked_fill( - input_ids >= MM_PAD_SHIFT_VALUE, self.config.image_token_id - ) + pp_proxy_tensors: Optional[PPProxyTensors] = None, + ) -> torch.Tensor: + input_ids, input_embeds = self.prepare_model_inputs( + input_ids=input_ids, forward_batch=forward_batch, input_embeds=input_embeds + ) return input_ids, input_embeds diff --git a/python/sglang/srt/models/deepseek_v4_nextn.py b/python/sglang/srt/models/deepseek_v4_nextn.py index 6694abb60..d35f88454 100644 --- a/python/sglang/srt/models/deepseek_v4_nextn.py +++ b/python/sglang/srt/models/deepseek_v4_nextn.py @@ -222,6 +222,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM): self.quant_config = quant_config self.wo_a_fp8 = wo_a_fp8_gemm_enabled(quant_config) self.determine_num_fused_shared_experts() + self.vision = None self.model = DeepseekV4ModelNextN( config, quant_config, prefix=add_prefix("model", prefix) diff --git a/test/registered/unit/managers/test_embed_overrides.py b/test/registered/unit/managers/test_embed_overrides.py index e00a1d1fc..d2af78ec1 100644 --- a/test/registered/unit/managers/test_embed_overrides.py +++ b/test/registered/unit/managers/test_embed_overrides.py @@ -9,6 +9,7 @@ Covers: """ import unittest +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import torch @@ -17,10 +18,16 @@ from sglang.srt.constants import MIS_DELIMITER_TOKEN_ID from sglang.srt.entrypoints.openai.utils import convert_embeds_to_tensors from sglang.srt.managers.embed_types import PositionalEmbeds from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalInputs, +) from sglang.srt.managers.tokenizer_manager import TokenizerManager from sglang.srt.managers.tokenizer_manager_score_mixin import ( TokenizerManagerScoreMixin, ) +from sglang.srt.model_executor.model_runner import ModelRunner from sglang.srt.runtime_context import publish, reset_context from sglang.srt.server_args import ServerArgs from sglang.test.ci.ci_register import register_cpu_ci @@ -642,5 +649,87 @@ class TestScoreRequestValidation(CustomTestCase): ) +class TestEmbedOverridesRejectMultimodal(CustomTestCase): + def setUp(self): + reset_context() + self.addCleanup(reset_context) + publish(ServerArgs(model_path="dummy"), role="tokenizer") + self.manager = TokenizerManager.__new__(TokenizerManager) + self.manager.context_len = 128 + self.manager.num_reserved_tokens = 0 + self.manager.allow_auto_truncate = False + self.manager.validate_total_tokens = False + self.manager.is_generation = True + + def _request(self, **fields): + return GenerateReqInput( + input_ids=[10, 50, 20], + sampling_params={}, + positional_embed_overrides=PositionalEmbeds(embeds=[_vec()], positions=[1]), + **fields, + ) + + def test_request_with_image_is_rejected(self): + req = self._request(image_data=["image.png"]) + with self.assertRaisesRegex(ValueError, "overrides cannot be combined"): + self.manager._validate_one_request(req, req.input_ids) + text_only = self._request() + self.manager._validate_one_request(text_only, text_only.input_ids) + + def test_unresolved_embedding_overrides_with_image_are_rejected(self): + """EmbeddingReqInput resolves embed_overrides only after validation, so + the unresolved form must be caught at admission too.""" + self.manager.is_generation = False + req = EmbeddingReqInput( + input_ids=[10, 50, 20], + sampling_params={}, + embed_override_token_id=50, + embed_overrides=[_vec()], + image_data=["image.png"], + ) + with self.assertRaisesRegex(ValueError, "overrides cannot be combined"): + self.manager._validate_one_request(req, req.input_ids) + req.image_data = None + self.manager._validate_one_request(req, req.input_ids) + + def test_mixed_extend_batch_is_rejected_before_embedding_lookup(self): + """Placeholder rows hold hash IDs, so the base lookup must never run + on a batch whose chunk also covers multimodal placeholders.""" + embed_layer = MagicMock( + side_effect=AssertionError("embedding lookup must not run") + ) + runner = SimpleNamespace( + _pp_kwargs=lambda pp_proxy_tensors: {}, + model=SimpleNamespace(get_input_embeddings=lambda: embed_layer), + is_generation=True, + ) + image = MultimodalDataItem( + modality=Modality.IMAGE, feature=torch.zeros(1), offsets=[(0, 1)] + ) + image.set_hash(1234) + forward_batch = SimpleNamespace( + input_embeds=None, + input_ids=torch.tensor([1, 2, image.pad_value, image.pad_value]), + replace_embeds=torch.full((1, HIDDEN_DIM), 5.0), + replace_positions=torch.tensor([0]), + mm_inputs=[None, MultimodalInputs(mm_items=[image])], + extend_prefix_lens_cpu=[0, 0], + extend_seq_lens_cpu=[2, 2], + ) + + with self.assertRaisesRegex(ValueError, "cannot share an extend batch"): + ModelRunner._extend_forward_kwargs(runner, forward_batch, None) + embed_layer.assert_not_called() + + # A decoding image request in a mixed chunk has no placeholder rows here. + forward_batch.input_ids = torch.tensor([1, 2, 3]) + forward_batch.extend_prefix_lens_cpu = [0, 5] + forward_batch.extend_seq_lens_cpu = [2, 1] + embed_layer.side_effect = None + embed_layer.return_value = torch.zeros(3, HIDDEN_DIM) + kwargs = ModelRunner._extend_forward_kwargs(runner, forward_batch, None) + self.assertTrue(torch.equal(kwargs["input_embeds"][0], _vec(5.0))) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/managers/test_prefill_adder.py b/test/registered/unit/managers/test_prefill_adder.py index 03b21a95a..c56769f12 100644 --- a/test/registered/unit/managers/test_prefill_adder.py +++ b/test/registered/unit/managers/test_prefill_adder.py @@ -28,7 +28,13 @@ from sglang.srt.runtime_context import get_context from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.srt.utils.common import Range from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +import sglang.srt.managers.scheduler as scheduler_module +from sglang.srt.disaggregation.utils import DisaggregationMode +from sglang.srt.managers.scheduler import Scheduler register_cpu_ci(est_time=11, suite="base-a-test-cpu") @@ -296,6 +302,139 @@ class TestPrefillAdder(CustomTestCase): ) self.assertEqual(adder.can_run_list, [first]) + def test_embed_override_and_multimodal_requests_never_share_a_batch(self): + def tagged(rid, *, multimodal=False, overrides=False): + req = self.create_shared_req(rid) + req.multimodal_inputs = object() if multimodal else None + req.positional_embed_overrides = object() if overrides else None + return req + + for first, second in ( + (tagged("image", multimodal=True), tagged("override", overrides=True)), + (tagged("override", overrides=True), tagged("image", multimodal=True)), + ): + with self.subTest(first=first.rid): + adder = self.create_shared_adder() + self.assertTrue(adder.can_share_extend_batch(first)) + adder.add_one_req( + first, has_chunked_req=False, truncation_align_size=None + ) + self.assertEqual(adder.can_run_list, [first]) + self.assertFalse(adder.can_share_extend_batch(second)) + self.assertTrue(adder.can_share_extend_batch(tagged("text"))) + + adder = self.create_shared_adder() + chunked = tagged("chunked-image", multimodal=True) + chunked.full_untruncated_fill_ids = list(range(64)) + self.assertIs(adder.add_chunked_req(chunked), chunked) + self.assertFalse( + adder.can_share_extend_batch(tagged("override", overrides=True)) + ) + + def create_admission_scheduler(self, *, chunked_req) -> Scheduler: + allocator = self.create_token_allocator(available_size=4096) + allocator.page_size = 1 + self.mock_tree_cache.supports_mamba.return_value = False + self.mock_tree_cache.is_tree_cache.return_value = False + self.mock_tree_cache.supports_fast_match_prefix.return_value = False + self.mock_tree_cache.storage_prefetch_retries = None + scheduler = Scheduler.__new__(Scheduler) + scheduler.grammar_manager = SimpleNamespace(has_waiting_grammars=lambda: False) + scheduler.enable_priority_preemption = False + scheduler.enable_priority_scheduling = False + scheduler.is_hybrid_swa = False + scheduler.min_free_slots_delayer = None + scheduler.get_num_allocatable_reqs = lambda *args, **kwargs: 64 + scheduler.policy = SchedulePolicy( + policy="fcfs", + tree_cache=self.mock_tree_cache, + enable_hierarchical_cache=False, + enable_priority_scheduling=False, + schedule_low_priority_values_first=False, + ) + scheduler.processed_tokens_counter = 0 + scheduler.chunked_prefill_size = 16 + scheduler.dynamic_chunk_sizer = None + scheduler.tp_worker = SimpleNamespace( + model_runner=SimpleNamespace(attn_backend=object(), prefill_aware_swa=False) + ) + scheduler.page_size = 1 + scheduler.tree_cache = self.mock_tree_cache + scheduler.token_to_kv_pool_allocator = allocator + scheduler.new_token_ratio_tracker = SimpleNamespace(current=1.0) + scheduler.max_prefill_tokens = 16384 + scheduler.is_mixed_chunk = False + scheduler.priority_scheduling_preemption_threshold = 0 + scheduler.max_prefill_bs = 64 + scheduler.max_running_requests = 64 + scheduler.dllm_config = None + scheduler.enable_lora = False + scheduler.req_to_token_pool = SimpleNamespace() + scheduler.disaggregation_mode = DisaggregationMode.NULL + scheduler.enable_hicache_storage = False + scheduler.enable_hierarchical_cache = False + scheduler.enable_unified_cache_external_linker = False + scheduler.truncation_align_size = None + scheduler.model_config = None + scheduler.enable_overlap = False + scheduler.spec_algorithm = None + scheduler.load_inquirer = MagicMock() + scheduler.chunked_req = chunked_req + scheduler.waiting_queue = [] + return scheduler + + def run_admission_pass(self, scheduler: Scheduler) -> list: + running_batch = self.create_running_batch() + running_batch.batch_is_full = False + with ( + patch.object(scheduler_module, "ScheduleBatch") as schedule_batch, + patch.object(scheduler_module, "PrefillStats"), + patch.object(scheduler_module, "set_time_batch"), + ): + new_batch, _ = scheduler._get_new_batch_prefill_raw(None, running_batch) + if new_batch is None: + return [] + admitted = list(schedule_batch.init_new.call_args.args[0]) + for req in admitted: + req.prefix_indices = list(range(req.extend_range.end)) + return admitted + + def test_fcfs_admits_override_request_once_image_continuation_drains(self): + """An override request at the queue head must be admitted once the image + chunk ahead of it drains, even while more image requests keep arriving.""" + + def tagged(rid, length, *, multimodal=False, overrides=False): + req = self.create_shared_req(rid) + req.origin_input_ids = list(range(length)) + req.full_untruncated_fill_ids = list(range(length)) + req.multimodal_inputs = object() if multimodal else None + req.positional_embed_overrides = object() if overrides else None + req.beam_group = None + req.inflight_middle_chunks = 0 + return req + + continuation = tagged("image-continuation", 20, multimodal=True) + continuation.prefix_indices = list(range(16)) + scheduler = self.create_admission_scheduler(chunked_req=continuation) + override = tagged("override", 4, overrides=True) + scheduler.waiting_queue = [override] + + admitted_at = None + for pass_index in range(6): + scheduler.waiting_queue.append( + tagged(f"image-{pass_index}", 16, multimodal=True) + ) + admitted = self.run_admission_pass(scheduler) + self.assertFalse( + any(r.multimodal_inputs is not None for r in admitted) + and any(r.positional_embed_overrides is not None for r in admitted) + ) + if any(r is override for r in admitted): + admitted_at = pass_index + break + self.assertIsNotNone(admitted_at) + self.assertNotIn(override, scheduler.waiting_queue) + def test_shared_admission_rechecks_after_prefix_lock(self): adder = self.create_shared_adder() self.assertIsNotNone(adder.token_to_kv_pool_allocator.alloc(24)) diff --git a/test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py b/test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py new file mode 100644 index 000000000..01b09f578 --- /dev/null +++ b/test/registered/unit/models/test_deepseek_v41_vision_cp_inputs.py @@ -0,0 +1,275 @@ +"""Vision inputs under prefill CP merge on the full extend layout before the shard.""" + +import unittest +from contextlib import contextmanager +from types import SimpleNamespace +from unittest.mock import patch + +import torch +from torch import nn + +from sglang.srt.layers.cp.base import init_cp_strategy +from sglang.srt.layers.cp.utils import prepare_cp_forward +from sglang.srt.managers import mm_schedule +from sglang.srt.managers.schedule_batch import ( + Modality, + MultimodalDataItem, + MultimodalInputs, +) +from sglang.srt.model_executor.forward_batch_info import ForwardMode +from sglang.srt.model_executor.runner.eager_runner import EagerRunner +from sglang.srt.models.deepseek_v4 import DeepseekV4ForCausalLM +from sglang.srt.runtime_context import get_parallel +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=15, suite="base-a-test-cpu") + +HIDDEN = 8 +VOCAB = 64 +IMAGE_TOKEN_ID = 7 +CP_SIZE = 4 +# (prefix_len, extend_len) per request. Request 1 carries one image whose span +# [2, 8] starts inside its prefix, so only span rows 1..6 land in this chunk. +CHUNKS = [(0, 7), (3, 9), (1, 5)] +IMAGE_OFFSET = (2, 8) +IMAGE_HASH = 12345 +NUM_TOKENS = sum(extend_len for _, extend_len in CHUNKS) +# 21 tokens over 4 ranks give logical [6, 5, 5, 5], padded to the CP alignment. +PHYSICAL_ROWS = 8 +IMAGE_ROWS = torch.arange(7, 13) +POSITIONS = torch.cat([torch.arange(p, p + n) for p, n in CHUNKS]) + + +def _image_span(item: MultimodalDataItem) -> torch.Tensor: + start, end = item.offsets[0] + rows = end - start + 1 + return torch.arange(rows * HIDDEN, dtype=torch.float32).view(rows, HIDDEN) + 100.0 + + +def _pad(x: torch.Tensor) -> torch.Tensor: + return torch.cat([x, x.new_zeros(PHYSICAL_ROWS - x.shape[0], *x.shape[1:])]) + + +class _RecordingBody: + def __init__(self, embed: nn.Embedding): + self.embed = embed + self.calls = [] + + def get_input_embeddings(self): + return self.embed + + def __call__(self, input_ids, positions, forward_batch, input_embeds=None): + self.calls.append( + SimpleNamespace( + input_ids=input_ids, + positions=positions, + input_embeds=input_embeds, + input_ids_global=forward_batch.input_ids_global, + ) + ) + return input_embeds, input_embeds + + +class _VisionStub(DeepseekV4ForCausalLM): + def __init__(self, embed: nn.Embedding): + nn.Module.__init__(self) + self.config = SimpleNamespace(image_token_id=IMAGE_TOKEN_ID) + self.vision = object() + self.tp_size = 1 + self.model = _RecordingBody(embed) + self.pp_group = SimpleNamespace(is_last_rank=True) + self.lm_head = object() + self.capture_aux_hidden_states = False + self.logits_calls = [] + + def get_image_feature(self, items): + return [_image_span(item) for item in items] + + def logits_processor( + self, + input_ids, + hidden_states, + lm_head, + logits_metadata, + aux_hidden_states=None, + hidden_states_before_norm=None, + ): + self.logits_calls.append( + SimpleNamespace( + input_ids=input_ids, + hidden_states=hidden_states, + logits_metadata=logits_metadata, + hidden_states_before_norm=hidden_states_before_norm, + ) + ) + return object() + + +def _build_batch(): + item = MultimodalDataItem( + modality=Modality.IMAGE, feature=torch.zeros(1), offsets=[IMAGE_OFFSET] + ) + item.set_hash(IMAGE_HASH) + ids = list(range(10, 17)) + ids += [item.pad_value] * len(IMAGE_ROWS) + [20, 21, 22] + ids += list(range(30, 35)) + forward_batch = SimpleNamespace( + forward_mode=ForwardMode.EXTEND, + mm_inputs=[ + MultimodalInputs(mm_items=[]), + MultimodalInputs(mm_items=[item], im_token_id=IMAGE_TOKEN_ID), + None, + ], + extend_prefix_lens_cpu=[prefix for prefix, _ in CHUNKS], + extend_seq_lens_cpu=[extend_len for _, extend_len in CHUNKS], + seq_lens_cpu=[prefix + extend_len for prefix, extend_len in CHUNKS], + input_ids=torch.tensor(ids, dtype=torch.long), + positions=POSITIONS.clone(), + mm_input_embeds=None, + attn_cp_metadata=None, + global_num_tokens_cpu=None, + out_cache_loc=None, + input_ids_global=torch.zeros(1, dtype=torch.long), + ) + return forward_batch, item + + +def _expected_embeds(embed, scheduler_ids, item): + with torch.no_grad(): + full = embed(scheduler_ids.clamp(max=VOCAB - 1)) + full[IMAGE_ROWS] = _image_span(item)[1:7] + return full + + +def _canonical(scheduler_ids): + canonical = scheduler_ids.clone() + canonical[IMAGE_ROWS] = IMAGE_TOKEN_ID + return canonical + + +class TestDeepseekV41VisionPrefillCPInputs(CustomTestCase): + def setUp(self): + mm_schedule.init_mm_embedding_cache(1 << 20) + init_cp_strategy( + enable_prefill_cp=True, cp_size=CP_SIZE, cp_strategy="interleave" + ) + torch.manual_seed(0) + self.embed = nn.Embedding(VOCAB, HIDDEN) + self.model = _VisionStub(self.embed) + + def tearDown(self): + init_cp_strategy(enable_prefill_cp=False, cp_size=1, cp_strategy="interleave") + + @contextmanager + def _cp_collectives(self, full: torch.Tensor, rank: int): + def all_gather(output, input_tensor): + # Peers contribute their expected shards; this rank's rows come from + # what the runner actually handed to the collective. + output.zero_() + for peer in range(CP_SIZE): + rows = full[peer::CP_SIZE] + output[peer * PHYSICAL_ROWS : peer * PHYSICAL_ROWS + rows.shape[0]] = ( + rows + ) + output[rank * PHYSICAL_ROWS : (rank + 1) * PHYSICAL_ROWS] = input_tensor + + with ( + patch("torch.cuda.current_stream", return_value=None), + patch( + "sglang.srt.layers.cp.interleave.attn_cp_all_gather_into_tensor", + side_effect=all_gather, + ), + patch( + "sglang.srt.layers.cp.interleave.is_allocation_symmetric", + return_value=False, + ), + patch( + "sglang.srt.layers.cp.interleave.use_symmetric_memory", + return_value=torch.no_grad(), + ), + ): + yield + + def _prepare(self, forward_batch, input_embeds=None): + with torch.no_grad(): + return self.model.prepare_model_inputs( + input_ids=forward_batch.input_ids, + forward_batch=forward_batch, + input_embeds=input_embeds, + ) + + def test_prepare_model_inputs_merges_on_full_layout(self): + forward_batch, item = _build_batch() + scheduler_ids = forward_batch.input_ids.clone() + + model_ids, embeds = self._prepare(forward_batch) + + self.assertTrue(torch.equal(forward_batch.input_ids, scheduler_ids)) + self.assertIs(forward_batch.mm_input_embeds, embeds) + self.assertTrue(torch.equal(model_ids, _canonical(scheduler_ids))) + self.assertTrue(torch.equal(embeds[IMAGE_ROWS], _image_span(item)[1:7])) + text_rows = model_ids != IMAGE_TOKEN_ID + with torch.no_grad(): + text_embeds = self.embed(scheduler_ids[text_rows]) + self.assertTrue(torch.equal(embeds[text_rows], text_embeds)) + + def test_cp_runner_merges_before_shard(self): + runner = EagerRunner.__new__(EagerRunner) + runner.model_runner = SimpleNamespace(model=self.model) + padded = torch.zeros(CP_SIZE * PHYSICAL_ROWS, dtype=torch.long) + + for rank in range(CP_SIZE): + forward_batch, item = _build_batch() + routing_sentinel = forward_batch.input_ids_global + scheduler_ids = forward_batch.input_ids.clone() + canonical = _canonical(scheduler_ids) + full = _expected_embeds(self.embed, scheduler_ids, item) + padded[:NUM_TOKENS] = canonical + rank_major_ids = padded.view(-1, CP_SIZE).T.flatten() + self.model.model.calls.clear() + self.model.logits_calls.clear() + + with ( + get_parallel().override( + attn_cp_rank=rank, attn_cp_size=CP_SIZE, attn_cp_group=object() + ), + self._cp_collectives(full, rank), + torch.no_grad(), + ): + prepare_cp_forward(forward_batch) + runner._execute_extend_cp(forward_batch, {}) + + with self.subTest(rank=rank): + metadata = forward_batch.attn_cp_metadata + self.assertEqual(metadata.per_rank_actual_token, [PHYSICAL_ROWS] * 4) + (body,) = self.model.model.calls + self.assertTrue( + torch.equal(body.input_ids, _pad(canonical[rank::CP_SIZE])) + ) + self.assertTrue( + torch.equal(body.positions, _pad(POSITIONS[rank::CP_SIZE])) + ) + self.assertTrue( + torch.equal(body.input_embeds, _pad(full[rank::CP_SIZE])) + ) + self.assertTrue(torch.equal(body.input_ids_global, rank_major_ids)) + + (logits,) = self.model.logits_calls + self.assertTrue(torch.equal(logits.input_ids, canonical)) + self.assertTrue(torch.equal(logits.hidden_states, full)) + self.assertTrue(torch.equal(logits.hidden_states_before_norm, full)) + self.assertIs(logits.logits_metadata, forward_batch) + + self.assertTrue(torch.equal(forward_batch.mm_input_embeds, full)) + self.assertTrue(torch.equal(forward_batch.input_ids, scheduler_ids)) + self.assertIs(forward_batch.input_ids_global, routing_sentinel) + + def test_external_embeddings_with_images_are_rejected(self): + forward_batch, _ = _build_batch() + with self.assertRaisesRegex(ValueError, "Cannot combine"): + self._prepare(forward_batch, input_embeds=torch.zeros(NUM_TOKENS, HIDDEN)) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index fed3493e5..90dd7be50 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -29,6 +29,7 @@ from sglang.srt.arg_groups.cuda_graph_hook import ( finalize_cuda_graph_prefill_max_context, handle_cuda_graph_config, ) +from sglang.srt.arg_groups.deepseek_v4_hook import validate_deepseek_v41_features from sglang.srt.arg_groups.hicache_hook import ( handle_hicache, handle_hicache_ratio_default, @@ -45,6 +46,7 @@ from sglang.srt.arg_groups.kv_cache_hook import ( ) from sglang.srt.arg_groups.mamba_hook import handle_mamba_backend from sglang.srt.arg_groups.memory_hook import handle_gpu_memory_settings +from sglang.srt.arg_groups.model_hook import handle_model_specific_adjustments from sglang.srt.arg_groups.model_path_hook import handle_load_format from sglang.srt.arg_groups.moe_hook import ( handle_a2a_moe, @@ -4069,5 +4071,94 @@ class TestLazyReexports(CustomTestCase): server_args_module.NotAThing +class TestDeepseekV41VisionPrefillCPArgs(CustomTestCase): + def _args( + self, + *, + vision_n_layers=2, + prefill_backend=Backend.DISABLED, + lock_prefill_backend=False, + **overrides, + ): + fields = dict( + model_path="dummy", + enable_prefill_cp=True, + cp_strategy="interleave", + tp_size=2, + ) + fields.update(overrides) + server_args = ServerArgs(**fields) + server_args._model_config = SimpleNamespace( + hf_config=SimpleNamespace( + architectures=["DeepseekV4ForCausalLM"], + model_type="deepseek_v41", + vision_n_layers=vision_n_layers, + ), + nvfp4_moe_meta=None, + is_fp4_experts=False, + ) + # The dummy path does not initialize phase configs. + server_args.cuda_graph_config = CudaGraphConfig( + decode=PhaseConfig(backend=Backend.FULL, max_bs=512), + prefill=PhaseConfig(backend=prefill_backend, max_bs=512), + ) + server_args._resolved_overrides = [] + server_args._cuda_graph_config_locked = ( + {(Phase.PREFILL, "backend")} if lock_prefill_backend else set() + ) + return server_args + + @override_platform(is_cuda=True, is_hip=False) + def test_encoder_swa_replay_is_rejected_in_model_hook_order(self): + """The V4.1 validator runs before the CP validator declares attn_cp_size, + so encoder SWA replay used to pass resolution with vision prefill CP.""" + args = self._args( + enable_encoder_swa_bounded_replay=True, + max_running_requests=4, + chunked_prefill_size=128, + ) + with self.assertRaisesRegex( + ValueError, + "encoder-swa-bounded-replay does not support context parallelism", + ): + handle_model_specific_adjustments(args) + + def test_interleave_eager_prefill_is_accepted(self): + args = self._args() + validate_deepseek_v41_features(args) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.DISABLED, + ) + + def test_zigzag_is_rejected_only_with_vision(self): + with self.assertRaisesRegex(ValueError, "requires --cp-strategy interleave"): + validate_deepseek_v41_features(self._args(cp_strategy="zigzag")) + validate_deepseek_v41_features( + self._args(cp_strategy="zigzag", vision_n_layers=0) + ) + + def test_prefill_graph_explicit_rejects_and_default_resolves_eager(self): + with self.assertRaisesRegex(ValueError, "runs eager prefill"): + validate_deepseek_v41_features( + self._args(prefill_backend=Backend.BREAKABLE, lock_prefill_backend=True) + ) + args = self._args(prefill_backend=Backend.BREAKABLE) + validate_deepseek_v41_features(args) + self.assertEqual( + resolution_result(args, "cuda_graph_config").prefill.backend, + Backend.DISABLED, + ) + + def test_dspark_with_decoder_swa_bounded_replay_is_rejected(self): + with self.assertRaisesRegex(ValueError, "DSpark.*decoder-swa-bounded-replay"): + validate_deepseek_v41_features( + self._args( + speculative_algorithm="DSPARK", + enable_decoder_swa_bounded_replay=True, + ) + ) + + if __name__ == "__main__": unittest.main()