Support FP8 MLA prefill and 128k context. (#14395)

This commit is contained in:
weiliang
2025-12-19 17:24:00 +08:00
committed by GitHub
parent af780c59f3
commit 6559e43f30

View File

@@ -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):