Fix gpu-fault when running mtp in eager mode (#15233)

This commit is contained in:
kk
2025-12-17 17:25:06 +08:00
committed by GitHub
parent feb8e30b9d
commit 888594333e
2 changed files with 172 additions and 30 deletions

View File

@@ -40,6 +40,7 @@ except ImportError:
)
from sglang.srt.configs.model_config import AttentionArch
from sglang.srt.layers.attention.utils import pad_sequence_with_mask
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
from sglang.srt.utils import get_bool_env_var
@@ -77,7 +78,7 @@ class ForwardMetadata:
reduce_final_map: Optional[torch.Tensor] = None
reduce_partial_map: Optional[torch.Tensor] = None
num_kv_splits: Optional[int] = None
# num_kv_splits_indptr: Optional[torch.Tensor] = None
run_graph: Optional[bool] = True
global_workspace_buffer = None
@@ -390,6 +391,7 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
run_graph=False,
)
elif forward_batch.forward_mode.is_draft_extend():
@@ -446,7 +448,7 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
run_graph=False,
)
else:
self.indices_updater_prefill.update(
@@ -541,7 +543,7 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
# num_kv_splits_indptr=num_kv_splits_indptr,
run_graph=False,
)
else:
self.indices_updater_prefill.update(
@@ -1176,10 +1178,6 @@ class AiterAttnBackend(AttentionBackend):
)
return o
elif forward_batch.forward_mode.is_draft_extend():
o = q.new_empty(
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
dtype=self.input_dtype,
)
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
@@ -1207,29 +1205,75 @@ class AiterAttnBackend(AttentionBackend):
intra_batch_mode=intra_batch_mode,
)
mla_decode_fwd(
q,
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
o,
self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices,
self.forward_metadata.kv_last_page_len,
self.forward_metadata.max_q_len,
layer.scaling,
layer.logit_cap,
work_meta_data=work_metadata,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=layer.k_scale,
kv_scale=layer.k_scale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
return o
if self.forward_metadata.run_graph is not True:
bs, q_pad, q_mask = pad_sequence_with_mask(
q.view(q.shape[0], -1),
qo_indptr[:-1],
forward_batch.extend_seq_lens,
self.forward_metadata.max_q_len,
)
o = q.new_empty(
(
bs * self.forward_metadata.max_q_len,
layer.tp_q_head_num,
layer.v_head_dim,
),
dtype=self.input_dtype,
)
mla_decode_fwd(
q_pad.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
o,
self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices,
self.forward_metadata.kv_last_page_len,
self.forward_metadata.max_q_len,
layer.scaling,
layer.logit_cap,
work_meta_data=work_metadata,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=layer.k_scale,
kv_scale=layer.k_scale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
return o[q_mask]
else:
o = q.new_empty(
(q.shape[0], layer.tp_q_head_num, layer.v_head_dim),
dtype=self.input_dtype,
)
mla_decode_fwd(
q,
K_Buffer.view(-1, 1, 1, layer.qk_head_dim),
o,
self.forward_metadata.qo_indptr,
self.forward_metadata.kv_indptr,
self.forward_metadata.kv_indices,
self.forward_metadata.kv_last_page_len,
self.forward_metadata.max_q_len,
layer.scaling,
layer.logit_cap,
work_meta_data=work_metadata,
work_indptr=work_indptr,
work_info_set=work_info_set,
reduce_indptr=reduce_indptr,
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
q_scale=layer.k_scale,
kv_scale=layer.k_scale,
intra_batch_mode=intra_batch_mode,
num_kv_splits=num_kv_splits,
)
return o
else:
raise ValueError(
f"Invalid forward mode for MLA prefill: {forward_batch.forward_mode=}"

View File

@@ -179,3 +179,101 @@ def concat_and_cast_mha_k_triton(
nope_dim,
rope_dim,
)
@triton.jit
def pad_sequence_with_mask_kernel(
input_ptr, # (total_tokens, hidden)
offsets_ptr, # (B,)
lengths_ptr, # (B,)
output_ptr, # (B, max_len, hidden)
mask_ptr, # (B, max_len)
max_len,
hidden_dim,
BLOCK_M: tl.constexpr, # seq block
BLOCK_D: tl.constexpr, # hidden block
):
b = tl.program_id(0) # batch index
m = tl.program_id(1) # seq block index
offset = tl.load(offsets_ptr + b)
length = tl.load(lengths_ptr + b)
seq_ids = m * BLOCK_M + tl.arange(0, BLOCK_M)
hid_ids = tl.arange(0, BLOCK_D)
seq_mask = seq_ids < max_len
valid_token = seq_ids < length
# input index
in_token = offset + seq_ids
in_ptr = input_ptr + in_token[:, None] * hidden_dim + hid_ids[None, :]
# output index
out_ptr = (
output_ptr
+ b * max_len * hidden_dim
+ seq_ids[:, None] * hidden_dim
+ hid_ids[None, :]
)
values = tl.load(
in_ptr,
mask=valid_token[:, None] & (hid_ids[None, :] < hidden_dim),
other=0.0,
)
tl.store(
out_ptr,
values,
mask=seq_mask[:, None] & (hid_ids[None, :] < hidden_dim),
)
# attention mask
if tl.program_id(2) == 0:
mask_out_ptr = mask_ptr + b * max_len + seq_ids
tl.store(mask_out_ptr, valid_token, mask=seq_mask)
def pad_sequence_with_mask(
input_emb, # (total_tokens, hidden)
offsets, # (B,)
lengths, # (B,)
max_len,
):
B = offsets.shape[0]
hidden_dim = input_emb.shape[1]
output = torch.zeros(
(B, max_len, hidden_dim),
device=input_emb.device,
dtype=input_emb.dtype,
)
attn_mask = torch.empty(
(B * max_len),
device=input_emb.device,
dtype=torch.bool,
)
BLOCK_M = 32
BLOCK_D = triton.next_power_of_2(hidden_dim)
grid = (
B,
triton.cdiv(max_len, BLOCK_M),
1,
)
pad_sequence_with_mask_kernel[grid](
input_emb,
offsets,
lengths,
output,
attn_mask,
max_len,
hidden_dim,
BLOCK_M=BLOCK_M,
BLOCK_D=BLOCK_D,
)
return B, output, attn_mask