Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 2 additions & 23 deletions python/sglang/kernels/ops/attention/dsa/tilelang_kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -271,7 +271,6 @@ def sparse_attention_fwd_kernel_v1(
block_I=64,
num_stages=2,
threads=256,
return_lse=False,
):
assert (
dim == tilelang.math.next_power_of_2(dim) or dim % 64 == 0
Expand All @@ -297,7 +296,6 @@ def sparse_attention_fwd_kernel_v1(
q_shape = [batch, seq_len, num_heads, dim + tail_dim]
kv_shape = [batch, seq_len_kv, kv_group, dim + tail_dim]
o_shape = [batch, seq_len, num_heads, dim]
lse_shape = [batch, seq_len, num_heads]
indices_shape = [batch, seq_len, kv_group, topk]
indices_dtype = "int32"
dtype = "bfloat16"
Expand Down Expand Up @@ -325,7 +323,6 @@ def main(
Q: T.Tensor(q_shape, dtype), # type: ignore
KV: T.Tensor(kv_shape, dtype), # type: ignore
Indices: T.Tensor(indices_shape, indices_dtype), # type: ignore
LSE: T.Tensor(lse_shape, accum_dtype), # type: ignore
Output: T.Tensor(o_shape, dtype), # type: ignore
):
with T.Kernel(seq_len * REPLICATE_H, batch, kv_group, threads=threads) as (
Expand Down Expand Up @@ -426,8 +423,6 @@ def main(
acc_o[h_i, d_i] /= sumexp[h_i]
for h_i in T.Parallel(H_per_block):
sumexp[h_i] = T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale
if return_lse:
T.copy(sumexp, LSE[b_i, s_i, H0:H1])

T.copy(acc_o, O_shared)
T.copy(acc_o, Output[b_i, s_i, H0:H1, :])
Expand Down Expand Up @@ -459,7 +454,6 @@ def sparse_attention_fwd_kernel_v2(
kv_group: int = 1,
sm_scale: Optional[float] = None,
block_I: int = 64,
return_lse: bool = False,
):
assert dim == tilelang.math.next_power_of_2(
dim
Expand All @@ -483,7 +477,6 @@ def sparse_attention_fwd_kernel_v2(
q_shape = [batch, qo_len, num_heads, dim + tail_dim]
kv_shape = [batch, num_pages, kv_group, dim + tail_dim]
o_shape = [batch, qo_len, num_heads, dim]
lse_shape = [batch, qo_len, num_heads]
indices_shape = [batch, qo_len, kv_group, topk]

indices_dtype = "int32"
Expand Down Expand Up @@ -512,7 +505,6 @@ def main(
Q: T.Tensor(q_shape, dtype), # type: ignore
KV: T.Tensor(kv_shape, dtype), # type: ignore
Indices: T.Tensor(indices_shape, indices_dtype), # type: ignore
LSE: T.Tensor(lse_shape, accum_dtype), # type: ignore
Output: T.Tensor(o_shape, dtype), # type: ignore
):
"""
Expand Down Expand Up @@ -680,8 +672,6 @@ def main(
acc_o_l[h_i, d_i] /= sumexp[h_i]
for h_i in T.Parallel(H_per_block):
sumexp[h_i] = T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale
if return_lse:
T.copy(sumexp, LSE[b_i, s_i, H0:H1])
T.copy(acc_o_l, O_shared_l)
T.copy(O_shared_l, Output[b_i, s_i, H0:H1, 0 : D // 2])
elif tx >= 128 and tx < 256:
Expand Down Expand Up @@ -1339,7 +1329,6 @@ def tilelang_sparse_fwd(
indices: torch.Tensor,
sm_scale: float,
d_v: int = 512,
return_lse: bool = False,
) -> torch.Tensor:
assert q.dim() == 3 and kv.dim() == 3 and indices.dim() == 3
num_heads = q.shape[1]
Expand All @@ -1349,7 +1338,6 @@ def tilelang_sparse_fwd(
assert topk % 64 == 0, "topk must be padded to a multiple of 64"

if _is_hip:
assert not return_lse, "tilelang HIP sparse fwd does not return LSE"
is_fp8_kv = kv.dtype in (torch.float8_e4m3fn, torch.float8_e4m3fnuz)
if is_fp8_kv:
if q.dtype != kv.dtype:
Expand Down Expand Up @@ -1406,17 +1394,8 @@ def tilelang_sparse_fwd(
if tail_dim == 0
else sparse_attention_fwd_kernel_v2
)
kernel = kernel_factory(
num_heads, d_v, tail_dim, topk, sm_scale=sm_scale, return_lse=return_lse
)
# Caller-allocated LSE (in-place kernel arg): written only by kernels
# traced with return_lse=True, but the prim_func signature always has it.
lse = torch.empty(
(1, q.shape[0], num_heads), dtype=torch.float32, device=q.device
)
out = kernel(q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0), lse) # type: ignore
if return_lse:
return out, lse
kernel = kernel_factory(num_heads, d_v, tail_dim, topk, sm_scale=sm_scale)
out = kernel(q.unsqueeze(0), kv.unsqueeze(0), indices.unsqueeze(0)) # type: ignore
return out


Expand Down
28 changes: 0 additions & 28 deletions python/sglang/srt/layers/attention/dsa_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,10 +154,6 @@ def _should_all_gather_dsa_trtllm_fp8_kv(
return save_kv_cache and cos_sin_cache is not None and dsa_prefill_cp


def _should_return_dsa_dcp_lse(*, forward_mode: ForwardMode, dcp_enabled: bool) -> bool:
return dcp_enabled and (forward_mode.is_decode() or forward_mode.is_target_verify())


def materialize_full_kv_cp(
attn_mla,
forward_batch: ForwardBatch,
Expand Down Expand Up @@ -3094,10 +3090,6 @@ def forward_extend(
page_table_1=page_table_1,
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
return_lse=_should_return_dsa_dcp_lse(
forward_mode=forward_batch.forward_mode,
dcp_enabled=get_parallel().dcp_enabled,
),
)
elif dsa_impl in ("flashmla_sparse", "flashmla_sparse_q8"):
if topk_transform_method == TopkTransformMethod.RAGGED:
Expand Down Expand Up @@ -3396,10 +3388,6 @@ def forward_decode(
page_table_1=page_table_1,
sm_scale=layer.scaling,
v_head_dim=layer.v_head_dim,
return_lse=_should_return_dsa_dcp_lse(
forward_mode=forward_batch.forward_mode,
dcp_enabled=get_parallel().dcp_enabled,
),
)
elif dsa_impl == "fa3":
return self._forward_fa3(
Expand Down Expand Up @@ -4015,7 +4003,6 @@ def _forward_tilelang(
v_head_dim: int,
page_table_1: torch.Tensor,
sm_scale: float,
return_lse: bool = False,
) -> torch.Tensor:
from sglang.kernels.ops.attention.dsa.tilelang_kernel import tilelang_sparse_fwd

Expand All @@ -4032,17 +4019,6 @@ def _forward_tilelang(
dim=-1,
)

if return_lse:
out, lse = tilelang_sparse_fwd(
q=q_all,
kv=kv_cache,
indices=page_table_1.unsqueeze(1),
sm_scale=sm_scale,
d_v=v_head_dim,
return_lse=True,
)
# [1, tokens, H] -> [tokens, H], the [B, H] dcp_pack_a2a_send expects.
return out, lse.squeeze(0)
return tilelang_sparse_fwd(
q=q_all,
kv=kv_cache,
Expand Down Expand Up @@ -4400,10 +4376,6 @@ def _forward_trtllm(
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
sparse_mla_top_k_lens=sparse_mla_top_k_lens,
multi_ctas_kv_counter_buffer=self._multi_ctas_kv_counter_buffer,
return_lse=_should_return_dsa_dcp_lse(
forward_mode=forward_batch.forward_mode,
dcp_enabled=get_parallel().dcp_enabled,
),
)

return out
Expand Down
5 changes: 1 addition & 4 deletions python/sglang/srt/layers/dcp/comm.py
Original file line number Diff line number Diff line change
Expand Up @@ -305,16 +305,13 @@ def all_gather_kv_cache_for_mla_extend(
# all gather kv cache and re-org to query orders
def all_gather_kv_cache_for_dcp(
prefix_kv_a: torch.Tensor,
prefix_k_pe: Optional[torch.Tensor],
prefix_k_pe: torch.Tensor,
prefix_kv_lens_cpu: torch.Tensor,
prefix_starts_cpu: torch.Tensor = None,
):
"""
prefix_kv_a and prefix_k_pe should have same shape, expect for last dim
"""
if prefix_k_pe is None:
prefix_k_pe = prefix_kv_a.new_empty((*prefix_kv_a.shape[:-1], 0))

parallel = get_parallel()
if not parallel.dcp_enabled:
return torch.cat([prefix_kv_a, prefix_k_pe], dim=-1)
Expand Down
11 changes: 1 addition & 10 deletions python/sglang/srt/mem_cache/kv_cache_configurator.py
Original file line number Diff line number Diff line change
Expand Up @@ -361,12 +361,6 @@ def loc_space_scale(self) -> int:
def pool_page_size(self) -> int:
return get_schedule().page_size * self.loc_space_scale

def _dsa_pool_geometry(self, max_total_num_tokens: int) -> tuple[int, int]:
physical_page_size = get_schedule().page_size
# Physical page stays 64 for CUDA DSA kernels; only the token pool is grown.
pool_size = max_total_num_tokens + self.pool_page_size - physical_page_size
return pool_size, physical_page_size

def _derive_pool_sizes(self, *, config: MemoryPoolConfig) -> _PoolSizes:
max_total_num_tokens = config.max_total_num_tokens
max_running_requests = config.max_running_requests
Expand Down Expand Up @@ -1519,9 +1513,6 @@ def _build_dsa_kv_pool(
) -> KVCache:
from sglang.srt.layers.cp.utils import get_glm_dsa_cp_layer_shard_info

max_total_num_tokens, pool_page_size = self._dsa_pool_geometry(
max_total_num_tokens
)
(
dsa_cp_layer_shard_rank,
dsa_cp_layer_shard_size,
Expand Down Expand Up @@ -1554,7 +1545,7 @@ def _build_dsa_kv_pool(
]
token_to_kv_pool = PoolCls(
max_total_num_tokens,
page_size=pool_page_size,
page_size=self.pool_page_size,
dtype=self.kv_cache_dtype,
kv_lora_rank=self.model_config.kv_lora_rank,
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
Expand Down
23 changes: 6 additions & 17 deletions python/sglang/srt/model_executor/pool_configurator.py
Original file line number Diff line number Diff line change
Expand Up @@ -228,23 +228,16 @@ def __init__(self, kvc: KVCacheConfigurator):
target_kv_num_layers = get_glm_dsa_layer_split_effective_num_layers(
kvc, num_layers
)
# Draft pools are DCP-replicated, not sharded: budget all copies.
dcp_size = kvc.ps.attn_dcp_size
draft_kv_size = (
int(target_kv_size * draft_num_layers / target_kv_num_layers)
* dcp_size
draft_kv_size = int(
target_kv_size * draft_num_layers / target_kv_num_layers
)
draft_indexer_size = (
self._compute_dsa_indexer_cell_size(
kvc=kvc,
num_layers=draft_num_layers,
allocate_all_layers=True,
)
* dcp_size
draft_indexer_size = self._compute_dsa_indexer_cell_size(
kvc=kvc,
num_layers=draft_num_layers,
allocate_all_layers=True,
)
self._cell_size += draft_kv_size + draft_indexer_size
else:
draft_num_layers *= kvc.ps.attn_dcp_size
self._cell_size = int(
self._cell_size * (1 + draft_num_layers / int(num_layers))
)
Expand Down Expand Up @@ -586,10 +579,6 @@ def __init__(self, kvc: KVCacheConfigurator):
- self._draft_swa_layers_num
- self._draft_swa_full_layers_num
)
dcp_size = kvc.ps.attn_dcp_size
self._draft_swa_layers_num *= dcp_size
self._draft_swa_full_layers_num *= dcp_size
self._draft_full_layers_num *= dcp_size

self._draft_cell_size = _dflash_draft_cell_size(kvc)

Expand Down
8 changes: 2 additions & 6 deletions python/sglang/srt/model_executor/runner/eager_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -295,12 +295,8 @@ def _execute_extend(
or cp_v2_active
or forward_batch.forward_mode.is_target_verify()
):
if (
model_runner.ps.attn_dcp_size > 1
and not forward_batch.forward_mode.is_target_verify()
and hasattr(
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
)
if model_runner.ps.attn_dcp_size > 1 and hasattr(
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
):
# prepare kv cache buffer for dcp to gather kv cache
forward_batch.attn_dcp_metadata = (
Expand Down
29 changes: 0 additions & 29 deletions python/sglang/srt/models/glm5_next.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,6 @@
get_attn_tp_context,
)
from sglang.srt.layers.communicator_mhc import MHCLayerCommunicator
from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
ColumnParallelBatchedLinear,
Expand Down Expand Up @@ -1254,34 +1253,6 @@ def set_dflash_layers_to_capture(self, layer_ids: List[int]):
# Capturing before layer k + 1 gives the completed output of layer k.
self.model.layers_to_capture = [val + 1 for val in layer_ids]

def prepare_context_parallel_metadata_for_dcp(
self,
seq_lens: torch.Tensor,
extend_prefix_lens: torch.Tensor,
extend_prefix_lens_cpu: torch.Tensor,
extend_seq_lens: torch.Tensor,
req_pool_indices: torch.Tensor,
req_to_token: torch.Tensor,
seq_lens_sum: int,
kv_buffer_shape: torch.Size,
kv_cache_dtype,
kv_cache_device,
create_chunked_prefix_cache_kv_indices_fn,
):
return prepare_decode_context_parallel_metadata(
seq_lens=seq_lens,
extend_prefix_lens=extend_prefix_lens,
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
extend_seq_lens=extend_seq_lens,
req_pool_indices=req_pool_indices,
req_to_token=req_to_token,
seq_lens_sum=seq_lens_sum,
kv_buffer_shape=kv_buffer_shape,
kv_cache_dtype=kv_cache_dtype,
kv_cache_device=kv_cache_device,
create_chunked_prefix_cache_kv_indices_fn=create_chunked_prefix_cache_kv_indices_fn,
)

def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
return pattern.pad_input_tokens(input_ids, mm_inputs)
Expand Down
30 changes: 0 additions & 30 deletions test/registered/cp/test_dsa_trtllm_fp8_cp.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,7 @@

from sglang.srt.layers.attention.dsa_backend import (
_should_all_gather_dsa_trtllm_fp8_kv,
_should_return_dsa_dcp_lse,
)
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase

Expand Down Expand Up @@ -49,34 +47,6 @@ def test_non_cp_and_no_cache_paths_do_not_gather(self):
)
)

def test_dcp_decode_and_verify_request_lse(self):
self.assertTrue(
_should_return_dsa_dcp_lse(
forward_mode=ForwardMode.DECODE,
dcp_enabled=True,
)
)
self.assertTrue(
_should_return_dsa_dcp_lse(
forward_mode=ForwardMode.TARGET_VERIFY,
dcp_enabled=True,
)
)

def test_non_dcp_and_prefill_do_not_request_lse(self):
self.assertFalse(
_should_return_dsa_dcp_lse(
forward_mode=ForwardMode.DECODE,
dcp_enabled=False,
)
)
self.assertFalse(
_should_return_dsa_dcp_lse(
forward_mode=ForwardMode.EXTEND,
dcp_enabled=True,
)
)


if __name__ == "__main__":
unittest.main()
24 changes: 0 additions & 24 deletions test/registered/dcp/test_dcp_layout_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -437,30 +437,6 @@ def test_configurator_scales_only_the_virtual_dcp_allocator(self):
self.assertEqual(dcp4_allocator.page_size, 256)
self.assertEqual(dcp4_allocator.num_pages, 16)

def test_dsa_draft_pool_preserves_physical_page_and_backs_virtual_tail(self):
physical_page_size = 64
max_total_num_tokens = 4096
self._sa_override = rc.get_context().override_server_args(
page_size=physical_page_size
)
self._sa_override.install()
self.addCleanup(self._sa_override.restore)

dcp1 = SimpleNamespace(pool_page_size=physical_page_size)
dcp4_draft = SimpleNamespace(pool_page_size=physical_page_size * 4)

self.assertEqual(
KVCacheConfigurator._dsa_pool_geometry(dcp1, max_total_num_tokens),
(max_total_num_tokens, physical_page_size),
)
self.assertEqual(
KVCacheConfigurator._dsa_pool_geometry(dcp4_draft, max_total_num_tokens),
(
max_total_num_tokens + physical_page_size * 3,
physical_page_size,
),
)

def test_live_cell_and_page_ownership_formulas(self):
dcp_size = 4
physical_page_size = 64
Expand Down
Loading
Loading