From 888594333e6034c07536d15509a6da051d405f4f Mon Sep 17 00:00:00 2001 From: kk <43161300+kkHuang-amd@users.noreply.github.com> Date: Wed, 17 Dec 2025 17:25:06 +0800 Subject: [PATCH] Fix gpu-fault when running mtp in eager mode (#15233) --- .../srt/layers/attention/aiter_backend.py | 104 +++++++++++++----- python/sglang/srt/layers/attention/utils.py | 98 +++++++++++++++++ 2 files changed, 172 insertions(+), 30 deletions(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 4cca37d97..a2c4a2e19 100644 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -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=}" diff --git a/python/sglang/srt/layers/attention/utils.py b/python/sglang/srt/layers/attention/utils.py index a8d8cc4b4..dc9973788 100644 --- a/python/sglang/srt/layers/attention/utils.py +++ b/python/sglang/srt/layers/attention/utils.py @@ -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