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:
Xinyuan Tong
2026-09-23 14:36:00 +08:00
committed by minke.yu
parent ddf5207630
commit bfeb7cd9b2
13 changed files with 731 additions and 37 deletions
@@ -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()
@@ -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))
@@ -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,
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()