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,
|
get_num_page_per_block_flashmla,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
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.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.server_args import get_global_server_args
|
from sglang.srt.server_args import get_global_server_args
|
||||||
from sglang.srt.utils import is_cuda, is_flashinfer_available, is_float4_e2m1fn_x2
|
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
|
from sgl_kernel import concat_mla_absorb_q
|
||||||
|
|
||||||
# Constants
|
# 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
|
# Block constraint from flashinfer requirements
|
||||||
# From flashinfer.decode._check_trtllm_gen_mla_shape:
|
# 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
|
global_zero_init_workspace_buffer = None
|
||||||
|
|
||||||
|
|
||||||
@@ -1035,6 +1066,25 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
k = torch.cat([k, k_rope], dim=-1)
|
k = torch.cat([k, k_rope], dim=-1)
|
||||||
k = k.view(-1, layer.tp_k_head_num, layer.head_dim)
|
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)
|
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.
|
# When chunked prefix cache is enabled, dispatch to different path for ragged attention.
|
||||||
if forward_batch.attn_attend_prefix_cache:
|
if forward_batch.attn_attend_prefix_cache:
|
||||||
# MHA for chunked prefix kv cache when running model with MLA
|
# MHA for chunked prefix kv cache when running model with MLA
|
||||||
@@ -1044,46 +1094,40 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
assert k_rope is None
|
assert k_rope is None
|
||||||
chunk_idx = forward_batch.prefix_chunk_idx
|
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(
|
return flashinfer.prefill.trtllm_ragged_attention_deepseek(
|
||||||
query=q,
|
**common_trtllm_args,
|
||||||
key=k,
|
|
||||||
value=v,
|
|
||||||
workspace_buffer=self.workspace_buffer,
|
|
||||||
seq_lens=forward_batch.prefix_chunk_seq_lens[chunk_idx],
|
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],
|
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,
|
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],
|
cum_seq_lens_kv=forward_batch.prefix_chunk_cu_seq_lens[chunk_idx],
|
||||||
enable_pdl=False,
|
|
||||||
is_causal=False,
|
is_causal=False,
|
||||||
return_lse=True,
|
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(
|
return flashinfer.prefill.trtllm_ragged_attention_deepseek(
|
||||||
query=q,
|
**common_trtllm_args,
|
||||||
key=k,
|
|
||||||
value=v,
|
|
||||||
workspace_buffer=self.workspace_buffer,
|
|
||||||
seq_lens=self.forward_prefill_metadata.seq_lens,
|
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,
|
max_kv_len=self.forward_prefill_metadata.max_seq_len,
|
||||||
bmm1_scale=layer.scaling,
|
|
||||||
bmm2_scale=1.0,
|
|
||||||
o_sf_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,
|
cum_seq_lens_kv=self.forward_prefill_metadata.cum_seq_lens,
|
||||||
enable_pdl=False,
|
|
||||||
is_causal=True,
|
is_causal=True,
|
||||||
return_lse=forward_batch.mha_return_lse,
|
return_lse=forward_batch.mha_return_lse,
|
||||||
|
out=out,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user