[BugFix] Resolve adaptive speculative decoding conflicts for Qwen3.5 (hybrid GDN) (#23331)
Signed-off-by: EanWang211123 <wangyiheng@sangfor.com.cn> Co-authored-by: shuwenn <47200617+alphabetc1@users.noreply.github.com> Co-authored-by: shuwenn <2508695655@qq.com>
This commit is contained in:
co-authored by
shuwenn
shuwenn
parent
16bcc4583e
commit
b9c2bf717b
@@ -258,7 +258,9 @@ def fused_sigmoid_gating_delta_rule_update(
|
|||||||
disable_state_update: bool = False,
|
disable_state_update: bool = False,
|
||||||
intermediate_states_buffer: Optional[torch.Tensor] = None,
|
intermediate_states_buffer: Optional[torch.Tensor] = None,
|
||||||
intermediate_state_indices: Optional[torch.Tensor] = None,
|
intermediate_state_indices: Optional[torch.Tensor] = None,
|
||||||
cache_steps: Optional[int] = None,
|
cache_steps: Optional[
|
||||||
|
int
|
||||||
|
] = None, # kept for API compat; stride is derived from ``intermediate_states_buffer.shape[1]``
|
||||||
retrieve_parent_token: Optional[torch.Tensor] = None,
|
retrieve_parent_token: Optional[torch.Tensor] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -307,6 +309,14 @@ def fused_sigmoid_gating_delta_rule_update(
|
|||||||
|
|
||||||
grid = (NK, NV, N * HV)
|
grid = (NK, NV, N * HV)
|
||||||
|
|
||||||
|
# Per-req stride must match the buffer's allocated dim, not runtime steps
|
||||||
|
# (they can differ under --speculative-adaptive).
|
||||||
|
cache_stride_steps = (
|
||||||
|
intermediate_states_buffer.shape[1]
|
||||||
|
if intermediate_states_buffer is not None
|
||||||
|
else 0
|
||||||
|
)
|
||||||
|
|
||||||
fused_sigmoid_gating_delta_rule_update_kernel[grid](
|
fused_sigmoid_gating_delta_rule_update_kernel[grid](
|
||||||
A_log=A_log,
|
A_log=A_log,
|
||||||
a=a,
|
a=a,
|
||||||
@@ -323,7 +333,7 @@ def fused_sigmoid_gating_delta_rule_update(
|
|||||||
cu_seqlens=cu_seqlens,
|
cu_seqlens=cu_seqlens,
|
||||||
intermediate_states_buffer=intermediate_states_buffer,
|
intermediate_states_buffer=intermediate_states_buffer,
|
||||||
intermediate_state_indices=intermediate_state_indices,
|
intermediate_state_indices=intermediate_state_indices,
|
||||||
cache_steps=0 if cache_steps is None else cache_steps,
|
cache_steps=cache_stride_steps,
|
||||||
retrieve_parent_token_ptr=retrieve_parent_token,
|
retrieve_parent_token_ptr=retrieve_parent_token,
|
||||||
stride_retrieve_parent_token_seq=stride_retrieve_parent_token_seq,
|
stride_retrieve_parent_token_seq=stride_retrieve_parent_token_seq,
|
||||||
stride_retrieve_parent_token_token=stride_retrieve_parent_token_token,
|
stride_retrieve_parent_token_token=stride_retrieve_parent_token_token,
|
||||||
|
|||||||
@@ -16,6 +16,7 @@ from sglang.srt.layers.attention.mamba.mamba_state_scatter_triton import (
|
|||||||
fused_mamba_state_scatter_with_mask,
|
fused_mamba_state_scatter_with_mask,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
|
from sglang.srt.layers.radix_linear_attention import RadixLinearAttention
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
@@ -763,6 +764,20 @@ class HybridLinearAttnBackend(AttentionBackend):
|
|||||||
self.linear_attn_backend = linear_attn_backend
|
self.linear_attn_backend = linear_attn_backend
|
||||||
self.attn_backend_list = [full_attn_backend, linear_attn_backend]
|
self.attn_backend_list = [full_attn_backend, linear_attn_backend]
|
||||||
|
|
||||||
|
def _is_full_attn(
|
||||||
|
self, layer: Optional[RadixAttention], layer_id: Optional[int] = None
|
||||||
|
) -> bool:
|
||||||
|
# Dispatch by the layer's runtime type
|
||||||
|
if isinstance(layer, RadixLinearAttention):
|
||||||
|
return False
|
||||||
|
if isinstance(layer, RadixAttention):
|
||||||
|
return True
|
||||||
|
|
||||||
|
if layer is not None:
|
||||||
|
layer_id = layer.layer_id
|
||||||
|
assert layer_id is not None, "either layer or layer_id must be provided"
|
||||||
|
return layer_id in self.full_attn_layers
|
||||||
|
|
||||||
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
def init_forward_metadata(self, forward_batch: ForwardBatch):
|
||||||
for attn_backend in self.attn_backend_list:
|
for attn_backend in self.attn_backend_list:
|
||||||
attn_backend.init_forward_metadata(forward_batch)
|
attn_backend.init_forward_metadata(forward_batch)
|
||||||
@@ -859,8 +874,7 @@ class HybridLinearAttnBackend(AttentionBackend):
|
|||||||
b: Optional[torch.Tensor] = None, # For GDN linear attention
|
b: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
layer_id = layer.layer_id if layer else kwargs["layer_id"]
|
if self._is_full_attn(layer, kwargs.get("layer_id")):
|
||||||
if layer_id in self.full_attn_layers:
|
|
||||||
return self.full_attn_backend.forward_decode(
|
return self.full_attn_backend.forward_decode(
|
||||||
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
)
|
)
|
||||||
@@ -891,8 +905,7 @@ class HybridLinearAttnBackend(AttentionBackend):
|
|||||||
b: Optional[torch.Tensor] = None, # For GDN linear attention
|
b: Optional[torch.Tensor] = None, # For GDN linear attention
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
layer_id = layer.layer_id if layer else kwargs["layer_id"]
|
if self._is_full_attn(layer, kwargs.get("layer_id")):
|
||||||
if layer_id in self.full_attn_layers:
|
|
||||||
return self.full_attn_backend.forward_extend(
|
return self.full_attn_backend.forward_extend(
|
||||||
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
q, k, v, layer, forward_batch, save_kv_cache, **kwargs
|
||||||
)
|
)
|
||||||
@@ -923,8 +936,7 @@ class HybridLinearAttnBackend(AttentionBackend):
|
|||||||
b: Optional[torch.Tensor] = None, # For linear attention
|
b: Optional[torch.Tensor] = None, # For linear attention
|
||||||
**kwargs,
|
**kwargs,
|
||||||
):
|
):
|
||||||
layer_id = layer.layer_id if layer else kwargs["layer_id"]
|
is_linear_attn = not self._is_full_attn(layer, kwargs.get("layer_id"))
|
||||||
is_linear_attn = layer_id not in self.full_attn_layers
|
|
||||||
|
|
||||||
if forward_batch.forward_mode.is_idle():
|
if forward_batch.forward_mode.is_idle():
|
||||||
if is_linear_attn:
|
if is_linear_attn:
|
||||||
|
|||||||
@@ -242,9 +242,12 @@ class ModelRunnerKVCacheMixin:
|
|||||||
# Initialize req_to_token_pool
|
# Initialize req_to_token_pool
|
||||||
if self.req_to_token_pool is None:
|
if self.req_to_token_pool is None:
|
||||||
# FIXME(lsyin): this is the temporary fix for the context length issue when using speculative decoding
|
# FIXME(lsyin): this is the temporary fix for the context length issue when using speculative decoding
|
||||||
|
max_spec_draft_tokens = (
|
||||||
|
self.server_args.effective_max_speculative_num_draft_tokens()
|
||||||
|
)
|
||||||
extra_max_context_len = 4
|
extra_max_context_len = 4
|
||||||
if self.server_args.speculative_num_draft_tokens is not None:
|
if max_spec_draft_tokens is not None:
|
||||||
extra_max_context_len += self.server_args.speculative_num_draft_tokens
|
extra_max_context_len += max_spec_draft_tokens
|
||||||
|
|
||||||
if self.server_args.disaggregation_mode == "decode":
|
if self.server_args.disaggregation_mode == "decode":
|
||||||
from sglang.srt.disaggregation.decode import (
|
from sglang.srt.disaggregation.decode import (
|
||||||
@@ -274,7 +277,7 @@ class ModelRunnerKVCacheMixin:
|
|||||||
if self.start_layer <= i < self.end_layer
|
if self.start_layer <= i < self.end_layer
|
||||||
]
|
]
|
||||||
),
|
),
|
||||||
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
speculative_num_draft_tokens=max_spec_draft_tokens,
|
||||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
pre_alloc_size=pre_alloc_size,
|
pre_alloc_size=pre_alloc_size,
|
||||||
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
||||||
@@ -308,7 +311,7 @@ class ModelRunnerKVCacheMixin:
|
|||||||
]
|
]
|
||||||
),
|
),
|
||||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
speculative_num_draft_tokens=max_spec_draft_tokens,
|
||||||
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
||||||
start_layer=self.start_layer,
|
start_layer=self.start_layer,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@
|
|||||||
|
|
||||||
"""Inference-only Qwen3_5 MTP model."""
|
"""Inference-only Qwen3_5 MTP model."""
|
||||||
|
|
||||||
|
import copy
|
||||||
import logging
|
import logging
|
||||||
from typing import Iterable, Optional, Tuple
|
from typing import Iterable, Optional, Tuple
|
||||||
|
|
||||||
@@ -51,6 +52,9 @@ class Qwen3_5ForCausalLMMTP(nn.Module):
|
|||||||
if self.is_multimodal:
|
if self.is_multimodal:
|
||||||
config = config.text_config
|
config = config.text_config
|
||||||
|
|
||||||
|
# Deep-copy so MTP mutations below don't leak into the target's config.
|
||||||
|
config = copy.deepcopy(config)
|
||||||
|
|
||||||
# The MTP model is unquantized in the nvfp4 checkpoint.
|
# The MTP model is unquantized in the nvfp4 checkpoint.
|
||||||
if quant_config and quant_config.get_name() == "modelopt_fp4":
|
if quant_config and quant_config.get_name() == "modelopt_fp4":
|
||||||
quant_config = None
|
quant_config = None
|
||||||
|
|||||||
@@ -14,6 +14,7 @@
|
|||||||
|
|
||||||
"""Inference-only Qwen3Next MTP Speculative Decoding."""
|
"""Inference-only Qwen3Next MTP Speculative Decoding."""
|
||||||
|
|
||||||
|
import copy
|
||||||
import logging
|
import logging
|
||||||
from typing import Iterable, Optional, Tuple
|
from typing import Iterable, Optional, Tuple
|
||||||
|
|
||||||
@@ -44,6 +45,8 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
|
|||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
) -> None:
|
) -> None:
|
||||||
nn.Module.__init__(self)
|
nn.Module.__init__(self)
|
||||||
|
# Deep-copy so MTP mutations below don't leak into the target's config.
|
||||||
|
config = copy.deepcopy(config)
|
||||||
self.config = config
|
self.config = config
|
||||||
self.tp_size = get_tensor_model_parallel_world_size()
|
self.tp_size = get_tensor_model_parallel_world_size()
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -6846,6 +6846,25 @@ class ServerArgs:
|
|||||||
def enable_mamba_extra_buffer(self) -> bool:
|
def enable_mamba_extra_buffer(self) -> bool:
|
||||||
return self.mamba_scheduler_strategy == "extra_buffer"
|
return self.mamba_scheduler_strategy == "extra_buffer"
|
||||||
|
|
||||||
|
def effective_max_speculative_num_draft_tokens(self) -> Optional[int]:
|
||||||
|
"""Return the maximum draft-token count runtime speculative decoding may use."""
|
||||||
|
if self.speculative_num_draft_tokens is None:
|
||||||
|
return None
|
||||||
|
if not self.speculative_adaptive:
|
||||||
|
return self.speculative_num_draft_tokens
|
||||||
|
|
||||||
|
from sglang.srt.speculative.adaptive_spec_params import (
|
||||||
|
resolve_candidate_steps_from_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
candidate_steps = resolve_candidate_steps_from_config(
|
||||||
|
initial_steps=self.speculative_num_steps,
|
||||||
|
cfg_path=self.speculative_adaptive_config,
|
||||||
|
)
|
||||||
|
# TODO: adaptive spec currently requires topk=1, so each runtime state
|
||||||
|
# needs steps + 1 draft-token slots. Revisit this if topk>1 is supported.
|
||||||
|
return max(candidate_steps) + 1
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def mamba_cache_chunk_size(self) -> int:
|
def mamba_cache_chunk_size(self) -> int:
|
||||||
# For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE and page_size.
|
# For mamba cache with extra buffer, the chunk size is the max of FLA_CHUNK_SIZE and page_size.
|
||||||
|
|||||||
@@ -2,10 +2,7 @@ import logging
|
|||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, Protocol
|
from typing import TYPE_CHECKING, Protocol
|
||||||
|
|
||||||
from sglang.srt.speculative.adaptive_spec_params import (
|
from sglang.srt.speculative.adaptive_spec_params import AdaptiveSpeculativeParams
|
||||||
AdaptiveSpeculativeParams,
|
|
||||||
load_adaptive_config,
|
|
||||||
)
|
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
@@ -76,10 +73,9 @@ class AdaptiveController:
|
|||||||
|
|
||||||
def __init__(self, worker: AdaptiveSpecWorker, config_path: str | None = None):
|
def __init__(self, worker: AdaptiveSpecWorker, config_path: str | None = None):
|
||||||
self.worker = worker
|
self.worker = worker
|
||||||
cfg = load_adaptive_config(config_path)
|
|
||||||
self.params = AdaptiveSpeculativeParams(
|
self.params = AdaptiveSpeculativeParams(
|
||||||
initial_steps=worker.speculative_num_steps,
|
initial_steps=worker.speculative_num_steps,
|
||||||
config=cfg,
|
cfg_path=config_path,
|
||||||
)
|
)
|
||||||
self._states: dict[int, SpecRuntimeState] = {}
|
self._states: dict[int, SpecRuntimeState] = {}
|
||||||
|
|
||||||
|
|||||||
@@ -73,6 +73,35 @@ def load_adaptive_config(path: str | None) -> dict[str, object]:
|
|||||||
return cfg
|
return cfg
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_candidate_steps(initial_steps: int, cfg: dict[str, object]) -> list[int]:
|
||||||
|
"""Return sorted, deduplicated candidate steps; inserts *initial_steps* when missing."""
|
||||||
|
raw = cfg.get("candidate_steps") or (1, 3, 7)
|
||||||
|
candidates: set[int] = set(raw)
|
||||||
|
|
||||||
|
# Ensure the worker's initial speculative_num_steps is itself a candidate.
|
||||||
|
# Otherwise AdaptiveController.register() would store the worker's pre-built
|
||||||
|
# runtime state under a key that _activate() never queries, leaking that
|
||||||
|
# state's draft attn backend and cuda graph buffers for the process lifetime.
|
||||||
|
if initial_steps not in candidates:
|
||||||
|
log_info_on_rank0(
|
||||||
|
logger,
|
||||||
|
f"Adding initial speculative_num_steps={initial_steps} to "
|
||||||
|
f"candidate_steps={sorted(candidates)} so the pre-built "
|
||||||
|
f"runtime state is reused.",
|
||||||
|
)
|
||||||
|
candidates.add(initial_steps)
|
||||||
|
|
||||||
|
return sorted(candidates)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_candidate_steps_from_config(
|
||||||
|
initial_steps: int, cfg_path: str | None
|
||||||
|
) -> list[int]:
|
||||||
|
"""Load adaptive config and resolve candidate steps."""
|
||||||
|
cfg = load_adaptive_config(cfg_path)
|
||||||
|
return _resolve_candidate_steps(initial_steps, cfg)
|
||||||
|
|
||||||
|
|
||||||
class AdaptiveSpeculativeParams:
|
class AdaptiveSpeculativeParams:
|
||||||
"""Tracks acceptance rate via EMA and adapts num_steps accordingly.
|
"""Tracks acceptance rate via EMA and adapts num_steps accordingly.
|
||||||
|
|
||||||
@@ -88,26 +117,11 @@ class AdaptiveSpeculativeParams:
|
|||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
initial_steps: int,
|
initial_steps: int,
|
||||||
config: dict[str, object] | None = None,
|
cfg_path: str | None = None,
|
||||||
):
|
):
|
||||||
cfg = config or {}
|
cfg = load_adaptive_config(cfg_path)
|
||||||
# TODO: Wider range of candidate_steps (once lazy init is supported).
|
# TODO: Wider range of candidate_steps (once lazy init is supported).
|
||||||
candidates = set(cfg.get("candidate_steps", [1, 3, 7]))
|
self.candidate_steps = _resolve_candidate_steps(initial_steps, cfg)
|
||||||
|
|
||||||
# Ensure the worker's initial speculative_num_steps is itself a candidate.
|
|
||||||
# Otherwise AdaptiveController.register() would store the worker's pre-built
|
|
||||||
# runtime state under a key that _activate() never queries, leaking that
|
|
||||||
# state's draft attn backend and cuda graph buffers for the process lifetime.
|
|
||||||
if initial_steps not in candidates:
|
|
||||||
log_info_on_rank0(
|
|
||||||
logger,
|
|
||||||
f"Adding initial speculative_num_steps={initial_steps} to "
|
|
||||||
f"candidate_steps={sorted(candidates)} so the pre-built "
|
|
||||||
f"runtime state is reused.",
|
|
||||||
)
|
|
||||||
candidates.add(initial_steps)
|
|
||||||
|
|
||||||
self.candidate_steps = sorted(candidates)
|
|
||||||
assert (
|
assert (
|
||||||
len(self.candidate_steps) >= 2
|
len(self.candidate_steps) >= 2
|
||||||
), "candidate_steps must have at least 2 distinct values"
|
), "candidate_steps must have at least 2 distinct values"
|
||||||
|
|||||||
@@ -1,26 +1,53 @@
|
|||||||
|
import json
|
||||||
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
|
|
||||||
from sglang.srt.speculative.adaptive_spec_params import AdaptiveSpeculativeParams
|
from sglang.srt.speculative.adaptive_spec_params import (
|
||||||
|
AdaptiveSpeculativeParams,
|
||||||
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=6, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
||||||
|
def _make_params_from_config(self, initial_steps: int, config: dict[str, object]):
|
||||||
|
with tempfile.NamedTemporaryFile("w", suffix=".json") as f:
|
||||||
|
json.dump(config, f)
|
||||||
|
f.flush()
|
||||||
|
return AdaptiveSpeculativeParams(
|
||||||
|
initial_steps=initial_steps, cfg_path=f.name
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_params_loads_config_path(self):
|
||||||
|
with tempfile.NamedTemporaryFile("w", suffix=".json") as f:
|
||||||
|
json.dump(
|
||||||
|
{
|
||||||
|
"candidate_steps": [1, 5],
|
||||||
|
"ema_alpha": 0.75,
|
||||||
|
"warmup_batches": 2,
|
||||||
|
},
|
||||||
|
f,
|
||||||
|
)
|
||||||
|
f.flush()
|
||||||
|
|
||||||
|
params = AdaptiveSpeculativeParams(initial_steps=3, cfg_path=f.name)
|
||||||
|
|
||||||
|
self.assertEqual(params.candidate_steps, [1, 3, 5])
|
||||||
|
self.assertEqual(params.ema_alpha, 0.75)
|
||||||
|
self.assertEqual(params.warmup_batches, 2)
|
||||||
|
|
||||||
def test_initial_steps_added_to_candidates_when_missing(self):
|
def test_initial_steps_added_to_candidates_when_missing(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(2, {"candidate_steps": [1, 3, 7]})
|
||||||
initial_steps=2,
|
|
||||||
config={"candidate_steps": [1, 3, 7]},
|
|
||||||
)
|
|
||||||
|
|
||||||
self.assertEqual(params.candidate_steps, [1, 2, 3, 7])
|
self.assertEqual(params.candidate_steps, [1, 2, 3, 7])
|
||||||
self.assertEqual(params.current_steps, 2)
|
self.assertEqual(params.current_steps, 2)
|
||||||
self.assertEqual(params.ema_accept_len, 1.0)
|
self.assertEqual(params.ema_accept_len, 1.0)
|
||||||
|
|
||||||
def test_update_respects_warmup_and_interval(self):
|
def test_update_respects_warmup_and_interval(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=3,
|
3,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 1,
|
"warmup_batches": 1,
|
||||||
@@ -38,9 +65,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.current_steps, 1)
|
self.assertEqual(params.current_steps, 1)
|
||||||
|
|
||||||
def test_empty_batches_do_not_consume_warmup_or_shift_steps(self):
|
def test_empty_batches_do_not_consume_warmup_or_shift_steps(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=3,
|
3,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 1,
|
"warmup_batches": 1,
|
||||||
@@ -59,9 +86,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.current_steps, 1)
|
self.assertEqual(params.current_steps, 1)
|
||||||
|
|
||||||
def test_update_scales_up_across_candidates(self):
|
def test_update_scales_up_across_candidates(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=1,
|
1,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 0,
|
"warmup_batches": 0,
|
||||||
@@ -77,9 +104,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.current_steps, 7)
|
self.assertEqual(params.current_steps, 7)
|
||||||
|
|
||||||
def test_update_can_scale_down_across_candidates_in_one_recompute(self):
|
def test_update_can_scale_down_across_candidates_in_one_recompute(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=7,
|
7,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 0,
|
"warmup_batches": 0,
|
||||||
@@ -91,9 +118,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.current_steps, 1)
|
self.assertEqual(params.current_steps, 1)
|
||||||
|
|
||||||
def test_exact_rise_threshold_does_not_upshift(self):
|
def test_exact_rise_threshold_does_not_upshift(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=3,
|
3,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 0,
|
"warmup_batches": 0,
|
||||||
@@ -110,9 +137,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.current_steps, 7)
|
self.assertEqual(params.current_steps, 7)
|
||||||
|
|
||||||
def test_exact_drop_threshold_does_downshift(self):
|
def test_exact_drop_threshold_does_downshift(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=3,
|
3,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 0,
|
"warmup_batches": 0,
|
||||||
@@ -127,9 +154,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.ema_accept_len, 0.5)
|
self.assertEqual(params.ema_accept_len, 0.5)
|
||||||
|
|
||||||
def test_hysteresis_can_prevent_premature_upshift(self):
|
def test_hysteresis_can_prevent_premature_upshift(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=3,
|
3,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 0,
|
"warmup_batches": 0,
|
||||||
@@ -145,9 +172,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.current_steps, 7)
|
self.assertEqual(params.current_steps, 7)
|
||||||
|
|
||||||
def test_down_hysteresis_can_prevent_premature_downshift(self):
|
def test_down_hysteresis_can_prevent_premature_downshift(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=7,
|
7,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 1.0,
|
"ema_alpha": 1.0,
|
||||||
"warmup_batches": 0,
|
"warmup_batches": 0,
|
||||||
@@ -163,9 +190,9 @@ class TestAdaptiveSpeculativeParams(unittest.TestCase):
|
|||||||
self.assertEqual(params.current_steps, 3)
|
self.assertEqual(params.current_steps, 3)
|
||||||
|
|
||||||
def test_multi_batch_sequence_can_ramp_up_then_back_down(self):
|
def test_multi_batch_sequence_can_ramp_up_then_back_down(self):
|
||||||
params = AdaptiveSpeculativeParams(
|
params = self._make_params_from_config(
|
||||||
initial_steps=3,
|
3,
|
||||||
config={
|
{
|
||||||
"candidate_steps": [1, 3, 7],
|
"candidate_steps": [1, 3, 7],
|
||||||
"ema_alpha": 0.5,
|
"ema_alpha": 0.5,
|
||||||
"warmup_batches": 0,
|
"warmup_batches": 0,
|
||||||
|
|||||||
Reference in New Issue
Block a user