[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:
co-authored by
Hao Phan
Claude Opus 4.8
parent
db40fd83d2
commit
7bc343470f
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user