Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
29 commits
Select commit Hold shift + click to select a range
0071867
[fix] fix step0 dsd support for dspkv4-dspark
EanWang211123 Jul 6, 2026
4b48898
[Spec Decode] Harden DSpark metadata and TP sampling state
voipmonitor Jul 10, 2026
357fa29
[Bugfix][Spec Decode] Mask cache-restored tokens out of DFlash draft …
giorgiopiatti-caffeinated Jul 7, 2026
ed5f7c8
Prefer FlashAttn over FlashInfer for SM100f non-causal attention
mgoin Jul 9, 2026
4b248b0
[Spec Decode] Never full-graph-capture non-causal FlashInfer draft at…
mgoin Jul 9, 2026
1d13bf6
fix(spec decode): preserve explicit zero adaptive depth
voipmonitor Jul 17, 2026
7916aaf
spec_decode: retain DFlash CUDA graph backbone outputs
voipmonitor Jul 10, 2026
7a8830e
test(spec decode): use current dynamic-depth predicate
voipmonitor Jul 17, 2026
3571604
fix(dflash): serialize overlapping block-table shift loads
voipmonitor Jul 17, 2026
5d06d3d
style: normalize DSpark correctness tests
voipmonitor Jul 17, 2026
7da3d41
[Spec Decode] DSpark capacity reallocation with varlen full-CG verifi…
LucasWilkinson Jul 8, 2026
1e11f2d
tests: adapt DSpark capacity fixtures to fork InputBatch
voipmonitor Jul 10, 2026
77f3a30
spec_decode: report DSpark capacity verification mode
voipmonitor Jul 10, 2026
f0e7c35
spec_decode: pass draft cache state through capacity warmup
voipmonitor Jul 10, 2026
eedbe48
spec_decode: keep varlen DSpark verification on full graphs
voipmonitor Jul 10, 2026
c12769f
spec_decode: add load-aware DSpark physical depth control
voipmonitor Jul 10, 2026
2f28fde
fix(dspark): harden capacity verification edge paths
voipmonitor Jul 17, 2026
b4f6e92
spec_decode: gate DSpark capacity below load knee
voipmonitor Jul 12, 2026
313a77a
fix: canonicalize DSpark capacity across TP ranks
voipmonitor Jul 12, 2026
fed18d4
Merge DSpark correctness prerequisite for capacity validation
voipmonitor Jul 17, 2026
43a0a20
fix(dspark): preserve default draft generation width
voipmonitor Jul 17, 2026
caa795e
fix(dflash): fail closed on partial restored KV blocks
voipmonitor Jul 17, 2026
7c0d639
Merge branch 'codex/ff-dspark-core-canonical-20260717' into codex/ff-…
voipmonitor Jul 17, 2026
edc4976
fix(dspark): harden capacity opt-in and calibration
voipmonitor Jul 17, 2026
21c08ff
fix(indexer): remove duplicate RoPE quant helper
voipmonitor Jul 17, 2026
4cd7055
Merge remote-tracking branch 'origin/codex/ff-sparse-indexer-dedup-20…
voipmonitor Jul 17, 2026
c8bff14
style: format touched DSpark helpers
voipmonitor Jul 17, 2026
8fbc196
test(indexer): align B12X fixtures with FF contracts
voipmonitor Jul 17, 2026
8c55386
Merge remote-tracking branch 'lil/codex/ff-sparse-indexer-dedup-20260…
voipmonitor Jul 17, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
165 changes: 165 additions & 0 deletions benchmarks/profile_dspark_sps_curve.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,165 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""Profile the engine step-rate curve for the DSpark prefix scheduler.

Times the captured FULL cudagraph replays of the target verification step at
every captured batch token count and emits ``dspark_sps_curve`` breakpoints,
one per capture size. The scheduler linearly interpolates between
breakpoints, which amortizes cudagraph padding smoothly instead of
concentrating it into thresholds at capture-size boundaries.

Example:
python benchmarks/profile_dspark_sps_curve.py <target-model> \\
--speculative-config '{"method": "dspark", "model": "...", ...}' \\
--engine-args '{"tensor_parallel_size": 4, "max_num_seqs": 32}'

Paste the printed ``dspark_sps_curve`` entry into --speculative-config.

Caveats: replays run on whatever (dummy) buffer state capture left behind, so
data-dependent kernels (e.g. MoE routing) may be timed on unrepresentative
inputs, and per-step CPU/draft overhead is modeled only through the constant
``--overhead-ms``. Only the curve's shape matters to the scheduler.
"""

import argparse
import json


def _time_fullgraph_replays(worker, iters: int, warmup: int) -> dict[int, float]:
"""Worker-side: time FULL graph replay per batch token count (ms/step).

Runs on every TP rank via collective_rpc so the collectives captured in
the graphs stay matched; every rank replays the same descs in the same
sorted order. Before timing each descriptor the input buffers are
refreshed into the same coherent dummy state capture used, so replays
never read stale metadata.
"""
import torch

from vllm.v1.worker.gpu.cudagraph_utils import prepare_inputs_to_capture

runner = worker.model_runner
mgr = runner.cudagraph_manager
assert mgr is not None and mgr.graphs, (
"No FULL cudagraphs captured; run with a cudagraph_mode that captures "
"FULL decode graphs."
)
# Prefer varlen spec-decode descs; fall back to all captured graphs.
descs = [d for d in mgr.graphs if d.max_req_tokens is not None]
if not descs:
descs = list(mgr.graphs.keys())
# One desc per token count: the largest request count is the most
# representative shape under load.
by_tokens: dict[int, object] = {}
for d in descs:
cur = by_tokens.get(d.num_tokens)
if cur is None or (d.num_reqs or 0) > (cur.num_reqs or 0):
by_tokens[d.num_tokens] = d

results: dict[int, float] = {}
for num_tokens in sorted(by_tokens):
desc = by_tokens[num_tokens]
num_reqs = desc.num_reqs or min(num_tokens, mgr.max_num_reqs)
prepare_inputs_to_capture(
num_reqs,
num_tokens,
runner.model_state,
runner.input_buffers,
runner.block_tables,
runner.attn_groups,
runner.kv_cache_config,
max_req_tokens=desc.max_req_tokens,
)
graph = mgr.graphs[desc]
for _ in range(warmup):
graph.replay()
torch.accelerator.synchronize()
start = torch.Event(enable_timing=True)
end = torch.Event(enable_timing=True)
start.record()
for _ in range(iters):
graph.replay()
end.record()
torch.accelerator.synchronize()
results[num_tokens] = start.elapsed_time(end) / iters
return results


def curve_breakpoints(
ms_per_step: dict[int, float], overhead_ms: float
) -> list[list[float]]:
"""Convert per-capture-size step times into ``dspark_sps_curve``
breakpoints, one per capture size. The scheduler's table linearly
interpolates between them (and clamps at the ends)."""
return [
[size, round(1000.0 / (ms_per_step[size] + overhead_ms), 3)]
for size in sorted(ms_per_step)
]
Comment thread
voipmonitor marked this conversation as resolved.


def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("model", help="Target model (path or HF id)")
parser.add_argument(
"--speculative-config",
required=True,
help="JSON speculative config (same value you pass to vllm serve)",
)
parser.add_argument(
"--engine-args",
default="{}",
help="JSON dict of extra vllm.LLM kwargs "
'(e.g. \'{"tensor_parallel_size": 4, "max_num_seqs": 32}\')',
)
parser.add_argument("--iters", type=int, default=50)
parser.add_argument("--warmup", type=int, default=5)
parser.add_argument(
"--overhead-ms",
type=float,
default=0.0,
help="Constant per-step overhead (draft forward, sampling, CPU gap) "
"added to every measured step time before converting to a rate.",
)
parser.add_argument("--output", help="Write the curve JSON to this file")
args = parser.parse_args()
if args.iters <= 0:
parser.error("--iters must be greater than zero")
if args.warmup < 0:
parser.error("--warmup must be non-negative")
if args.overhead_ms < 0:
parser.error("--overhead-ms must be non-negative")

# The timing callable is shipped to the workers via collective_rpc, which
# requires the pickle fallback. Local profiling tool, trusted input.
import os

os.environ.setdefault("VLLM_ALLOW_INSECURE_SERIALIZATION", "1")

from vllm import LLM

llm = LLM(
model=args.model,
speculative_config=json.loads(args.speculative_config),
**json.loads(args.engine_args),
)
per_rank = llm.collective_rpc(
_time_fullgraph_replays, kwargs={"iters": args.iters, "warmup": args.warmup}
)
ms_per_step = per_rank[0]

print("\nMeasured FULL-graph step times (rank 0):")
for size in sorted(ms_per_step):
print(f" B={size:5d} tokens: {ms_per_step[size]:8.3f} ms/step")

curve = curve_breakpoints(ms_per_step, args.overhead_ms)
entry = {"dspark_sps_curve": curve}
print("\nAdd to --speculative-config:")
print(json.dumps(entry))
if args.output:
with open(args.output, "w") as f:
json.dump(entry, f, indent=2)
print(f"\nWritten to {args.output}")


if __name__ == "__main__":
main()
37 changes: 37 additions & 0 deletions tests/engine/test_arg_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -225,6 +225,43 @@ def test_jit_monitor_mode_arg(mode):
assert engine_args.create_observability_config().jit_monitor_mode == mode


def test_dspark_capacity_verification_mode_arg():
parser = EngineArgs.add_cli_args(FlexibleArgumentParser())
args = parser.parse_args(
[
"--spec-method",
"ngram",
"--spec-tokens",
"1",
"--dspark-capacity-verification-mode",
"mask",
]
)

engine_args = EngineArgs.from_cli_args(args)
assert engine_args.dspark_capacity_verification_mode == "mask"
speculative_config = engine_args.create_speculative_config(None, None)
assert speculative_config is not None
assert speculative_config.dspark_capacity_verification_mode == "mask"


def test_dspark_capacity_verification_mode_conflicts_with_speculative_config():
parser = EngineArgs.add_cli_args(FlexibleArgumentParser())
args = parser.parse_args(
[
"--speculative-config",
'{"method":"ngram","num_speculative_tokens":1,'
'"dspark_capacity_verification_mode":"mask"}',
"--dspark-capacity-verification-mode",
"varlen",
]
)

engine_args = EngineArgs.from_cli_args(args)
with pytest.raises(ValueError, match="dspark_capacity_verification_mode"):
engine_args.create_speculative_config(None, None)


def test_hf_token_get_kwargs():
kwargs = get_kwargs(ModelConfig)["hf_token"]

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,22 @@
model_name: "deepseek-ai/DeepSeek-V4-Flash-DSpark"
accuracy_threshold: 0.92
num_questions: 1319
num_fewshot: 5
startup_max_wait_seconds: 1800
server_args: >-
--tokenizer-mode deepseek_v4
--trust-remote-code
--dtype bfloat16
--max-model-len 8192
--tensor-parallel-size 4
--enable-expert-parallel
--block-size 256
--gpu-memory-utilization 0.5
--kv-cache-dtype fp8
--max-num-batched-tokens 16384
--max-num-seqs 32
--speculative-config '{"method":"dspark",
"model":"deepseek-ai/DeepSeek-V4-Flash-DSpark",
"attention_backend":"FLASH_ATTN","num_speculative_tokens":7,
"draft_sample_method":"probabilistic","dspark_confidence_threshold":0.0,
"dspark_budget_frac":0.5,"dspark_capacity_verification_mode":"varlen"}'
89 changes: 51 additions & 38 deletions tests/model_executor/layers/test_sparse_attn_indexer_b12x.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,6 +153,44 @@ def build_paged_mqa_schedule_metadata(seq_lens, block_size, num_sms, *, out):
monkeypatch.setitem(sys.modules, "b12x.integration", integration_mod)


def _install_fake_b12x_dcp_merge(monkeypatch, run_row_topk, *, world_size: int):
tiled_topk_mod = types.ModuleType("b12x.attention.indexer.tiled_topk")
tiled_topk_mod.run_row_topk = run_row_topk
monkeypatch.setitem(
sys.modules,
"b12x.attention.indexer.tiled_topk",
tiled_topk_mod,
)

class FakeDCPGroup:
def __init__(self) -> None:
self.world_size = world_size

def all_gather(self, tensor, dim):
assert dim == 0
return torch.cat([tensor.clone() for _ in range(world_size)], dim=dim)

import vllm.distributed.parallel_state as parallel_state
import vllm.v1.attention.backends.mla.sparse_utils as sparse_utils

monkeypatch.setattr(parallel_state, "get_dcp_group", lambda: FakeDCPGroup())
monkeypatch.setattr(
sparse_utils,
"triton_convert_dcp_local_topk_to_global",
lambda *args, **kwargs: None,
)

def gather_topk_ids_by_position(candidate_ids, positions, out):
gathered = torch.gather(candidate_ids, 1, positions.to(torch.int64))
out.copy_(gathered.to(out.dtype))

monkeypatch.setattr(
sparse_utils,
"triton_gather_topk_ids_by_position",
gather_topk_ids_by_position,
)


@pytest.mark.parametrize(
"page_stride0",
[
Expand Down Expand Up @@ -560,9 +598,6 @@ def fake_merge(**kwargs):
schedule_metadata=None,
active_width=None,
),
dcp_world_size=2,
dcp_rank=1,
cp_interleave_size=16,
)
layer_name = "layers.0.attn"
metadata_key = indexer_mod._resolve_layer_name(layer_name)
Expand Down Expand Up @@ -599,6 +634,9 @@ def fake_merge(**kwargs):
use_fp4_cache=False,
use_b12x_sparse_indexer=True,
topk_scores_buffer=topk_scores_buffer,
dcp_world_size=2,
dcp_rank=1,
cp_kv_cache_interleave_size=16,
)

assert result is topk_indices_buffer
Expand All @@ -625,7 +663,6 @@ def fake_merge(**kwargs):


def test_b12x_dcp_merge_passes_contiguous_scores_to_topk(monkeypatch):
tiled_topk_mod = types.ModuleType("b12x.attention.indexer.tiled_topk")
run_row_topk_calls: list[tuple[bool, tuple[int, ...]]] = []

def run_row_topk(*, row_logits, lengths, topk, output_values, output_indices):
Expand All @@ -638,39 +675,7 @@ def run_row_topk(*, row_logits, lengths, topk, output_values, output_indices):
output_indices.copy_(positions)
output_values.copy_(row_logits[:, :topk])

tiled_topk_mod.run_row_topk = run_row_topk
monkeypatch.setitem(
sys.modules,
"b12x.attention.indexer.tiled_topk",
tiled_topk_mod,
)

class FakeDCPGroup:
world_size = 2

def all_gather(self, tensor, dim):
assert dim == 0
return torch.cat([tensor, tensor.clone()], dim=dim)

import vllm.distributed.parallel_state as parallel_state
import vllm.v1.attention.backends.mla.sparse_utils as sparse_utils

monkeypatch.setattr(parallel_state, "get_dcp_group", lambda: FakeDCPGroup())
monkeypatch.setattr(
sparse_utils,
"triton_convert_dcp_local_topk_to_global",
lambda *args, **kwargs: None,
)

def gather_topk_ids_by_position(candidate_ids, positions, out):
gathered = torch.gather(candidate_ids, 1, positions.to(torch.int64))
out.copy_(gathered.to(out.dtype))

monkeypatch.setattr(
sparse_utils,
"triton_gather_topk_ids_by_position",
gather_topk_ids_by_position,
)
_install_fake_b12x_dcp_merge(monkeypatch, run_row_topk, world_size=2)
workspace_manager = _FakeWorkspaceManager()
monkeypatch.setattr(
indexer_mod, "current_workspace_manager", lambda: workspace_manager
Expand All @@ -695,6 +700,14 @@ def gather_topk_ids_by_position(candidate_ids, positions, out):


def test_b12x_dcp_merge_warmup_reserves_workspace(monkeypatch):
def run_row_topk(*, row_logits, lengths, topk, output_values, output_indices):
positions = torch.arange(topk, dtype=output_indices.dtype).expand_as(
output_indices
)
output_indices.copy_(positions)
output_values.copy_(row_logits[:, :topk])

_install_fake_b12x_dcp_merge(monkeypatch, run_row_topk, world_size=4)
workspace_manager = _FakeWorkspaceManager()
monkeypatch.setattr(
indexer_mod, "current_workspace_manager", lambda: workspace_manager
Expand Down Expand Up @@ -1054,7 +1067,7 @@ def test_b12x_schedule_metadata_uses_canonical_indexer_import(monkeypatch):
)
builder = object.__new__(mla_indexer_mod.DeepseekV32IndexerMetadataBuilder)
builder.scheduler_metadata_buffer = torch.zeros((5, 2), dtype=torch.int32)
builder.storage_block_size = 64
builder.kv_cache_spec = types.SimpleNamespace(storage_block_size=64)
builder.num_sms = 4

seq_lens = torch.tensor([64, 128], dtype=torch.int32)
Expand Down
Loading
Loading