[Qwen3.5] Enable MTP spec_v2 and add test for nvidia/Qwen3.5-397B-A17B-NVFP4 (#19391)
This commit is contained in:
@@ -183,11 +183,6 @@ class HybridMambaDecodeReqToTokenPool(HybridReqToTokenPool):
|
||||
pre_alloc_size=pre_alloc_size,
|
||||
)
|
||||
|
||||
if envs.SGLANG_ENABLE_SPEC_V2.get() and not enable_mamba_extra_buffer:
|
||||
raise ValueError(
|
||||
"Spec v2 requires mamba scheduler strategy `extra_buffer` for mamba models. "
|
||||
"Please set `--mamba-scheduler-strategy extra_buffer`."
|
||||
)
|
||||
self.mamba_ping_pong_track_buffer_size = 2 if enable_overlap_schedule else 1
|
||||
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
|
||||
self.enable_memory_saver = enable_memory_saver
|
||||
|
||||
@@ -465,11 +465,6 @@ class HybridReqToTokenPool(ReqToTokenPool):
|
||||
device=device,
|
||||
enable_memory_saver=enable_memory_saver,
|
||||
)
|
||||
if envs.SGLANG_ENABLE_SPEC_V2.get() and not enable_mamba_extra_buffer:
|
||||
raise ValueError(
|
||||
"Spec v2 requires mamba scheduler strategy `extra_buffer` for mamba models. "
|
||||
"Please set `--mamba-scheduler-strategy extra_buffer`."
|
||||
)
|
||||
|
||||
self.mamba_ping_pong_track_buffer_size = 2 if enable_overlap_schedule else 1
|
||||
self.enable_mamba_extra_buffer = enable_mamba_extra_buffer
|
||||
|
||||
@@ -684,7 +684,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
mm_input, self.seq_lens_cpu[batch_idx]
|
||||
)
|
||||
mrope_positions_list[batch_idx] = mrope_positions
|
||||
elif self.forward_mode.is_extend():
|
||||
elif self.forward_mode.is_extend(include_draft_extend_v2=True):
|
||||
extend_seq_len, extend_prefix_len = (
|
||||
batch.extend_seq_lens[batch_idx],
|
||||
batch.extend_prefix_lens[batch_idx],
|
||||
|
||||
@@ -111,7 +111,7 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
|
||||
if (
|
||||
forward_batch.forward_mode.is_extend()
|
||||
and forward_batch.contains_mm_inputs()
|
||||
and not forward_batch.forward_mode.is_draft_extend()
|
||||
and not forward_batch.forward_mode.is_draft_extend(include_v2=True)
|
||||
):
|
||||
assert input_embeds is not None
|
||||
input_embeds = torch.cat(
|
||||
|
||||
@@ -1896,10 +1896,11 @@ class ServerArgs:
|
||||
self.disable_radix_cache = True
|
||||
self.disable_overlap_schedule = False
|
||||
else:
|
||||
logger.warning(
|
||||
f"Disabling radix cache since speculative decoding for {model_arch} is not supported with radix cache yet."
|
||||
)
|
||||
self.disable_radix_cache = True
|
||||
if not self.disable_radix_cache:
|
||||
raise ValueError(
|
||||
f"Speculative decoding for {model_arch} is not compatible with radix cache when using --mamba-scheduler-strategy no_buffer."
|
||||
"To use radix cache with speculative decoding, please use --mamba-scheduler-strategy extra_buffer and set SGLANG_ENABLE_SPEC_V2=1."
|
||||
)
|
||||
|
||||
def _handle_sampling_backend(self):
|
||||
if self.sampling_backend is None:
|
||||
|
||||
@@ -465,6 +465,7 @@ class EagleDraftWorker(BaseDraftWorker):
|
||||
batch: ModelWorkerBatch,
|
||||
target_hidden_states: torch.Tensor,
|
||||
next_token_ids: torch.Tensor,
|
||||
mm_input_embeds: Optional[torch.Tensor] = None,
|
||||
):
|
||||
"""
|
||||
Run draft model extend to correctly fill the KV cache.
|
||||
@@ -498,6 +499,8 @@ class EagleDraftWorker(BaseDraftWorker):
|
||||
|
||||
# Run forward
|
||||
forward_batch = ForwardBatch.init_new(batch, self.draft_runner)
|
||||
if mm_input_embeds is not None:
|
||||
forward_batch.mm_input_embeds = mm_input_embeds
|
||||
logits_output = self.draft_runner.forward(forward_batch).logits_output
|
||||
|
||||
# Update spec_info for the next draft step
|
||||
@@ -668,6 +671,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
||||
model_worker_batch,
|
||||
batch_output.logits_output.hidden_states,
|
||||
batch_output.next_token_ids,
|
||||
batch_output.logits_output.mm_input_embeds,
|
||||
)
|
||||
)
|
||||
return batch_output
|
||||
|
||||
Reference in New Issue
Block a user