Support FP8 MLA prefill and 128k context. (#14395)
This commit is contained in:
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user