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.
This commit is contained in:
@@ -181,12 +181,15 @@ def validate_deepseek_v41_features(server_args: ServerArgs) -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
cfg = resolving_view(server_args)
|
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:
|
if cfg.enable_encoder_swa_bounded_replay:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--enable-encoder-swa-bounded-replay requires DeepSeek-V4.1"
|
"--enable-encoder-swa-bounded-replay requires DeepSeek-V4.1"
|
||||||
)
|
)
|
||||||
return
|
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:
|
if cfg.enable_encoder_swa_bounded_replay:
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
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,
|
cfg.cuda_graph_config.prefill.backend != Backend.DISABLED,
|
||||||
),
|
),
|
||||||
("DP attention", cfg.enable_dp_attention),
|
("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),
|
("external cache linker", cfg.enable_unified_cache_external_linker),
|
||||||
("unified memory", cfg.enable_unified_memory),
|
("unified memory", cfg.enable_unified_memory),
|
||||||
("PD disaggregation", cfg.disaggregation_mode != "null"),
|
("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 "
|
"--enable-decoder-swa-bounded-replay cannot be combined with "
|
||||||
f"{feature} yet; disable one of them."
|
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."
|
||||||
|
)
|
||||||
|
|||||||
@@ -1019,6 +1019,15 @@ class PrefillAdder:
|
|||||||
else AddReqResult.CONTINUE
|
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):
|
def add_chunked_req(self, req: Req):
|
||||||
if self.dllm_config is not None:
|
if self.dllm_config is not None:
|
||||||
_rem_tokens = self._get_dllm_remain_tokens()
|
_rem_tokens = self._get_dllm_remain_tokens()
|
||||||
|
|||||||
@@ -3940,6 +3940,8 @@ class Scheduler(
|
|||||||
for req in self.waiting_queue:
|
for req in self.waiting_queue:
|
||||||
if self.enable_lora and not self.can_schedule_lora_req(req, running_loras):
|
if self.enable_lora and not self.can_schedule_lora_req(req, running_loras):
|
||||||
continue
|
continue
|
||||||
|
if not adder.can_share_extend_batch(req):
|
||||||
|
break
|
||||||
|
|
||||||
running_bs = len(running_batch.reqs)
|
running_bs = len(running_batch.reqs)
|
||||||
candidate_beam_width = (
|
candidate_beam_width = (
|
||||||
|
|||||||
@@ -1279,6 +1279,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
|||||||
raise ValueError(
|
raise ValueError(
|
||||||
"encoder SWA replay cannot return cached prompt logprobs"
|
"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
|
_max_req_len = self.context_len
|
||||||
input_token_num = len(input_ids) if input_ids is not None else 0
|
input_token_num = len(input_ids) if input_ids is not None else 0
|
||||||
input_token_num += self.num_reserved_tokens
|
input_token_num += self.num_reserved_tokens
|
||||||
|
|||||||
@@ -1646,6 +1646,7 @@ class ModelRunner:
|
|||||||
forward_batch.replace_embeds is not None
|
forward_batch.replace_embeds is not None
|
||||||
and forward_batch.replace_positions 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
|
# Token embedding overrides: get base embeddings, scatter replacements
|
||||||
if "input_embeds" not in kwargs:
|
if "input_embeds" not in kwargs:
|
||||||
embed_layer = self.model.get_input_embeddings()
|
embed_layer = self.model.get_input_embeddings()
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from sglang.srt.server_args import CHUNKED_PREFIX_CACHE_SUPPORTED_ATTENTION_BACK
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -105,3 +106,24 @@ def resolve_pp_proxy_dspark_hidden_size(
|
|||||||
if isinstance(model, _SupportsDSparkPPProxy):
|
if isinstance(model, _SupportsDSparkPPProxy):
|
||||||
return model.get_pp_proxy_dspark_hidden_size()
|
return model.get_pp_proxy_dspark_hidden_size()
|
||||||
return 0
|
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"
|
||||||
|
)
|
||||||
|
|||||||
@@ -387,12 +387,13 @@ class EagerRunner(BaseRunner):
|
|||||||
|
|
||||||
input_ids = forward_batch.input_ids
|
input_ids = forward_batch.input_ids
|
||||||
input_embeds = kwargs.get("input_embeds")
|
input_embeds = kwargs.get("input_embeds")
|
||||||
# Multimodal spans must be embedded in global token order, before CP
|
if hasattr(model, "prepare_model_inputs"):
|
||||||
# slicing. The model may also normalize image hash IDs for its router.
|
# Multimodal offsets are request-global, so the merge and the
|
||||||
prepare_inputs = getattr(model, "prepare_language_model_inputs", None)
|
# placeholder-ID remap must see the full extend layout first.
|
||||||
if prepare_inputs is not None:
|
input_ids, input_embeds = model.prepare_model_inputs(
|
||||||
input_ids, input_embeds = prepare_inputs(
|
input_ids=input_ids,
|
||||||
input_ids, forward_batch, input_embeds
|
forward_batch=forward_batch,
|
||||||
|
input_embeds=input_embeds,
|
||||||
)
|
)
|
||||||
if input_embeds is None:
|
if input_embeds is None:
|
||||||
input_embeds = model.get_input_embeddings()(input_ids)
|
input_embeds = model.get_input_embeddings()(input_ids)
|
||||||
|
|||||||
@@ -74,6 +74,7 @@ from sglang.srt.layers.communicator_dsa_cp import (
|
|||||||
dsa_cp_gather_hidden_states,
|
dsa_cp_gather_hidden_states,
|
||||||
dsa_cp_reduce_scatter_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.cp_decode_attn_tp import get_cp_decode_attn_tp_ctx
|
||||||
from sglang.srt.layers.cp.utils import (
|
from sglang.srt.layers.cp.utils import (
|
||||||
cp_gather_full_sequence_states,
|
cp_gather_full_sequence_states,
|
||||||
@@ -4897,14 +4898,18 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
and not getattr(config, "language_model_only", False)
|
and not getattr(config, "language_model_only", False)
|
||||||
):
|
):
|
||||||
if (
|
if (
|
||||||
get_parallel().attn_cp_size != 1
|
get_pp_group().world_size != 1
|
||||||
or get_pp_group().world_size != 1
|
|
||||||
or not _v41_vision_a2a_supported()
|
or not _v41_vision_a2a_supported()
|
||||||
):
|
):
|
||||||
raise ValueError(
|
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"
|
"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)
|
args = SimpleNamespace(**vars(config), dim=config.hidden_size)
|
||||||
self.vision = ViT(args)
|
self.vision = ViT(args)
|
||||||
@@ -5065,6 +5070,34 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
def get_input_embeddings(self) -> nn.Module:
|
def get_input_embeddings(self) -> nn.Module:
|
||||||
return self.model.get_input_embeddings()
|
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:
|
def set_dspark_layers_to_capture(self, layer_ids: List[int]) -> None:
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
return
|
return
|
||||||
@@ -5115,30 +5148,10 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
input_ids: torch.Tensor,
|
input_ids: torch.Tensor,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
input_embeds: Optional[torch.Tensor] = None,
|
input_embeds: Optional[torch.Tensor] = None,
|
||||||
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
|
pp_proxy_tensors: Optional[PPProxyTensors] = None,
|
||||||
"""Prepare full-sequence image embeddings and model IDs before CP splits.
|
) -> torch.Tensor:
|
||||||
|
input_ids, input_embeds = self.prepare_model_inputs(
|
||||||
Scheduler hash IDs stay intact for multimodal cache keys; the language
|
input_ids=input_ids, forward_batch=forward_batch, input_embeds=input_embeds
|
||||||
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
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return input_ids, input_embeds
|
return input_ids, input_embeds
|
||||||
|
|||||||
@@ -222,6 +222,7 @@ class DeepseekV4ForCausalLMNextN(DeepseekV4ForCausalLM):
|
|||||||
self.quant_config = quant_config
|
self.quant_config = quant_config
|
||||||
self.wo_a_fp8 = wo_a_fp8_gemm_enabled(quant_config)
|
self.wo_a_fp8 = wo_a_fp8_gemm_enabled(quant_config)
|
||||||
self.determine_num_fused_shared_experts()
|
self.determine_num_fused_shared_experts()
|
||||||
|
self.vision = None
|
||||||
|
|
||||||
self.model = DeepseekV4ModelNextN(
|
self.model = DeepseekV4ModelNextN(
|
||||||
config, quant_config, prefix=add_prefix("model", prefix)
|
config, quant_config, prefix=add_prefix("model", prefix)
|
||||||
|
|||||||
@@ -9,6 +9,7 @@ Covers:
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
import torch
|
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.entrypoints.openai.utils import convert_embeds_to_tensors
|
||||||
from sglang.srt.managers.embed_types import PositionalEmbeds
|
from sglang.srt.managers.embed_types import PositionalEmbeds
|
||||||
from sglang.srt.managers.io_struct import EmbeddingReqInput, GenerateReqInput
|
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 import TokenizerManager
|
||||||
from sglang.srt.managers.tokenizer_manager_score_mixin import (
|
from sglang.srt.managers.tokenizer_manager_score_mixin import (
|
||||||
TokenizerManagerScoreMixin,
|
TokenizerManagerScoreMixin,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
from sglang.srt.runtime_context import publish, reset_context
|
from sglang.srt.runtime_context import publish, reset_context
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -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.server_args import ServerArgs, set_global_server_args_for_scheduler
|
||||||
from sglang.srt.utils.common import Range
|
from sglang.srt.utils.common import Range
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
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")
|
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])
|
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):
|
def test_shared_admission_rechecks_after_prefix_lock(self):
|
||||||
adder = self.create_shared_adder()
|
adder = self.create_shared_adder()
|
||||||
self.assertIsNotNone(adder.token_to_kv_pool_allocator.alloc(24))
|
self.assertIsNotNone(adder.token_to_kv_pool_allocator.alloc(24))
|
||||||
|
|||||||
@@ -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()
|
||||||
@@ -29,6 +29,7 @@ from sglang.srt.arg_groups.cuda_graph_hook import (
|
|||||||
finalize_cuda_graph_prefill_max_context,
|
finalize_cuda_graph_prefill_max_context,
|
||||||
handle_cuda_graph_config,
|
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 (
|
from sglang.srt.arg_groups.hicache_hook import (
|
||||||
handle_hicache,
|
handle_hicache,
|
||||||
handle_hicache_ratio_default,
|
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.mamba_hook import handle_mamba_backend
|
||||||
from sglang.srt.arg_groups.memory_hook import handle_gpu_memory_settings
|
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.model_path_hook import handle_load_format
|
||||||
from sglang.srt.arg_groups.moe_hook import (
|
from sglang.srt.arg_groups.moe_hook import (
|
||||||
handle_a2a_moe,
|
handle_a2a_moe,
|
||||||
@@ -4069,5 +4071,94 @@ class TestLazyReexports(CustomTestCase):
|
|||||||
server_args_module.NotAThing
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user