Fix gpu-fault when running mtp in eager mode (#15233)
This commit is contained in:
@@ -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=}"
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user