diff --git a/vllm/model_executor/models/deepseek_v2.py b/vllm/model_executor/models/deepseek_v2.py index cd28fb0192f..17ddd5edece 100644 --- a/vllm/model_executor/models/deepseek_v2.py +++ b/vllm/model_executor/models/deepseek_v2.py @@ -66,10 +66,6 @@ from vllm.model_executor.layers.quantization import QuantizationConfig from vllm.model_executor.layers.quantization.utils.fp8_utils import ( per_token_group_quant_fp8, ) -from vllm.model_executor.layers.quantization.utils.quant_utils import ( - GroupShape, - scaled_dequantize, -) from vllm.model_executor.layers.rotary_embedding import get_rope from vllm.model_executor.layers.sparse_attn_indexer import ( SparseAttnIndexer, @@ -632,6 +628,10 @@ class Indexer(nn.Module): self.vllm_config = vllm_config self.config = config self.quant_config = quant_config + self.is_fp4_ckpt = ( + self.quant_config is not None + and self.quant_config.get_name() == "modelopt_fp4" + ) # self.indexer_cfg = config.attn_module_list_cfg[0]["attn_index"] self.topk_tokens = config.index_topk self.n_head = config.index_n_heads # 64 @@ -646,16 +646,36 @@ class Indexer(nn.Module): quant_config=quant_config, prefix=f"{prefix}.wq_b", ) - # Fused wk + weights_proj: single GEMM producing [head_dim + n_head]. - # FP8 wk weights are upcasted to BF16 during loading to maintain fusion. - self.wk_weights_proj = MergedColumnParallelLinear( - hidden_size, - [self.head_dim, self.n_head], - bias=False, - quant_config=None, - disable_tp=True, - prefix=f"{prefix}.wk_weights_proj", - ) + if self.is_fp4_ckpt: + # Fused wk + weights_proj: single GEMM producing [head_dim + n_head]. + # weights_proj does not get quantized, + # so we run both with quant_config=None + # wk may be upcasted from the default quant; + # experiments show fusion is always faster unless WK proj is in FP4, + # which is not the case for all known quants. + self.wk_weights_proj = MergedColumnParallelLinear( + hidden_size, + [self.head_dim, self.n_head], + bias=False, + quant_config=None, + disable_tp=True, + prefix=f"{prefix}.wk_weights_proj", + ) + else: + self.wk = ReplicatedLinear( + hidden_size, + self.head_dim, + bias=False, + quant_config=quant_config, + prefix=f"{prefix}.wk", + ) + self.weights_proj = ReplicatedLinear( + hidden_size, + self.n_head, + bias=False, + quant_config=None, + prefix=f"{prefix}.weights_proj", + ) self.k_norm = LayerNorm(self.head_dim, eps=1e-6) self.softmax_scale = self.head_dim**-0.5 @@ -696,10 +716,14 @@ class Indexer(nn.Module): q_pe, q_nope = torch.split( q, [self.rope_dim, self.head_dim - self.rope_dim], dim=-1 ) - # Fused wk + weights_proj: one GEMM, then split - kw, _ = self.wk_weights_proj(hidden_states) - k = kw[:, : self.head_dim] - weights = kw[:, self.head_dim :] + if self.is_fp4_ckpt: + # Fused wk + weights_proj: one GEMM, then split + kw, _ = self.wk_weights_proj(hidden_states) + k = kw[:, : self.head_dim] + weights = kw[:, self.head_dim :] + else: + k, _ = self.wk(hidden_states) + weights, _ = self.weights_proj(hidden_states) k = self.k_norm(k) k_pe, k_nope = torch.split( @@ -737,46 +761,6 @@ class Indexer(nn.Module): return self.indexer_op(hidden_states, q_fp8, k, weights) -def _try_load_fp8_indexer_wk(name, tensor, buf, params_dict, loaded_params): - """ - We fuse the WK and weights_proj projections, but in some checkpoints WK is stored - in FP8 with a separate weight_scale_inv, while weights_proj is stored in BF16. - Upcasting to BF16 during loading enables the fusion. This function loads the FP8 WK - weights and scale, and when both are available, dequantizes to BF16 and stores into - the fused wk_weights_proj.weight parameter. - """ - if "indexer.wk." not in name or "wk_weights" in name: - return False # Weight is not an isolated WK weight for the indexer, ignore. - is_weight = name.endswith(".weight") and tensor.dtype == torch.float8_e4m3fn - is_scale = "weight_scale_inv" in name - if not is_weight and not is_scale: - return False # WK is not in FP8 format, ignore. - # Buffer this tensor (weight or scale) until both have arrived. - layer_prefix = name.rsplit(".wk.", 1)[0] # e.g. "model.layers.0.self_attn.indexer" - entry = buf.setdefault(layer_prefix, {}) - entry["weight" if is_weight else "scale"] = tensor - if "weight" not in entry or "scale" not in entry: - return True # still waiting for the other param - - # We have both weight and scale: dequantize FP8 to BF16. - weight_fp8, scale_inv = entry["weight"], entry["scale"] - del buf[layer_prefix] - block_size = weight_fp8.shape[1] // scale_inv.shape[1] - weight_bf16 = scaled_dequantize( - weight_fp8, - scale_inv, - group_shape=GroupShape(block_size, block_size), - out_dtype=torch.bfloat16, - ) - - # Load the dequantized weight into shard 0 of the fused buffer. - fused_name = f"{layer_prefix}.wk_weights_proj.weight" - param = params_dict[fused_name] - param.weight_loader(param, weight_bf16, 0) - loaded_params.add(fused_name) - return True - - def _min_latency_fused_qkv_a_proj_impl( input_: torch.Tensor, weight: torch.Tensor, @@ -1360,6 +1344,10 @@ class DeepseekV2ForCausalLM( quant_config = vllm_config.quant_config self.config = config self.quant_config = quant_config + self.is_fp4_ckpt = ( + self.quant_config is not None + and self.quant_config.get_name() == "modelopt_fp4" + ) qk_nope_head_dim = getattr(config, "qk_nope_head_dim", 0) qk_rope_head_dim = getattr(config, "qk_rope_head_dim", 0) @@ -1485,13 +1473,13 @@ class DeepseekV2ForCausalLM( ("qkv_proj", "k_proj", "k"), ("qkv_proj", "v_proj", "v"), ] - # Fused indexer wk + weights_proj (shard 0 = wk, shard 1 = weights_proj) - _pending_wk_fp8: dict = {} # When WK is in FP8, we dequant to BF16 for fusion - indexer_fused_mapping = [ - ("wk_weights_proj", "wk", 0), - ("wk_weights_proj", "weights_proj", 1), - ] - stacked_params_mapping.extend(indexer_fused_mapping) + if self.is_fp4_ckpt: + # Fused indexer wk + weights_proj (shard 0 = wk, shard 1 = weights_proj) + indexer_fused_mapping = [ + ("wk_weights_proj", "wk", 0), + ("wk_weights_proj", "weights_proj", 1), + ] + stacked_params_mapping.extend(indexer_fused_mapping) if self.use_mha: stacked_params_mapping.extend(mha_params_mapping) @@ -1528,11 +1516,6 @@ class DeepseekV2ForCausalLM( rocm_aiter_moe_shared_expert_enabled and ("mlp.shared_experts" in name) ) - if _try_load_fp8_indexer_wk( - name, loaded_weight, _pending_wk_fp8, params_dict, loaded_params - ): - continue - for param_name, weight_name, shard_id in stacked_params_mapping: # Skip non-stacked layers and experts (experts handled below). if weight_name not in name: