diff --git a/.agents/skills/profile-dsv4-serving-strace/SKILL.md b/.agents/skills/profile-dsv4-serving-strace/SKILL.md index ae1bb7ec..5d0d8ec3 100644 --- a/.agents/skills/profile-dsv4-serving-strace/SKILL.md +++ b/.agents/skills/profile-dsv4-serving-strace/SKILL.md @@ -79,6 +79,8 @@ summarizes Effective when reprocessing an older log that contains complete devic Serving child processes can write several complete `[STRACE]` records on one physical log line. The analyzer splits at every marker before calling Simpler's built-in `parse_spans`, `group_invocations`, `to_chrome_trace`, and `_round_metrics` APIs. +It detects both the legacy four-submission decode layout and the current two-submission +main-plus-verify / MTP layout from the serving kernel spans and final MTP counters. ## Validate before reporting success @@ -91,8 +93,8 @@ Require all of the following: 4. `server.log` contains successful SA profiler start/stop, final request, and MTP acceptance lines. 5. Eight distinct `[chip_process pid=... dev=...] ready` mappings are present. -6. `serving-trace/trace.json` contains non-empty `traceEvents`, including framework and - all four DSV4 prefill/decode kernel spans. +6. `serving-trace/trace.json` contains non-empty `traceEvents`, including framework spans + and all prefill/decode kernel spans required by the detected decode layout. 7. `server.log` contains host `[STRACE]` records and no `clk=dev` records. 8. `simpler-swimlane.json`, `strace-8lane.json`, and `strace-8lane-host-clock.json` contain non-empty `traceEvents`. diff --git a/.agents/skills/profile-dsv4-serving-strace/scripts/analyze_profile.py b/.agents/skills/profile-dsv4-serving-strace/scripts/analyze_profile.py index f63e4e5b..603ea9eb 100644 --- a/.agents/skills/profile-dsv4-serving-strace/scripts/analyze_profile.py +++ b/.agents/skills/profile-dsv4-serving-strace/scripts/analyze_profile.py @@ -87,6 +87,32 @@ def main() -> None: server_log = server_log_path.read_text(encoding="utf-8", errors="replace") serving_trace = json.loads(serving_trace_path.read_text(encoding="utf-8")) + serving_events = serving_trace["traceEvents"] + kernel_durations_ms: dict[str, list[float]] = defaultdict(list) + for event in serving_events: + if ( + event.get("ph") == "X" + and event.get("cat") == "kernel" + and not event.get("name", "").endswith(".worker_run") + ): + kernel_durations_ms[event["args"]["kernel"]].append(event["dur"] / 1000.0) + + required_kernels = {"deepseek_v4_prefill", "deepseek_v4_mtp_prefill"} + if "deepseek_v4_decode_mtp_fused" in kernel_durations_ms: + decode_layout = "two_l2" + decode_width = 2 + required_kernels.add("deepseek_v4_decode_mtp_fused") + elif {"deepseek_v4_decode", "deepseek_v4_mtp_decode"}.issubset(kernel_durations_ms): + decode_layout = "split" + decode_width = 4 + required_kernels.update({"deepseek_v4_decode", "deepseek_v4_mtp_decode"}) + else: + raise RuntimeError( + "serving trace has neither the fused nor split DeepSeek V4 decode kernel layout" + ) + if not required_kernels.issubset(kernel_durations_ms): + raise RuntimeError(f"missing serving kernel spans: {required_kernels - kernel_durations_ms.keys()}") + split_log = server_log.replace("[STRACE]", "\n[STRACE]") spans = list(parse_spans(split_log.splitlines())) invocations = group_invocations(spans) @@ -130,11 +156,11 @@ def main() -> None: invocation_ids = sorted(hid_by_inv) if ( invocation_ids != list(range(1, invocation_ids[-1] + 1)) - or invocation_ids[-1] < 8 - or (invocation_ids[-1] - 4) % 4 + or invocation_ids[-1] < 4 + decode_width + or (invocation_ids[-1] - 4) % decode_width ): raise RuntimeError(f"unexpected invocation ids: {invocation_ids}") - decode_step_count = (invocation_ids[-1] - 4) // 4 + decode_step_count = (invocation_ids[-1] - 4) // decode_width if len(request_invocations) != len(devices) * invocation_ids[-1]: raise RuntimeError( f"incomplete rank data: got {len(request_invocations)} invocations, " @@ -177,10 +203,13 @@ def rank_phase_row(pid: int, main_ids: list[int], mtp_ids: list[int]) -> dict: decode_rows = [] critical_metric = "total_effective_us" if device_effective_available else "total_host_us" for step in range(decode_step_count): - base = 5 + step * 4 + base = 5 + step * decode_width per_rank = [] for pid in pids: - per_rank.append(rank_phase_row(pid, [base, base + 1], [base + 2, base + 3])) + if decode_layout == "two_l2": + per_rank.append(rank_phase_row(pid, [base], [base + 1])) + else: + per_rank.append(rank_phase_row(pid, [base, base + 1], [base + 2, base + 3])) critical = max(per_rank, key=lambda row: row[critical_metric]) decode_row = { "step": step + 1, @@ -200,25 +229,6 @@ def rank_phase_row(pid: int, main_ids: list[int], mtp_ids: list[int]) -> dict: ) decode_rows.append(decode_row) - serving_events = serving_trace["traceEvents"] - kernel_durations_ms: dict[str, list[float]] = defaultdict(list) - for event in serving_events: - if ( - event.get("ph") == "X" - and event.get("cat") == "kernel" - and not event.get("name", "").endswith(".worker_run") - ): - kernel_durations_ms[event["args"]["kernel"]].append(event["dur"] / 1000.0) - - required_kernels = { - "deepseek_v4_prefill", - "deepseek_v4_mtp_prefill", - "deepseek_v4_decode", - "deepseek_v4_mtp_decode", - } - if not required_kernels.issubset(kernel_durations_ms): - raise RuntimeError(f"missing serving kernel spans: {required_kernels - kernel_durations_ms.keys()}") - request_ms = next( ( event["dur"] / 1000.0 @@ -287,6 +297,7 @@ def kernel_summary(name: str) -> dict: summary = { "run_id": args.run_id, "devices": devices, + "decode_layout": decode_layout, "request": { "prompt_tokens": int(completion_match.group(1)), "completion_tokens": completion_tokens, @@ -325,10 +336,29 @@ def kernel_summary(name: str) -> dict: (artifact_dir / "profile-summary.json").write_text(json.dumps(summary, indent=2) + "\n") prefill = summary["simpler"]["prefill_critical"] - main_serving = summary["serving_kernel_ms"]["deepseek_v4_decode"]["steady_after_first"] - mtp_serving = summary["serving_kernel_ms"]["deepseek_v4_mtp_decode"]["steady_after_first"] - if main_serving is None or mtp_serving is None: - raise RuntimeError("need more than one decode kernel span for steady statistics") + if decode_layout == "two_l2": + fused_serving = summary["serving_kernel_ms"]["deepseek_v4_decode_mtp_fused"][ + "steady_after_first" + ] + if fused_serving is None: + raise RuntimeError("need more than one fused decode kernel span for steady statistics") + serving_decode_lines = ( + f"- Combined decode steady mean (steps 2-{decode_step_count}): " + f"{fused_serving['mean']:.3f} ms/iteration" + ) + main_phase_label = "Main+verify" + else: + main_serving = summary["serving_kernel_ms"]["deepseek_v4_decode"]["steady_after_first"] + mtp_serving = summary["serving_kernel_ms"]["deepseek_v4_mtp_decode"]["steady_after_first"] + if main_serving is None or mtp_serving is None: + raise RuntimeError("need more than one decode kernel span for steady statistics") + serving_decode_lines = ( + f"- Decode main steady mean (steps 2-{decode_step_count}): " + f"{main_serving['mean']:.3f} ms/iteration\n" + f"- Decode MTP steady mean (steps 2-{decode_step_count}): " + f"{mtp_serving['mean']:.3f} ms/iteration" + ) + main_phase_label = "Main" critical_rank_counts = Counter(row["critical_device"] for row in decode_rows) critical_rank_summary = ", ".join( f"device {device}: {count}/{decode_step_count}" @@ -366,9 +396,9 @@ def kernel_summary(name: str) -> dict: Sched windows. - Prefill critical rank: device {prefill["device"]}, main={prefill["main_effective_us"] / 1000:.3f} ms, MTP={prefill["mtp_effective_us"] / 1000:.3f} ms, total={prefill["total_effective_us"] / 1000:.3f} ms -- Decode steady critical rank mean: main={steady_effective["critical_main_effective_us"]["mean"] / 1000:.3f} ms, MTP={steady_effective["critical_mtp_effective_us"]["mean"] / 1000:.3f} ms, total={steady_effective["critical_total_effective_us"]["mean"] / 1000:.3f} ms/iteration +- Decode steady critical rank mean: {main_phase_label}={steady_effective["critical_main_effective_us"]["mean"] / 1000:.3f} ms, MTP={steady_effective["critical_mtp_effective_us"]["mean"] / 1000:.3f} ms, total={steady_effective["critical_total_effective_us"]["mean"] / 1000:.3f} ms/iteration -| Decode iteration | Critical device | Main Effective (ms) | MTP Effective (ms) | Total Effective (ms) | +| Decode iteration | Critical device | {main_phase_label} Effective (ms) | MTP Effective (ms) | Total Effective (ms) | | ---: | ---: | ---: | ---: | ---: | {effective_decode_table} """ @@ -384,16 +414,15 @@ def kernel_summary(name: str) -> dict: - Prefill main kernel span: {kernel_durations_ms["deepseek_v4_prefill"][0]:.3f} ms - Prefill MTP kernel span: {kernel_durations_ms["deepseek_v4_mtp_prefill"][0]:.3f} ms -- Decode main steady mean (steps 2-{decode_step_count}): {main_serving["mean"]:.3f} ms/iteration -- Decode MTP steady mean (steps 2-{decode_step_count}): {mtp_serving["mean"]:.3f} ms/iteration +{serving_decode_lines} ## Simpler Host STRACE - Prefill critical rank: device {prefill["device"]}, main={prefill["main_host_us"] / 1000:.3f} ms, MTP={prefill["mtp_host_us"] / 1000:.3f} ms, total={prefill["total_host_us"] / 1000:.3f} ms -- Decode steady critical rank mean (steps 2-{decode_step_count}): main={host_stats["critical_main_host_us"]["mean"] / 1000:.3f} ms, MTP={host_stats["critical_mtp_host_us"]["mean"] / 1000:.3f} ms, total={host_stats["critical_total_host_us"]["mean"] / 1000:.3f} ms/iteration +- Decode steady critical rank mean (steps 2-{decode_step_count}): {main_phase_label}={host_stats["critical_main_host_us"]["mean"] / 1000:.3f} ms, MTP={host_stats["critical_mtp_host_us"]["mean"] / 1000:.3f} ms, total={host_stats["critical_total_host_us"]["mean"] / 1000:.3f} ms/iteration - Critical-rank counts across decode: {critical_rank_summary} -| Decode iteration | Critical device | Main host (ms) | MTP host (ms) | Total host (ms) | +| Decode iteration | Critical device | {main_phase_label} host (ms) | MTP host (ms) | Total host (ms) | | ---: | ---: | ---: | ---: | ---: | {host_decode_table} diff --git a/.agents/skills/profile-dsv4-serving-strace/scripts/render_8lane.py b/.agents/skills/profile-dsv4-serving-strace/scripts/render_8lane.py index 7040ccd3..89aa9ed0 100644 --- a/.agents/skills/profile-dsv4-serving-strace/scripts/render_8lane.py +++ b/.agents/skills/profile-dsv4-serving-strace/scripts/render_8lane.py @@ -18,9 +18,10 @@ PROCESS_NAME_RE = re.compile(r"inv=(?P\d+) \(pid=(?P\d+)\)") DEVICE_READY_RE = re.compile(r"\[chip_process pid=(?P\d+) dev=(?P\d+)\] ready") +MTP_ACCEPTANCE_RE = re.compile(r"MTP acceptance for .* proposed=(?P\d+)") -def callable_label(invocation: int) -> tuple[str, int | None, str]: +def callable_label(invocation: int, decode_width: int) -> tuple[str, int | None, str]: prefill = { 1: ("prefill.main", None, "rail_response"), 2: ("prefill.main.lm_head", None, "rail_animation"), @@ -29,8 +30,15 @@ def callable_label(invocation: int) -> tuple[str, int | None, str]: } if invocation in prefill: return prefill[invocation] - step = (invocation - 5) // 4 + 1 - phase = (invocation - 5) % 4 + step = (invocation - 5) // decode_width + 1 + phase = (invocation - 5) % decode_width + if decode_width == 2: + decode = { + 0: ("decode.main+verify", "good"), + 1: ("decode.mtp", "cq_build_running"), + } + label, color = decode[phase] + return label, step, color decode = { 0: ("decode.main", "good"), 1: ("decode.main.lm_head", "rail_animation"), @@ -91,9 +99,10 @@ def main() -> None: source = json.loads(args.input.read_text()) source_events = source["traceEvents"] if isinstance(source, dict) else source + server_log = args.server_log.read_text(errors="replace") pid_to_device = { int(match.group("pid")): int(match.group("device")) - for match in DEVICE_READY_RE.finditer(args.server_log.read_text(errors="replace")) + for match in DEVICE_READY_RE.finditer(server_log) } devices = sorted(pid_to_device.values()) if len(devices) != 8 or len(set(devices)) != 8: @@ -119,6 +128,19 @@ def main() -> None: if virtual_pid in virtual_processes and event.get("ph") == "X": grouped.setdefault(int(virtual_pid), []).append(event) + invocation_ids = sorted({invocation for _device, invocation in virtual_processes.values()}) + if not invocation_ids or invocation_ids != list(range(1, invocation_ids[-1] + 1)): + raise ValueError(f"unexpected invocation ids: {invocation_ids}") + decode_invocations = invocation_ids[-1] - 4 + acceptance_matches = list(MTP_ACCEPTANCE_RE.finditer(server_log)) + proposed_steps = int(acceptance_matches[-1].group("steps")) if acceptance_matches else 0 + if proposed_steps and decode_invocations == proposed_steps * 2: + decode_width = 2 + elif decode_invocations % 4 == 0: + decode_width = 4 + else: + raise ValueError(f"cannot infer decode invocation width from {invocation_ids[-1]} invocations") + roots = [ event for events in grouped.values() @@ -169,7 +191,7 @@ def main() -> None: runner = one_event(events, "simpler_run.runner_run") validate = one_event(events, "simpler_run.validate") device_wall = one_event(events, "simpler_run.runner_run.device_wall") - label, step, color = callable_label(invocation) + label, step, color = callable_label(invocation, decode_width) event_name = label if step is None else f"D{step:02d} {label}" has_device_trace = device_wall is not None event_args = { diff --git a/pypto-lib b/pypto-lib index 95206fc4..fd91568d 160000 --- a/pypto-lib +++ b/pypto-lib @@ -1 +1 @@ -Subproject commit 95206fc4017577c7c975b67f128478eea946f060 +Subproject commit fd91568d7fff6ed956978a2070529c4d6cc21a5b diff --git a/pypto_serving/model/deepseek/npu_executor.py b/pypto_serving/model/deepseek/npu_executor.py index 79de74c0..9a4ee2e9 100644 --- a/pypto_serving/model/deepseek/npu_executor.py +++ b/pypto_serving/model/deepseek/npu_executor.py @@ -102,12 +102,14 @@ "decode_attention_hca", "decode_attention_swa", "decode_fwd", + "decode_fwd_mtp", "decode_input_pack", "decode_indexer", "decode_indexer_compressor", "decode_layer", "decode_metadata_device", "decode_mtp", + "decode_mtp_verify", "lookup_embedding", "decode_sparse_attn", "decode_sparse_attn_csa", @@ -438,9 +440,21 @@ def _compile_model(self, model: RuntimeModel) -> DeepSeekV4CompiledKernels: self._prefill_dummy_args(model, layout, modules["config"]), ) decode = self._compile_l3_callable( - "deepseek_v4_decode", - modules["decode_fwd"].l3_decode_fwd, - self._decode_dummy_args(model, layout, modules["config"]), + "deepseek_v4_decode_mtp_fused" if self._enable_mtp else "deepseek_v4_decode", + ( + modules["decode_fwd_mtp"].l3_decode_fwd_mtp + if self._enable_mtp + else modules["decode_fwd"].l3_decode_fwd + ), + ( + self._fused_mtp_dummy_args( + modules, + model=model, + layout=layout, + ) + if self._enable_mtp + else self._decode_dummy_args(model, layout, modules["config"]) + ), ) if self._enable_mtp: mtp_prefill = self._compile_l3_callable( @@ -453,16 +467,6 @@ def _compile_model(self, model: RuntimeModel) -> DeepSeekV4CompiledKernels: num_tokens=layout.prefill_seq, ), ) - mtp_decode = self._compile_l3_callable( - "deepseek_v4_mtp_decode", - modules["decode_mtp"].l3_mtp_decode_layer, - self._mtp_dummy_args( - modules["decode_mtp"], - model=model, - layout=layout, - num_tokens=layout.decode_tokens, - ), - ) freqs_cos, freqs_sin = self._build_rope_tables(modules["rope_tables"], modules["config"]) return DeepSeekV4CompiledKernels( @@ -488,6 +492,54 @@ def _compile_model(self, model: RuntimeModel) -> DeepSeekV4CompiledKernels: enable_mtp=self._enable_mtp, ) + def _fused_mtp_dummy_args( + self, + modules: dict[str, object], + *, + model: RuntimeModel, + layout: DeepSeekV4CacheLayout, + ) -> tuple[Any, ...]: + """Build the combined main-decode and MTP-decode compile signature.""" + main_args = self._decode_dummy_args(model, layout, modules["config"]) + mtp_args = self._mtp_dummy_args( + modules["decode_mtp"], + model=model, + layout=layout, + num_tokens=layout.decode_tokens, + ) + with _deepseek_v4_import_context( + self._kernel_dir, + pypto_root=self._kernel_dir.parents[2], + ep=len(self._device_ids), + lm_head_tp=DEEPSEEK_V4_LM_HEAD_TP_SIZE, + moe_shape="decode", + ): + mtp_specs = modules["decode_mtp"].build_tensor_specs( + num_tokens=layout.decode_tokens, + ) + shared_names = { + "embed_weight", + "main_pre_hc_hidden", + "freqs_cos", + "freqs_sin", + "ori_block_table", + "lm_head_weight", + } + tail_token_ids = torch.empty( + (layout.ranks, layout.decode_batch), + dtype=torch.int64, + ) + tail_positions = torch.empty( + (layout.ranks, layout.decode_batch), + dtype=torch.int32, + ) + fused_mtp_args = tuple( + arg + for spec, arg in zip(mtp_specs, mtp_args, strict=True) + if spec.name not in shared_names + ) + return (*main_args, tail_token_ids, tail_positions, *fused_mtp_args) + def _load_kernel_modules(self, layout: DeepSeekV4CacheLayout) -> dict[str, object]: """Import DeepSeekV4 pypto-lib modules with EP fixed to the serving world size.""" pypto_root = self._kernel_dir.parents[2] @@ -522,10 +574,13 @@ def _load_kernel_modules(self, layout: DeepSeekV4CacheLayout) -> dict[str, objec config.DECODE_RECV_MAX = ranks * layout.decode_tokens config.RECV_MAX = config.DECODE_RECV_MAX modules = {"config": config} + decode_module_names = ["decode_layer", "decode_fwd", "lm_head", "rope_tables"] + if self._enable_mtp: + decode_module_names.extend(("decode_mtp", "decode_fwd_mtp")) modules.update( { name: importlib.import_module(name) - for name in ("decode_layer", "decode_fwd", "decode_mtp", "lm_head", "rope_tables") + for name in decode_module_names } ) modules["prefill_layer"] = prefill_layer diff --git a/pypto_serving/model/deepseek/npu_runner.py b/pypto_serving/model/deepseek/npu_runner.py index 6b81001a..f93183f1 100644 --- a/pypto_serving/model/deepseek/npu_runner.py +++ b/pypto_serving/model/deepseek/npu_runner.py @@ -533,6 +533,17 @@ def deepseek_v4_physical_cache_blocks( "logit_row_indices", ) +_FUSED_MTP_SHARED_TENSORS = frozenset( + { + "embed_weight", + "main_pre_hc_hidden", + "freqs_cos", + "freqs_sin", + "ori_block_table", + "lm_head_weight", + } +) + _DECODE_INPUT_TENSOR_FIELDS = ( "input_ids", "position_ids", @@ -1248,6 +1259,8 @@ class _DeepSeekV4MtpSharedBuffers: decode_position_ids: torch.Tensor decode_accepted_counts: torch.Tensor decode_tail_slot_ids: torch.Tensor + decode_tail_token_ids: torch.Tensor + decode_tail_positions: torch.Tensor tail_init_hidden: torch.Tensor decode_kv_cache: torch.Tensor prefill_hidden_out: torch.Tensor @@ -2299,7 +2312,7 @@ def _run_autoregressive_decode(self, model: RuntimeModel, batch: DecodeBatch) -> return DecodeResult(hidden_states=None, logits=logits) def _run_mtp_decode(self, model: RuntimeModel, batch: DecodeBatch) -> DecodeResult: - """Verify request-local MTP drafts and advance the accepted windows.""" + """Verify and advance request-local MTP drafts in one L3 dispatch.""" if not batch.allow_device_greedy_sampling: raise RuntimeError("DeepSeekV4 MTP decode currently requires greedy device sampling") with profile_span("DeepSeekV4ModelRunner.decode.mtp_initialize", cat="executor"): @@ -2307,24 +2320,64 @@ def _run_mtp_decode(self, model: RuntimeModel, batch: DecodeBatch) -> DecodeResu with profile_span("DeepSeekV4ModelRunner.decode.mtp_build_speculative_batch", cat="executor"): draft_token_ids = self._mtp_drafts_for_requests(batch.request_ids) speculative_batch = self._main_speculative_batch(model, batch, draft_token_ids) - output = self._execute_main_decode( - model, - self.prepare_mtp_decode_inputs(model, speculative_batch), - active_seq=self._compiled.layout.decode_seq, - ) - inputs = output.inputs - decode_seq = self._compiled.layout.decode_seq + + with profile_span("DeepSeekV4ModelRunner.decode.prepare_inputs", cat="executor"): + inputs = self._stage_decode_inputs( + self.prepare_mtp_decode_inputs(model, speculative_batch) + ) + active_tokens = self._stage_fused_mtp_metadata(inputs) + layout = self._compiled.layout + decode_buffers = self._require_decode_buffers() + mtp_buffers = self._require_mtp_buffers() + hidden_buffer = self._require_decode_output_buffer(model.config.hidden_size) + pre_hc_hidden_buffer = self._materialize_main_pre_hc_device(model.config.hidden_size) + logits_buffer = self._require_decode_logits_buffer(model.config.vocab_size) + with profile_span( + "DeepSeekV4ModelRunner.decode.prepare_fwd_args", + cat="executor", + args={"actual_tokens": active_tokens}, + ): + main_args = self._decode_fwd_args( + inputs, + pre_hc_hidden_buffer, + hidden_buffer, + logits_buffer, + decode_buffers.sampled_ids, + ) + args = self._fused_mtp_decode_args(main_args, active_tokens) + self._debug_decode_dispatch(inputs, main_args) + try: + with profile_span( + "DeepSeekV4ModelRunner.decode.l3_dispatch", + cat="executor", + args={"actual_tokens": active_tokens, "fused_mtp": True}, + ): + self._run_l3(self._require_decode_callable(), *args) + except RuntimeError as exc: + raise RuntimeError( + "DeepSeekV4 fused main/MTP decode dispatch failed " + f"(actual_batch={inputs.actual_batch}, ranks={inputs.ranks})" + ) from exc + + decode_seq = layout.decode_seq with profile_span("DeepSeekV4ModelRunner.decode.mtp_accept", cat="executor"): main_ids = torch.stack( tuple( - output.sampled_ids[rank, local_row * decode_seq + offset, 0] + decode_buffers.sampled_ids[rank, local_row * decode_seq + offset, 0] for rank, local_row in zip(inputs.ranks, inputs.local_rows, strict=True) for offset in range(decode_seq) ) ).to(torch.long).reshape(inputs.actual_batch, decode_seq) - accepted = accept_mtp_tokens(main_ids, draft_token_ids) + accepted_counts = tuple( + int(mtp_buffers.decode_accepted_counts[rank, local_row].item()) + for rank, local_row in zip(inputs.ranks, inputs.local_rows, strict=True) + ) + accepted = [ + main_ids[index, :accepted_count].detach().cpu().tolist() + for index, accepted_count in enumerate(accepted_counts) + ] self._mtp_proposed_tokens += inputs.actual_batch - self._mtp_accepted_tokens += sum(len(tokens) == decode_seq for tokens in accepted) + self._mtp_accepted_tokens += sum(count == decode_seq for count in accepted_counts) for request_id, tokens in zip(inputs.request_ids, accepted, strict=True): state = self._require_mtp_request_state(request_id) state.proposed_tokens += 1 @@ -2341,18 +2394,26 @@ def _run_mtp_decode(self, model: RuntimeModel, batch: DecodeBatch) -> DecodeResu draft_token_ids.detach().cpu().tolist(), main_ids.detach().cpu().tolist(), ) - # Match the reference accepted_num flow: update the MTP window from - # committed main-model outputs immediately, even after rejection. with profile_span( - "DeepSeekV4ModelRunner.decode.mtp_advance", + "DeepSeekV4ModelRunner.decode.mtp_update_state", cat="executor", - args={"accepted_counts": tuple(len(tokens) for tokens in accepted)}, + args={"accepted_counts": accepted_counts}, ): - self._advance_mtp_drafts( - inputs, - main_ids, - accepted_counts=tuple(len(tokens) for tokens in accepted), - ) + for request_id, rank, local_row in zip( + inputs.request_ids, + inputs.ranks, + inputs.local_rows, + strict=True, + ): + row_end = (local_row + 1) * decode_seq - 1 + state = self._require_mtp_request_state(request_id) + state.draft_token_id = int( + mtp_buffers.decode_sampled_ids[rank, local_row, 0].item() + ) + state.tail_token_id = int(mtp_buffers.decode_input_ids[rank, row_end].item()) + state.tail_position = int( + mtp_buffers.decode_position_ids[rank, row_end].item() + ) return DecodeResult( hidden_states=None, logits=None, @@ -2704,6 +2765,27 @@ def _mtp_decode_args(self) -> tuple[Any, ...]: values = self._mark_resident_args(values, _MTP_DECODE_RESIDENT_POLICY) return self._ordered_layer_args(values, _MTP_DECODE_TENSOR_ORDER) + def _fused_mtp_decode_args( + self, + main_args: tuple[Any, ...], + active_tokens: int, + ) -> tuple[Any, ...]: + """Append the non-shared MTP arguments to the main decode arguments.""" + buffers = self._require_mtp_buffers() + mtp_args = self._mtp_decode_args() + fused_mtp_args = tuple( + arg + for name, arg in zip(_MTP_DECODE_TENSOR_ORDER, mtp_args, strict=True) + if name not in _FUSED_MTP_SHARED_TENSORS + ) + return ( + *main_args, + buffers.decode_tail_token_ids, + buffers.decode_tail_positions, + *fused_mtp_args, + self._int32_scalar(active_tokens), + ) + def _require_mtp_buffers(self) -> _DeepSeekV4MtpSharedBuffers: if self._mtp_buffers is None: raise RuntimeError("DeepSeekV4 MTP shared buffers are not staged") @@ -3007,6 +3089,37 @@ def _stage_mtp_decode_inputs( self._mtp_decode_inputs_initialized = True return active_tokens + def _stage_fused_mtp_metadata(self, inputs: DeepSeekV4PreparedDecodeInputs) -> int: + """Stage request tails consumed by the device-side verifier.""" + buffers = self._require_mtp_buffers() + layout = self._compiled.layout + buffers.decode_tail_slot_ids.fill_(-1) + buffers.decode_tail_token_ids.zero_() + buffers.decode_tail_positions.zero_() + buffers.decode_logit_row_indices.fill_(-1) + for request_id, rank, local_row in zip( + inputs.request_ids, + inputs.ranks, + inputs.local_rows, + strict=True, + ): + state = self._require_mtp_request_state(request_id) + if ( + state.tail_token_id is None + or state.tail_slot_id is None + or state.tail_position is None + ): + raise RuntimeError( + f"DeepSeekV4 MTP committed tail is not initialized for {request_id!r}" + ) + buffers.decode_tail_token_ids[rank, local_row] = state.tail_token_id + buffers.decode_tail_positions[rank, local_row] = state.tail_position + buffers.decode_tail_slot_ids[rank, local_row] = state.tail_slot_id + buffers.decode_logit_row_indices[rank, local_row] = ( + local_row * layout.decode_seq + layout.decode_seq - 1 + ) + return max(inputs.per_rank_counts) * layout.decode_seq + def _require_stacked_weights(self) -> DeepSeekV4StackedLayerWeights: tensors = self._stacked_device_weights or self._stacked_host_weights if tensors is None: @@ -3240,7 +3353,11 @@ def _ensure_decode_buffers(self, hidden_size: int) -> _DeepSeekV4DecodeSharedBuf def _ensure_mtp_buffers(self, hidden_size: int) -> _DeepSeekV4MtpSharedBuffers | None: """Load immutable MTP weights and allocate mutable shared buffers before worker fork.""" - if self._compiled.mtp_prefill is None or self._compiled.mtp_decode is None: + legacy_mtp = ( + self._compiled.mtp_prefill is not None + and self._compiled.mtp_decode is not None + ) + if not self._compiled.enable_mtp and not legacy_mtp: return None if self._mtp_buffers is not None: return self._mtp_buffers @@ -3294,6 +3411,12 @@ def _ensure_mtp_buffers(self, hidden_size: int) -> _DeepSeekV4MtpSharedBuffers | decode_tail_slot_ids=self._shared_empty( (ranks, layout.decode_batch), torch.int32, name="mtp_decode_tail_slot_ids" ), + decode_tail_token_ids=self._shared_empty( + (ranks, layout.decode_batch), torch.long, name="mtp_decode_tail_token_ids" + ), + decode_tail_positions=self._shared_empty( + (ranks, layout.decode_batch), torch.int32, name="mtp_decode_tail_positions" + ), tail_init_hidden=self._shared_empty( (ranks, layout.decode_batch, layout.hc_mult, hidden), torch.float32, diff --git a/tests/test_deepseek_v4.py b/tests/test_deepseek_v4.py index 6c5848d5..d2c593f6 100644 --- a/tests/test_deepseek_v4.py +++ b/tests/test_deepseek_v4.py @@ -157,6 +157,27 @@ def test_cli_selects_deepseek_executor_and_forces_prefix_cache_off(tmp_path): assert config.executor_kwargs["enable_mtp"] is True +def test_cli_keeps_deepseek_autoregressive_decode_when_mtp_is_disabled(tmp_path): + model_dir = _write_deepseek_model_dir(tmp_path) + args = cli.build_parser().parse_args( + [ + "--model", + str(model_dir), + "--devices", + "0,1,2,3,4,5,6,7", + "--dp", + "8", + "--ep", + "8", + "--no-enable-mtp", + ] + ) + + config = cli.build_serving_engine_config(args) + + assert config.executor_kwargs["enable_mtp"] is False + + def test_tokenizer_falls_back_when_deepseek_config_fails_strict_validation(tmp_path, monkeypatch): class StrictDataclassFieldValidationError(Exception): pass @@ -722,9 +743,17 @@ def test_deepseek_worker_registers_main_and_mtp_weights_for_inheritance(monkeypa captured = {} class FakeDistributedWorker: - def __init__(self, compiled, *, persistent, inherited_host_tensors): + def __init__( + self, + compiled, + *, + persistent, + reset_persistent_windows, + inherited_host_tensors, + ): captured["compiled"] = compiled captured["persistent"] = persistent + captured["reset_persistent_windows"] = reset_persistent_windows captured["inherited"] = inherited_host_tensors monkeypatch.setattr("pypto.runtime.DistributedWorker", FakeDistributedWorker) @@ -744,6 +773,7 @@ def __init__(self, compiled, *, persistent, inherited_host_tensors): assert isinstance(worker, FakeDistributedWorker) assert captured["compiled"] == [compiled_program] assert captured["persistent"] is True + assert captured["reset_persistent_windows"] is False assert captured["inherited"] == [main_weight, mtp_weight] @@ -1321,6 +1351,102 @@ def fake_run_l3(_callable, *args): assert result.logits.shape == (1, model.config.vocab_size) +def test_deepseek_mtp_decode_fuses_main_verify_and_draft_into_one_dispatch(): + runner, model = _runner_for_prepared_inputs() + runner._compiled.decode = DeepSeekV4L3Callable(compiled=object(), name="decode_mtp_fused") + runner._decode_flow = runner._run_mtp_decode + layout = runner._compiled.layout + main_sampled_ids = torch.zeros( + layout.ranks, + layout.decode_tokens, + 8, + dtype=torch.int32, + ) + mtp_buffers = SimpleNamespace( + decode_accepted_counts=torch.ones( + layout.ranks, + layout.decode_batch, + dtype=torch.int32, + ), + decode_input_ids=torch.zeros( + layout.ranks, + layout.decode_tokens, + dtype=torch.long, + ), + decode_position_ids=torch.zeros( + layout.ranks, + layout.decode_tokens, + dtype=torch.int32, + ), + decode_sampled_ids=torch.zeros( + layout.ranks, + layout.decode_batch, + 8, + dtype=torch.int32, + ), + ) + state = SimpleNamespace( + draft_token_id=5, + tail_token_id=3, + tail_slot_id=0, + tail_position=126, + proposed_tokens=0, + accepted_tokens=0, + ) + runner._mtp_request_states["req-a"] = state + staged = SimpleNamespace( + request_ids=("req-a",), + ranks=(0,), + local_rows=(0,), + actual_batch=1, + per_rank_counts=(1,) + (0,) * (layout.ranks - 1), + ) + dispatches = [] + + runner._ensure_l3_shared_buffers = lambda _model: None + runner.prepare_mtp_decode_inputs = lambda _model, _batch: staged + runner._stage_decode_inputs = lambda prepared: prepared + runner._stage_fused_mtp_metadata = lambda _inputs: layout.decode_seq + runner._require_decode_buffers = lambda: SimpleNamespace(sampled_ids=main_sampled_ids) + runner._require_mtp_buffers = lambda: mtp_buffers + runner._require_decode_output_buffer = lambda _hidden_size: torch.empty(0) + runner._materialize_main_pre_hc_device = lambda _hidden_size: torch.empty(0) + runner._require_decode_logits_buffer = lambda _vocab_size: torch.empty(0) + runner._decode_fwd_args = lambda *_args: () + runner._fused_mtp_decode_args = lambda _main_args, _active_tokens: ("fused",) + runner._debug_decode_dispatch = lambda *_args: None + + def fake_run_l3(callable_spec, *args): + dispatches.append((callable_spec.name, args)) + main_sampled_ids[0, 0, 0] = 5 + main_sampled_ids[0, 1, 0] = 9 + mtp_buffers.decode_accepted_counts[0, 0] = 2 + mtp_buffers.decode_input_ids[0, 1] = 9 + mtp_buffers.decode_position_ids[0, 1] = 128 + mtp_buffers.decode_sampled_ids[0, 0, 0] = 7 + + runner._run_l3 = fake_run_l3 + + result = runner.run_decode( + model, + DecodeBatch( + request_ids=["req-a"], + token_ids=torch.tensor([[3]], dtype=torch.long), + hidden_states=None, + seq_lens=torch.tensor([128], dtype=torch.int32), + allow_device_greedy_sampling=True, + ), + ) + + assert dispatches == [("decode_mtp_fused", ("fused",))] + assert result.accepted_token_ids == [[5, 9]] + assert state.draft_token_id == 7 + assert state.tail_token_id == 9 + assert state.tail_position == 128 + assert state.proposed_tokens == 1 + assert state.accepted_tokens == 1 + + def test_deepseek_prefill_staging_keeps_worker_resident_cache_tensors_out(): layout = DeepSeekV4CacheLayout( ranks=1, @@ -1536,6 +1662,7 @@ def _write_deepseek_kernel_dir( (kernel_dir / "prefill_mtp.py").write_text("") (kernel_dir / "decode_layer.py").write_text("") (kernel_dir / "decode_fwd.py").write_text("") + (kernel_dir / "decode_fwd_mtp.py").write_text("") (kernel_dir / "decode_mtp.py").write_text("") (kernel_dir / "config.py").write_text( "\n".join(