From 6559e43f306844c8aff9da704b173f178c27224f Mon Sep 17 00:00:00 2001 From: weiliang Date: Fri, 19 Dec 2025 17:24:00 +0800 Subject: [PATCH] Support FP8 MLA prefill and 128k context. (#14395) --- .../layers/attention/trtllm_mla_backend.py | 112 ++++++++++++------ 1 file changed, 78 insertions(+), 34 deletions(-) diff --git a/python/sglang/srt/layers/attention/trtllm_mla_backend.py b/python/sglang/srt/layers/attention/trtllm_mla_backend.py index 9740095a5..6afc12df8 100755 --- a/python/sglang/srt/layers/attention/trtllm_mla_backend.py +++ b/python/sglang/srt/layers/attention/trtllm_mla_backend.py @@ -21,6 +21,7 @@ from sglang.srt.layers.attention.utils import ( get_num_page_per_block_flashmla, ) from sglang.srt.layers.dp_attention import get_attention_tp_size +from sglang.srt.layers.quantization.fp8_kernel import scaled_fp8_quant from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.server_args import get_global_server_args from sglang.srt.utils import is_cuda, is_flashinfer_available, is_float4_e2m1fn_x2 @@ -39,7 +40,7 @@ if _is_cuda: from sgl_kernel import concat_mla_absorb_q # Constants -DEFAULT_WORKSPACE_SIZE_MB = 128 # Memory workspace size in MB +DEFAULT_WORKSPACE_SIZE_MB = 150 # Memory workspace size in MB # Block constraint from flashinfer requirements # From flashinfer.decode._check_trtllm_gen_mla_shape: @@ -194,6 +195,36 @@ def unpad_draft_extend_output_kernel( ) +def _quantize_fp8_qkv(q, k, v, layer): + q = q.to(torch.float8_e4m3fn) + + k_scale = getattr(layer, "k_scale_float", None) + if k_scale is None: + k_scale = 1.0 + if k_scale != 1.0: + assert hasattr(layer, "k_scale"), "k_scale is not set" + k_2d, _ = scaled_fp8_quant( + k.reshape(-1, k.shape[-1]).contiguous(), layer.k_scale + ) + k = k_2d.reshape(k.shape) + else: + k = k.to(torch.float8_e4m3fn) + + v_scale = getattr(layer, "v_scale_float", None) + if v_scale is None: + v_scale = 1.0 + if v_scale != 1.0: + assert hasattr(layer, "v_scale"), "v_scale is not set" + v_2d, _ = scaled_fp8_quant( + v.reshape(-1, v.shape[-1]).contiguous(), layer.v_scale + ) + v = v_2d.reshape(v.shape) + else: + v = v.to(torch.float8_e4m3fn) + + return q, k, v, k_scale, v_scale + + global_zero_init_workspace_buffer = None @@ -1035,6 +1066,25 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): k = torch.cat([k, k_rope], dim=-1) k = k.view(-1, layer.tp_k_head_num, layer.head_dim) v = v.view(-1, layer.tp_k_head_num, layer.v_head_dim) + + q_scale = k_scale = v_scale = 1.0 + if self.data_type == torch.float8_e4m3fn: + q, k, v, k_scale, v_scale = _quantize_fp8_qkv(q, k, v, layer) + + common_trtllm_args = { + "query": q, + "key": k, + "value": v, + "workspace_buffer": self.workspace_buffer, + "batch_size": forward_batch.batch_size, + "window_left": -1, + "enable_pdl": False, + "max_q_len": self.forward_prefill_metadata.max_seq_len, + "bmm1_scale": q_scale * k_scale * layer.scaling, + "bmm2_scale": v_scale, + "cum_seq_lens_q": self.forward_prefill_metadata.cum_seq_lens, + } + # When chunked prefix cache is enabled, dispatch to different path for ragged attention. if forward_batch.attn_attend_prefix_cache: # MHA for chunked prefix kv cache when running model with MLA @@ -1044,47 +1094,41 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend): assert k_rope is None chunk_idx = forward_batch.prefix_chunk_idx - output_shape = (q.shape[0], layer.tp_q_head_num, layer.v_head_dim) + out = torch.zeros( + q.shape[0], + layer.tp_q_head_num, + layer.v_head_dim, + dtype=self.q_data_type, + device=q.device, + ) return flashinfer.prefill.trtllm_ragged_attention_deepseek( - query=q, - key=k, - value=v, - workspace_buffer=self.workspace_buffer, + **common_trtllm_args, seq_lens=forward_batch.prefix_chunk_seq_lens[chunk_idx], - max_q_len=self.forward_prefill_metadata.max_seq_len, max_kv_len=forward_batch.prefix_chunk_max_seq_lens[chunk_idx], - bmm1_scale=layer.scaling, - bmm2_scale=1.0, o_sf_scale=-1.0, - batch_size=forward_batch.batch_size, - window_left=-1, - cum_seq_lens_q=self.forward_prefill_metadata.cum_seq_lens, cum_seq_lens_kv=forward_batch.prefix_chunk_cu_seq_lens[chunk_idx], - enable_pdl=False, is_causal=False, return_lse=True, - out=torch.zeros(*output_shape, dtype=q.dtype, device=q.device), + out=out, + ) + else: + out = torch.empty( + q.shape[0], + q.shape[1], + v.shape[2], + device=q.device, + dtype=self.q_data_type, + ) + return flashinfer.prefill.trtllm_ragged_attention_deepseek( + **common_trtllm_args, + seq_lens=self.forward_prefill_metadata.seq_lens, + max_kv_len=self.forward_prefill_metadata.max_seq_len, + o_sf_scale=1.0, + cum_seq_lens_kv=self.forward_prefill_metadata.cum_seq_lens, + is_causal=True, + return_lse=forward_batch.mha_return_lse, + out=out, ) - - return flashinfer.prefill.trtllm_ragged_attention_deepseek( - query=q, - key=k, - value=v, - workspace_buffer=self.workspace_buffer, - seq_lens=self.forward_prefill_metadata.seq_lens, - max_q_len=self.forward_prefill_metadata.max_seq_len, - max_kv_len=self.forward_prefill_metadata.max_seq_len, - bmm1_scale=layer.scaling, - bmm2_scale=1.0, - o_sf_scale=1.0, - batch_size=forward_batch.batch_size, - window_left=-1, - cum_seq_lens_q=self.forward_prefill_metadata.cum_seq_lens, - cum_seq_lens_kv=self.forward_prefill_metadata.cum_seq_lens, - enable_pdl=False, - is_causal=True, - return_lse=forward_batch.mha_return_lse, - ) class TRTLLMMLAMultiStepDraftBackend(FlashInferMLAMultiStepDraftBackend):