diff --git a/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py b/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py index 120f7f19652..965ef87e283 100644 --- a/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py +++ b/vllm/model_executor/models/deepseek_v3_2_monolithic/decoder_layer.py @@ -11,9 +11,6 @@ from torch import nn from vllm.config import VllmConfig from vllm.distributed import get_tensor_model_parallel_world_size -from vllm.model_executor.layers.quantization.utils.fp8_utils import ( - per_token_group_quant_fp8, -) from .attention import MonolithicMLAAttention from .ops import ( @@ -186,7 +183,7 @@ class MonolithicDecoderLayer(nn.Module): index_q = index_q.view(-1, self.attn.index_n_heads, self.attn.index_head_dim) # Step 4. Q RoPE + Index Q RoPE + Quantize + Index weights - fused_q( + index_q_fp8, index_weights = fused_q( positions, # Q RoPE q, @@ -195,29 +192,18 @@ class MonolithicDecoderLayer(nn.Module): # Index Q RoPE index_q, self.attn.indexer_rope_emb.cos_sin_cache, + # Index Q Quantize + 1e-10, # quant_eps + # Index weights + index_weights, + self.attn.indexer_softmax_scale, + self.attn.index_n_heads**-0.5, ) - index_q_fp8, index_q_scale = per_token_group_quant_fp8( - index_q.view(-1, self.attn.index_head_dim), - self.attn.indexer_quant_block_size, - column_major_scales=False, - use_ue8m0=True, - ) - index_q_fp8 = index_q_fp8.view( - -1, self.attn.index_n_heads, self.attn.index_head_dim - ) - index_q_scale = index_q_scale.view(-1, self.attn.index_n_heads, 1) - - index_weights = ( - index_weights.unsqueeze(-1) - * index_q_scale - * self.attn.indexer_softmax_scale - * self.attn.index_n_heads**-0.5 - ).squeeze(-1) - + # Step 5. Sparse indexer. self.attn.indexer_op(hidden_states, index_q_fp8, index_k, index_weights) - # 4-7. KV cache update + W_UK_T absorption + sparse attn + W_UV + # Step 6. MLA attention. attn_out = self.attn.mla_attn( q, kv_c, @@ -228,7 +214,7 @@ class MonolithicDecoderLayer(nn.Module): ), ) - # 8. Output projection (TP all-reduce) + # Step 7. Output projection. hidden_states, _ = self.attn.o_proj(attn_out) # Post-attn norm + residual diff --git a/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py b/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py index 159f05262f3..f442020a635 100644 --- a/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py +++ b/vllm/model_executor/models/deepseek_v3_2_monolithic/ops.py @@ -609,6 +609,19 @@ def _fused_q_kernel( index_q_cos_sin_ptr, index_q_cos_sin_stride, INDEX_Q_HALF_ROT_DIM: tl.constexpr, + # Index Q Quantize + index_q_fp8_ptr, + index_q_fp8_eps, + FP8_MIN: tl.constexpr, + FP8_MAX: tl.constexpr, + INDEX_Q_HEAD_DIM: tl.constexpr, + # Index weights + index_weights_ptr, + index_weights_stride, + index_weights_softmax_scale, + index_weights_head_scale, + index_weights_out_ptr, + index_weights_out_stride, ): tok_idx = tl.program_id(1) pos = tl.load(pos_ptr + tok_idx) @@ -658,6 +671,42 @@ def _fused_q_kernel( False, ) + # Index Q Quantize + index_q_block = tl.arange(0, INDEX_Q_HEAD_DIM) + index_q = tl.load( + index_q_ptr + + tok_idx * index_q_stride0 + + head_idx * index_q_stride1 + + index_q_block + ) + index_q = index_q.to(tl.float32) + + index_q_abs_max = tl.maximum(tl.max(tl.abs(index_q)), index_q_fp8_eps) + s = index_q_abs_max * (1.0 / FP8_MAX) + index_q_scale = tl.exp2(tl.ceil(tl.log2(s))) + + index_q_fp8 = tl.clamp(index_q / index_q_scale, FP8_MIN, FP8_MAX) + tl.store( + index_q_fp8_ptr + + tok_idx * index_q_stride0 + + head_idx * index_q_stride1 + + index_q_block, + index_q_fp8, + ) + + # Index weights update + index_weights = tl.load( + index_weights_ptr + tok_idx * index_weights_stride + head_idx + ) + index_weights = index_weights.to(tl.float32) + index_weights *= index_q_scale + index_weights *= index_weights_softmax_scale + index_weights *= index_weights_head_scale + tl.store( + index_weights_out_ptr + tok_idx * index_weights_out_stride + head_idx, + index_weights, + ) + def fused_q( positions: torch.Tensor, @@ -666,7 +715,13 @@ def fused_q( q_start_offset: int, index_q: torch.Tensor, index_q_cos_sin_cache: torch.Tensor, -) -> None: + # Index Q Quantize + quant_eps: float, + # Index weights + index_weights: torch.Tensor, + index_weights_softmax_scale: float, + index_weights_head_scale: float, +) -> tuple[torch.Tensor, torch.Tensor]: assert positions.ndim == 1 assert q.ndim == 3 assert q_cos_sin_cache.ndim == 2 @@ -676,7 +731,12 @@ def fused_q( num_tokens = positions.shape[0] num_q_heads = q.shape[1] num_index_q_heads = index_q.shape[1] + index_q_head_dim = index_q.shape[2] + FP8_DTYPE = torch.float8_e4m3fn + FP8_FINFO = torch.finfo(FP8_DTYPE) + index_q_fp8 = torch.empty_like(index_q, dtype=FP8_DTYPE) + index_weights_out = torch.empty_like(index_weights, dtype=torch.float32) _fused_q_kernel[(2, num_tokens, num_index_q_heads)]( positions, q, @@ -694,4 +754,17 @@ def fused_q( index_q_cos_sin_cache, index_q_cos_sin_cache.stride(0), index_q_cos_sin_cache.shape[-1] // 2, + index_q_fp8, + quant_eps, + FP8_FINFO.min, + FP8_FINFO.max, + index_q_head_dim, + index_weights, + index_weights.stride(0), + index_weights_softmax_scale, + index_weights_head_scale, + index_weights_out, + index_weights_out.stride(0), + num_warps=1, # TODO: Tune this ) + return index_q_fp8, index_weights_out