[Spec] Remove dead kernel params; fix stale comment in trtllm_mla (#25010)
This commit is contained in:
@@ -52,7 +52,6 @@ def fused_recurrent_kda_fwd(
|
|||||||
inplace_final_state: bool = True,
|
inplace_final_state: bool = True,
|
||||||
cu_seqlens: torch.LongTensor | None = None,
|
cu_seqlens: torch.LongTensor | None = None,
|
||||||
# ssm_state_indices: torch.Tensor | None = None,
|
# ssm_state_indices: torch.Tensor | None = None,
|
||||||
num_accepted_tokens: torch.Tensor | None = None,
|
|
||||||
use_qk_l2norm_in_kernel: bool = False,
|
use_qk_l2norm_in_kernel: bool = False,
|
||||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||||
B, T, H, K, V = *k.shape, v.shape[-1]
|
B, T, H, K, V = *k.shape, v.shape[-1]
|
||||||
@@ -92,7 +91,6 @@ def fused_recurrent_kda_fwd(
|
|||||||
ht=final_state,
|
ht=final_state,
|
||||||
cu_seqlens=cu_seqlens,
|
cu_seqlens=cu_seqlens,
|
||||||
# ssm_state_indices=ssm_state_indices,
|
# ssm_state_indices=ssm_state_indices,
|
||||||
# num_accepted_tokens=num_accepted_tokens,
|
|
||||||
scale=scale,
|
scale=scale,
|
||||||
# N=N,
|
# N=N,
|
||||||
T=T,
|
T=T,
|
||||||
@@ -155,7 +153,6 @@ def fused_recurrent_kda(
|
|||||||
inplace_final_state=inplace_final_state,
|
inplace_final_state=inplace_final_state,
|
||||||
cu_seqlens=cu_seqlens,
|
cu_seqlens=cu_seqlens,
|
||||||
# ssm_state_indices=ssm_state_indices,
|
# ssm_state_indices=ssm_state_indices,
|
||||||
num_accepted_tokens=None,
|
|
||||||
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
|
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
|
||||||
)
|
)
|
||||||
return o, final_state
|
return o, final_state
|
||||||
|
|||||||
@@ -18,8 +18,6 @@ def _fused_mamba_state_scatter_with_mask_kernel(
|
|||||||
# Raw index arrays (before index_select)
|
# Raw index arrays (before index_select)
|
||||||
dst_indices_raw_ptr, # [total_requests] - state_indices_tensor
|
dst_indices_raw_ptr, # [total_requests] - state_indices_tensor
|
||||||
step_indices_raw_ptr, # [total_requests] - accepted_steps or mamba_steps_to_track
|
step_indices_raw_ptr, # [total_requests] - accepted_steps or mamba_steps_to_track
|
||||||
# Total number of requests
|
|
||||||
total_requests,
|
|
||||||
elem_per_entry: tl.constexpr,
|
elem_per_entry: tl.constexpr,
|
||||||
src_layer_stride,
|
src_layer_stride,
|
||||||
src_req_stride,
|
src_req_stride,
|
||||||
@@ -176,7 +174,6 @@ def fused_mamba_state_scatter_with_mask(
|
|||||||
dst,
|
dst,
|
||||||
dst_indices_raw,
|
dst_indices_raw,
|
||||||
step_indices_raw,
|
step_indices_raw,
|
||||||
total_requests,
|
|
||||||
elem_per_entry,
|
elem_per_entry,
|
||||||
src_layer_stride,
|
src_layer_stride,
|
||||||
src_req_stride,
|
src_req_stride,
|
||||||
|
|||||||
@@ -59,7 +59,7 @@ def pad_draft_extend_query_kernel(
|
|||||||
q_ptr, # Input query tensor [total_seq_len, num_heads, head_dim]
|
q_ptr, # Input query tensor [total_seq_len, num_heads, head_dim]
|
||||||
padded_q_ptr, # Output padded query tensor [batch_size, max_seq_len, num_heads, head_dim]
|
padded_q_ptr, # Output padded query tensor [batch_size, max_seq_len, num_heads, head_dim]
|
||||||
seq_lens_q_ptr, # Sequence lengths for each sequence [batch_size]
|
seq_lens_q_ptr, # Sequence lengths for each sequence [batch_size]
|
||||||
cumsum_ptr, # Cumulative sum of accept lengths [batch_size + 1]
|
cumsum_ptr, # Cumulative sum of sequence lengths [batch_size + 1]
|
||||||
batch_size,
|
batch_size,
|
||||||
max_seq_len,
|
max_seq_len,
|
||||||
num_heads,
|
num_heads,
|
||||||
@@ -78,7 +78,7 @@ def pad_draft_extend_query_kernel(
|
|||||||
if batch_id >= batch_size:
|
if batch_id >= batch_size:
|
||||||
return
|
return
|
||||||
|
|
||||||
# Load accept length for this batch
|
# Load sequence length for this batch
|
||||||
seq_len = tl.load(seq_lens_q_ptr + batch_id)
|
seq_len = tl.load(seq_lens_q_ptr + batch_id)
|
||||||
|
|
||||||
if seq_pos >= seq_len:
|
if seq_pos >= seq_len:
|
||||||
|
|||||||
@@ -45,20 +45,20 @@ def _ref_update_like(
|
|||||||
intermediate_conv,
|
intermediate_conv,
|
||||||
*,
|
*,
|
||||||
state_indices_tensor,
|
state_indices_tensor,
|
||||||
accepted_steps,
|
step_indices_raw,
|
||||||
mamba_track_indices=None,
|
mamba_track_indices=None,
|
||||||
mamba_steps_to_track=None,
|
mamba_steps_to_track=None,
|
||||||
):
|
):
|
||||||
"""Reference implementation using PyTorch advanced indexing for correctness verification."""
|
"""Reference implementation using PyTorch advanced indexing for correctness verification."""
|
||||||
request_number = accepted_steps.shape[0]
|
total_requests = step_indices_raw.shape[0]
|
||||||
intermediate_state_indices = torch.arange(
|
intermediate_state_indices = torch.arange(
|
||||||
request_number, dtype=torch.int32, device=accepted_steps.device
|
total_requests, dtype=torch.int32, device=step_indices_raw.device
|
||||||
)
|
)
|
||||||
|
|
||||||
valid_mask = accepted_steps >= 0
|
valid_mask = step_indices_raw >= 0
|
||||||
dst_state_indices = state_indices_tensor[valid_mask].to(torch.int64)
|
dst_state_indices = state_indices_tensor[valid_mask].to(torch.int64)
|
||||||
src_state_indices = intermediate_state_indices[valid_mask].to(torch.int64)
|
src_state_indices = intermediate_state_indices[valid_mask].to(torch.int64)
|
||||||
last_steps = accepted_steps[valid_mask].to(torch.int64)
|
last_steps = step_indices_raw[valid_mask].to(torch.int64)
|
||||||
|
|
||||||
# Only scatter if there are valid indices (but don't early return -
|
# Only scatter if there are valid indices (but don't early return -
|
||||||
# mamba_track_indices processing is independent)
|
# mamba_track_indices processing is independent)
|
||||||
@@ -110,7 +110,7 @@ def _fused_update_like(
|
|||||||
intermediate_conv,
|
intermediate_conv,
|
||||||
*,
|
*,
|
||||||
state_indices_tensor,
|
state_indices_tensor,
|
||||||
accepted_steps,
|
step_indices_raw,
|
||||||
mamba_track_indices=None,
|
mamba_track_indices=None,
|
||||||
mamba_steps_to_track=None,
|
mamba_steps_to_track=None,
|
||||||
):
|
):
|
||||||
@@ -120,13 +120,13 @@ def _fused_update_like(
|
|||||||
ssm_states,
|
ssm_states,
|
||||||
intermediate_ssm,
|
intermediate_ssm,
|
||||||
state_indices_tensor,
|
state_indices_tensor,
|
||||||
accepted_steps,
|
step_indices_raw,
|
||||||
)
|
)
|
||||||
fused_mamba_state_scatter_with_mask(
|
fused_mamba_state_scatter_with_mask(
|
||||||
conv_states,
|
conv_states,
|
||||||
intermediate_conv,
|
intermediate_conv,
|
||||||
state_indices_tensor,
|
state_indices_tensor,
|
||||||
accepted_steps,
|
step_indices_raw,
|
||||||
)
|
)
|
||||||
|
|
||||||
if mamba_track_indices is not None:
|
if mamba_track_indices is not None:
|
||||||
@@ -199,10 +199,10 @@ class TestMambaStateScatterCorrectness(unittest.TestCase):
|
|||||||
:B
|
:B
|
||||||
].to(torch.int32)
|
].to(torch.int32)
|
||||||
|
|
||||||
accepted_steps = torch.randint(0, D, (B,), device=device, dtype=torch.int64)
|
step_indices_raw = torch.randint(0, D, (B,), device=device, dtype=torch.int64)
|
||||||
# set ~10% invalid
|
# set ~10% invalid
|
||||||
invalid = torch.rand((B,), device=device) < 0.1
|
invalid = torch.rand((B,), device=device) < 0.1
|
||||||
accepted_steps[invalid] = -1
|
step_indices_raw[invalid] = -1
|
||||||
|
|
||||||
# Optional track update
|
# Optional track update
|
||||||
mamba_track_indices = torch.randperm(C, device=device, dtype=torch.int64)[:B]
|
mamba_track_indices = torch.randperm(C, device=device, dtype=torch.int64)[:B]
|
||||||
@@ -223,7 +223,7 @@ class TestMambaStateScatterCorrectness(unittest.TestCase):
|
|||||||
conv_ref,
|
conv_ref,
|
||||||
intermediate_conv,
|
intermediate_conv,
|
||||||
state_indices_tensor=state_indices_tensor,
|
state_indices_tensor=state_indices_tensor,
|
||||||
accepted_steps=accepted_steps,
|
step_indices_raw=step_indices_raw,
|
||||||
mamba_track_indices=mamba_track_indices,
|
mamba_track_indices=mamba_track_indices,
|
||||||
mamba_steps_to_track=mamba_steps_to_track,
|
mamba_steps_to_track=mamba_steps_to_track,
|
||||||
)
|
)
|
||||||
@@ -233,7 +233,7 @@ class TestMambaStateScatterCorrectness(unittest.TestCase):
|
|||||||
conv_fused,
|
conv_fused,
|
||||||
intermediate_conv,
|
intermediate_conv,
|
||||||
state_indices_tensor=state_indices_tensor,
|
state_indices_tensor=state_indices_tensor,
|
||||||
accepted_steps=accepted_steps,
|
step_indices_raw=step_indices_raw,
|
||||||
mamba_track_indices=mamba_track_indices,
|
mamba_track_indices=mamba_track_indices,
|
||||||
mamba_steps_to_track=mamba_steps_to_track,
|
mamba_steps_to_track=mamba_steps_to_track,
|
||||||
)
|
)
|
||||||
@@ -290,10 +290,10 @@ class TestMambaStateScatterPerf(unittest.TestCase):
|
|||||||
state_indices_tensor = torch.randperm(C, device=device, dtype=torch.int64)[
|
state_indices_tensor = torch.randperm(C, device=device, dtype=torch.int64)[
|
||||||
:B
|
:B
|
||||||
].to(torch.int32)
|
].to(torch.int32)
|
||||||
accepted_steps = torch.randint(0, D, (B,), device=device, dtype=torch.int64)
|
step_indices_raw = torch.randint(0, D, (B,), device=device, dtype=torch.int64)
|
||||||
if invalid_ratio > 0:
|
if invalid_ratio > 0:
|
||||||
invalid = torch.rand((B,), device=device) < invalid_ratio
|
invalid = torch.rand((B,), device=device) < invalid_ratio
|
||||||
accepted_steps[invalid] = -1
|
step_indices_raw[invalid] = -1
|
||||||
|
|
||||||
mamba_track_indices = None
|
mamba_track_indices = None
|
||||||
mamba_steps_to_track = None
|
mamba_steps_to_track = None
|
||||||
@@ -314,7 +314,7 @@ class TestMambaStateScatterPerf(unittest.TestCase):
|
|||||||
conv_states,
|
conv_states,
|
||||||
intermediate_conv,
|
intermediate_conv,
|
||||||
state_indices_tensor=state_indices_tensor,
|
state_indices_tensor=state_indices_tensor,
|
||||||
accepted_steps=accepted_steps,
|
step_indices_raw=step_indices_raw,
|
||||||
mamba_track_indices=mamba_track_indices,
|
mamba_track_indices=mamba_track_indices,
|
||||||
mamba_steps_to_track=mamba_steps_to_track,
|
mamba_steps_to_track=mamba_steps_to_track,
|
||||||
)
|
)
|
||||||
@@ -326,7 +326,7 @@ class TestMambaStateScatterPerf(unittest.TestCase):
|
|||||||
conv_states,
|
conv_states,
|
||||||
intermediate_conv,
|
intermediate_conv,
|
||||||
state_indices_tensor=state_indices_tensor,
|
state_indices_tensor=state_indices_tensor,
|
||||||
accepted_steps=accepted_steps,
|
step_indices_raw=step_indices_raw,
|
||||||
mamba_track_indices=mamba_track_indices,
|
mamba_track_indices=mamba_track_indices,
|
||||||
mamba_steps_to_track=mamba_steps_to_track,
|
mamba_steps_to_track=mamba_steps_to_track,
|
||||||
)
|
)
|
||||||
@@ -339,7 +339,7 @@ class TestMambaStateScatterPerf(unittest.TestCase):
|
|||||||
ref_ms = _time_cuda_ms(ref_fn)
|
ref_ms = _time_cuda_ms(ref_fn)
|
||||||
fused_ms = _time_cuda_ms(fused_fn)
|
fused_ms = _time_cuda_ms(fused_fn)
|
||||||
|
|
||||||
num_valid = int((accepted_steps >= 0).sum().item())
|
num_valid = int((step_indices_raw >= 0).sum().item())
|
||||||
ratio = fused_ms / ref_ms if ref_ms > 0 else float("inf")
|
ratio = fused_ms / ref_ms if ref_ms > 0 else float("inf")
|
||||||
speedup = ref_ms / fused_ms if fused_ms > 0 else float("inf")
|
speedup = ref_ms / fused_ms if fused_ms > 0 else float("inf")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user