[Spec] DFlash: support pure-MLA targets with an fp8 KV cache (Kimi-K2.x-NVFP4) (#29218)

Co-authored-by: Hao Phan <htphan@nvidia.com>
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
Thanhhao
2026-07-07 19:52:14 -07:00
committed by GitHub
co-authored by Hao Phan Claude Opus 4.8
parent db40fd83d2
commit 7bc343470f
3 changed files with 96 additions and 8 deletions
@@ -2485,6 +2485,22 @@ class ModelRunner(ModelRunnerKVCacheMixin):
f"Unsupported kv_cache_dtype: {self.server_args.kv_cache_dtype}."
)
# DFLASH: fa4 draft attention can't read the target's fp8 KV (needs K.dtype == Q.dtype),
# so give the fa4 draft its own compute-dtype KV. fp8-capable backends keep the target dtype.
if (
self.is_draft_worker
and self.spec_algorithm.is_dflash()
and self.server_args.speculative_draft_attention_backend == "fa4"
and self.kv_cache_dtype != self.dtype
):
logger.info(
"DFLASH fa4 draft: overriding KV cache dtype %s -> %s "
"(fa4 needs K.dtype == Q.dtype; cannot read the target's quantized KV).",
self.kv_cache_dtype,
self.dtype,
)
self.kv_cache_dtype = self.dtype
def init_cublas(self):
"""We need to run a small matmul to init cublas. Otherwise, it will raise some errors later."""
dtype = torch.float16
@@ -117,6 +117,12 @@ class DFlashWorkerV2(BaseSpecWorker):
self.nccl_port = nccl_port
self._target_worker = target_worker
self.model_runner = target_worker.model_runner
self._need_mamba_verify_commit = (
self.model_runner.mambaish_config is not None
and hasattr(
self.model_runner.attn_backend, "update_mamba_state_after_mtp_verify"
)
)
self.page_size = server_args.page_size
# Normalized in arg_groups.speculative_hook.handle_speculative_decoding.
self.draft_window_size: Optional[int] = (
@@ -1099,9 +1105,9 @@ class DFlashWorkerV2(BaseSpecWorker):
cache per-step intermediate states. After acceptance, we need to commit the
state corresponding to each request's last accepted step.
"""
attn_backend = self.target_worker.model_runner.attn_backend
if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"):
if not self._need_mamba_verify_commit:
return
attn_backend = self.target_worker.model_runner.attn_backend
last_correct_step_indices = commit_lens.to(torch.int64) - 1
mamba_steps_to_track = None
@@ -1530,12 +1536,8 @@ class DFlashWorkerV2(BaseSpecWorker):
batch.out_cache_loc = verify_out_cache_loc
sampling_info = batch.sampling_info
need_mamba_verify_commit = hasattr(
self.target_worker.model_runner.attn_backend,
"update_mamba_state_after_mtp_verify",
)
seq_lens_pre_verify = (
batch.seq_lens.clone() if need_mamba_verify_commit else None
batch.seq_lens.clone() if self._need_mamba_verify_commit else None
)
seq_lens_cpu_backup = batch.seq_lens_cpu
seq_lens_sum_backup = batch.seq_lens_sum
@@ -1655,7 +1657,7 @@ class DFlashWorkerV2(BaseSpecWorker):
1, accept_len.to(torch.int64)[:, None], bonus[:, None]
)
if need_mamba_verify_commit:
if self._need_mamba_verify_commit:
assert seq_lens_pre_verify is not None
self._update_target_mamba_state_after_verify(
batch=batch,