Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@ server_args: >-
--block-size 256
--gpu-memory-utilization 0.5
--kv-cache-dtype fp8
--attention_config.use_fp4_indexer_cache=True
--attention_config.indexer_kv_dtype=mxfp4
--max-num-batched-tokens 16384
--max-num-seqs 128
--speculative-config '{"method":"dspark",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,4 +2,4 @@ model_name: "deepseek-ai/DeepSeek-V4-Flash"
accuracy_threshold: 0.95
num_questions: 1319
num_fewshot: 5
server_args: "--trust-remote-code --kv-cache-dtype fp8 --block-size 256 --enable-expert-parallel --tensor-parallel-size 2 --attention_config.use_fp4_indexer_cache=True --moe-backend deep_gemm_mega_moe --tokenizer-mode deepseek_v4 --tool-call-parser deepseek_v4 --enable-auto-tool-choice --reasoning-parser deepseek_v4 --speculative_config.method=mtp --speculative_config.num_speculative_tokens=2"
server_args: "--trust-remote-code --kv-cache-dtype fp8 --block-size 256 --enable-expert-parallel --tensor-parallel-size 2 --attention_config.indexer_kv_dtype=mxfp4 --moe-backend deep_gemm_mega_moe --tokenizer-mode deepseek_v4 --tool-call-parser deepseek_v4 --enable-auto-tool-choice --reasoning-parser deepseek_v4 --speculative_config.method=mtp --speculative_config.num_speculative_tokens=2"
41 changes: 34 additions & 7 deletions vllm/config/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,13 @@
from pydantic import field_validator

from vllm.config.utils import config
from vllm.logger import init_logger
from vllm.v1.attention.backends.mla.prefill.registry import MLAPrefillBackendEnum
from vllm.v1.attention.backends.registry import AttentionBackendEnum

IndexerKVDType = Literal["bf16", "fp8", "mxfp4", "nvfp4"]
logger = init_logger(__name__)

IndexerKVDType = Literal["auto", "bf16", "fp8", "mxfp4", "nvfp4"]
MiniMaxM3MSADecodeBackend = Literal["triton", "cutlass"]


Expand Down Expand Up @@ -65,12 +68,15 @@ class AttentionConfig:
use_prefill_query_quantization: bool = False
"""If set, quantize query for attention in prefill."""

use_fp4_indexer_cache: bool = False
"""If set, use fp4 indexer cache for dsv32 family model (not support yet)"""
use_fp4_indexer_cache: bool | None = None
"""Deprecated alias for `indexer_kv_dtype`; use that instead. True maps to
`mxfp4`, False is a no-op (it selected the model default already)."""

indexer_kv_dtype: IndexerKVDType = "bf16"
"""Data type for the sparse-attention indexer K cache. Quantized formats
(fp8, mxfp4, nvfp4) require indexer kernel support in the backend."""
indexer_kv_dtype: IndexerKVDType = "auto"
"""Data type for the sparse-attention indexer K cache. "auto" picks the
model's default (bf16 for MiniMax M3, fp8 for the DeepSeek sparse
indexer). Quantized formats (fp8, mxfp4, nvfp4) require indexer kernel
support in the backend."""

use_non_causal: bool = False
"""Whether to use non-causal (bidirectional) attention."""
Expand Down Expand Up @@ -114,6 +120,26 @@ def __post_init__(self) -> None:
# layers still use the platform's normal automatic backend.
self.backend = None

if self.use_fp4_indexer_cache is not None:
logger.warning(
"use_fp4_indexer_cache is deprecated and will be removed in "
"v0.19. Use indexer_kv_dtype instead (True -> 'mxfp4')."
)
if self.use_fp4_indexer_cache:
if self.indexer_kv_dtype not in ("auto", "mxfp4"):
raise ValueError(
"use_fp4_indexer_cache=True conflicts with "
f"indexer_kv_dtype={self.indexer_kv_dtype!r}. Set only "
"indexer_kv_dtype."
)
self.indexer_kv_dtype = "mxfp4"

def resolve_indexer_kv_dtype(self, default: IndexerKVDType) -> IndexerKVDType:
"""Resolve `indexer_kv_dtype`, substituting `default` for "auto"."""
if self.indexer_kv_dtype == "auto":
return default
return self.indexer_kv_dtype

def compute_hash(self) -> str:
"""
Provide a hash that uniquely identifies all the configs
Expand All @@ -124,7 +150,8 @@ def compute_hash(self) -> str:
"""
from vllm.config.utils import get_hash_factors, hash_factors

ignored_factors: set[str] = set()
# Folded into indexer_kv_dtype by __post_init__.
ignored_factors: set[str] = {"use_fp4_indexer_cache"}
factors = get_hash_factors(self, ignored_factors)
return hash_factors(factors)

Expand Down
3 changes: 2 additions & 1 deletion vllm/models/deepseek_v4/attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@
from vllm.v1.attention.backend import AttentionBackend, AttentionMetadata
from vllm.v1.attention.backends.mla.indexer import (
DeepseekV4IndexerBackend,
dsa_indexer_uses_fp4,
get_max_prefill_buffer_size,
)
from vllm.v1.attention.backends.mla.sparse_swa import DeepseekV4SWACache
Expand Down Expand Up @@ -805,7 +806,7 @@ def __init__(
self.q_lora_rank = q_lora_rank # 1536
self.compress_ratio = compress_ratio
self.eager_scratch_pool = eager_scratch_pool
self.use_fp4_kv = self.vllm_config.attention_config.use_fp4_indexer_cache
self.use_fp4_kv = dsa_indexer_uses_fp4(vllm_config)
logger.info_once(
"Using %s indexer cache for Lightning Indexer.",
"MXFP4" if self.use_fp4_kv else "FP8",
Expand Down
4 changes: 3 additions & 1 deletion vllm/models/minimax_m3/nvidia/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -498,7 +498,9 @@ def __init__(
set_default_quant_scales(self, register_buffer=True)
# Indexer side-cache dtype, mirroring --kv-cache-dtype for the main
# cache (--attention-config '{"indexer_kv_dtype": ...}').
self.indexer_kv_dtype = vllm_config.attention_config.indexer_kv_dtype
self.indexer_kv_dtype = vllm_config.attention_config.resolve_indexer_kv_dtype(
"bf16"
)

# Shared top-k buffer: the indexer writes the selected blocks into it and
# the attend impl reads them back (so nothing crosses the eager break as a
Expand Down
37 changes: 24 additions & 13 deletions vllm/v1/attention/backends/mla/indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,28 @@

logger = init_logger(__name__)

# The DSA indexer K cache is always quantized; "auto" means fp8 (V3.2 layout)
# and mxfp4 is the opt-in Blackwell path.
DSA_INDEXER_KV_DTYPES = ("fp8", "mxfp4")


def dsa_indexer_uses_fp4(vllm_config: VllmConfig) -> bool:
"""Whether the DeepSeek sparse indexer should use the MXFP4 K cache."""
kv_dtype = vllm_config.attention_config.resolve_indexer_kv_dtype("fp8")
if kv_dtype not in DSA_INDEXER_KV_DTYPES:
raise ValueError(
f"indexer_kv_dtype={kv_dtype!r} is not supported by the DeepSeek "
f"sparse indexer (expected one of {DSA_INDEXER_KV_DTYPES})."
)
use_fp4 = kv_dtype == "mxfp4"
if use_fp4 and not current_platform.is_device_capability_family(100):
raise ValueError(
"indexer_kv_dtype='mxfp4' requires Blackwell datacenter GPUs "
"(sm_10x, e.g. B200/GB200); sm_120 (consumer Blackwell) and "
"earlier architectures are not supported."
)
return use_fp4


@triton.jit
def _prepare_uniform_decode_kernel(
Expand Down Expand Up @@ -526,18 +548,7 @@ def __init__(self, *args, block_table_width: int, **kwargs) -> None:
if self.vllm_config.speculative_config
else 0
)
self.use_fp4_indexer_cache = (
self.vllm_config.attention_config.use_fp4_indexer_cache
)

assert (
current_platform.is_device_capability_family(100)
or not self.use_fp4_indexer_cache
), (
"use_fp4_indexer_cache requires Blackwell datacenter GPUs "
"(sm_10x, e.g. B200/GB200); sm_120 (consumer Blackwell) and "
"earlier architectures are not supported."
)
self.use_fp4_indexer_cache = dsa_indexer_uses_fp4(self.vllm_config)

next_n = self.num_speculative_tokens + 1
self.decode_threshold = next_n
Expand All @@ -546,7 +557,7 @@ def __init__(self, *args, block_table_width: int, **kwargs) -> None:
self.supports_varlen = _supports_varlen_paged_mqa_logits()
logger.info_once(
"DSA indexer decode path: use_flattening=%s supports_varlen=%s "
"(next_n=%d, use_fp4_indexer_cache=%s)",
"(next_n=%d, use_fp4_cache=%s)",
self.use_flattening,
self.supports_varlen,
next_n,
Expand Down
Loading