Skip to content
Open
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
5 changes: 5 additions & 0 deletions python/sglang/srt/environ.py
Original file line number Diff line number Diff line change
Expand Up @@ -1531,6 +1531,11 @@ class Envs:
SGLANG_DSA_MQA_LOGITS_FREE_MEM_FRACTION = EnvFloat(0.2)
SGLANG_ENABLE_PCG_DSV2_DUAL_STREAM = EnvBool(False)
SGLANG_DSA_TOPK_BROADCAST = EnvBool(False)
# ROCm only: split the prefill indexer's logits + top-k across attn-TP
# ranks in interleaved row stripes and AllReduce(MAX) the -1-filled
# result.
SGLANG_DSA_INDEXER_M_SPLIT = EnvBool(False)
SGLANG_DSA_INDEXER_M_SPLIT_STRIPE = EnvInt(512)
SGLANG_DISABLE_DSA_INDEXER_FUSION = EnvBool(False)
# Opt-in perf path for --dsa-prefill-backend flashmla_sparse_q8: fuse the
# absorbed q bmm with the nope/rope concat + fp8 cast so q is written
Expand Down
183 changes: 183 additions & 0 deletions python/sglang/srt/layers/attention/dsa/dsa_indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union

import torch
import torch.distributed as dist
from einops import rearrange

from sglang.kernels.fused_op import BaseFusedOp
Expand Down Expand Up @@ -211,6 +212,9 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
# aiter's fp8_mqa_logits only compiles below 2 GiB of logits (buffer_store).
_MQA_LOGITS_MAX_BYTES_ROCM = 2**31 - 1
_mqa_logits_budget_bytes: Dict[int, int] = {}
# Attn-TP-wide MIN of the budget above; M-split derives row ownership
# from it, so every rank must see the same value.
_m_split_budget_bytes: Dict[int, int] = {}

@staticmethod
def _mqa_logits_free_mem_fraction() -> float:
Expand Down Expand Up @@ -255,6 +259,18 @@ def __init__(
self.cp_size = get_parallel().attn_cp_size
else:
self.cp_size = None
# ROCm-only: split the prefill indexer's logits + top-k across attn-TP
# ranks. See _get_topk_ragged_m_split.
self.dsa_indexer_m_split = _is_hip and envs.SGLANG_DSA_INDEXER_M_SPLIT.get()
self.dsa_indexer_m_split_stripe = max(
1, envs.SGLANG_DSA_INDEXER_M_SPLIT_STRIPE.get()
)
if self.dsa_indexer_m_split and layer_id == 0:
logger.info(
"DSA indexer M-split enabled: prefill indexer rows are striped "
"across attn-TP ranks (stripe=%d).",
self.dsa_indexer_m_split_stripe,
)
if _is_cuda:
self.sm_count = deep_gemm.get_num_sms()
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8)
Expand Down Expand Up @@ -1176,6 +1192,25 @@ def _get_topk_ragged(
q_offset, k_offset, device_index
)

if self.dsa_indexer_m_split:
m_split_result = self._get_topk_ragged_m_split(
q_fp8=q_fp8,
weights=weights,
kv_fp8=kv_fp8,
ks=ks,
ke=ke,
seq_lens_expanded=seq_lens_expanded,
token_to_batch_idx=token_to_batch_idx,
metadata=metadata,
topk_result=topk_result,
q_offset=q_offset,
k_offset=k_offset,
)
# None: nothing to split (single rank / too few rows), use the
# replicated paths below.
if m_split_result is not None:
return m_split_result

if not need_chunk:
assert q_fp8[:q_offset].shape[0] != 0
with self._with_real_sm_count():
Expand Down Expand Up @@ -1312,6 +1347,154 @@ def _get_topk_ragged(

return topk_result

def _get_topk_ragged_m_split(
self,
q_fp8: torch.Tensor,
weights: torch.Tensor,
kv_fp8: Tuple[torch.Tensor, torch.Tensor],
ks: torch.Tensor,
ke: torch.Tensor,
seq_lens_expanded: torch.Tensor,
token_to_batch_idx: Optional[torch.Tensor],
metadata: BaseIndexerMetadata,
topk_result: torch.Tensor,
q_offset: int,
k_offset: int,
) -> Optional[torch.Tensor]:
"""M-split prefill indexer (ROCm): each attn-TP rank scores an interleaved
stripe subset of the query rows, writes its top-k into a -1-filled
buffer, and AllReduce(MAX) recovers the full result on every rank.

The Indexer weights are replicated, so without this every rank redoes
the full [M x K] logits + top-k. Interleaved stripes keep the causal
workload balanced (~M^2/2/tp per rank).

Returns None when there is nothing to split so the caller falls back
to the replicated path.
"""
assert _is_hip, "DSA indexer M-split is ROCm-only"
tp_group = get_attn_tp_group()
tp_size = tp_group.world_size
tp_rank = tp_group.rank_in_group
if tp_size <= 1 or q_offset < tp_size:
return None

from aiter.ops.triton.fp8_mqa_logits import fp8_mqa_logits

assert seq_lens_expanded.shape[0] == q_offset, (
f"seq_lens_expanded length mismatch: {seq_lens_expanded.shape[0]} != {q_offset}"
)
kv, scale = kv_fp8
device = q_fp8.device

# Every stripe must stay under the logits budget (aiter's 2 GiB
# buffer_store limit on ROCm). The per-rank budget depends on each GPU's
# free memory, and a different stripe per rank would leave rows that no
# rank owns at -1, so the budget is the attn-TP-wide MIN.
budget_bytes = self._get_m_split_logits_budget_bytes(tp_group, device.index)
bytes_per_row = k_offset * self._MQA_LOGITS_BYTES_PER_ELEM
max_rows = max(1, budget_bytes // max(bytes_per_row, 1))
stripe = max(
1, min(self.dsa_indexer_m_split_stripe, q_offset // tp_size, max_rows)
)
block = tp_size * stripe

# Same per-row-range top-k contract as the chunked path: RAGGED uses the
# global offset slice, PAGED treats each token as a length-1 sequence.
global_topk_offset = metadata.attn_metadata.topk_indices_offset
cu_seqlens_q_full = None
if global_topk_offset is None:
cu_seqlens_q_full = torch.ones(q_offset, dtype=torch.int32, device=device)
else:
assert global_topk_offset.shape[0] >= q_offset, (
f"topk_indices_offset too short: {global_topk_offset.shape[0]} < {q_offset}"
)

# Rows owned by other ranks must be -1 so that MAX recovers them.
topk_result[:q_offset].fill_(-1)

for block_start in range(0, q_offset, block):
start = block_start + tp_rank * stripe
if start >= q_offset:
break
end = min(start + stripe, q_offset)

with self._with_real_sm_count():
# clean_logits=False: topk transform handles masking via ks/ke.
logits_chunk = fp8_mqa_logits(
q_fp8[start:end],
kv,
scale,
weights[start:end],
ks[start:end],
ke[start:end],
clean_logits=False,
)

lengths_chunk = seq_lens_expanded[start:end]
self._mask_init_and_local_tokens(logits_chunk, lengths_chunk, ks[start:end])

if global_topk_offset is not None:
topk_offset_chunk = global_topk_offset[start:end]
cu_seqlens_q_chunk = None
batch_idx_chunk = None
else:
topk_offset_chunk = None
cu_seqlens_q_chunk = cu_seqlens_q_full[start:end]
batch_idx_chunk = token_to_batch_idx[start:end]

topk_result[start:end] = metadata.topk_transform(
logits_chunk,
self.index_topk,
ks=ks[start:end],
cu_seqlens_q=cu_seqlens_q_chunk,
ke_offset=lengths_chunk,
batch_idx_list=batch_idx_chunk,
topk_indices_offset_override=topk_offset_chunk,
)

self._m_split_all_reduce_max(tp_group, topk_result[:q_offset])
return topk_result

def _get_m_split_logits_budget_bytes(self, tp_group, device_index: int) -> int:
cached = self._m_split_budget_bytes.get(device_index)
if cached is not None:
return cached
# Same cap as _should_chunk_mqa_logits: aiter's fp8_mqa_logits only
# compiles below 2 GiB of logits (buffer_store), regardless of free memory.
local_budget = min(
self._get_mqa_logits_budget_bytes(device_index),
self._MQA_LOGITS_MAX_BYTES_ROCM,
)
budget = torch.tensor(
[local_budget], dtype=torch.int64, device=f"cuda:{device_index}"
)
dist.all_reduce(budget, op=dist.ReduceOp.MIN, group=tp_group.device_group)
synced = max(1, int(budget.item()))
# Under capture the local budget is the static estimate and is not
# cached either; cache only the real free-memory value.
if not get_is_capture_mode():
self._m_split_budget_bytes[device_index] = synced
return synced

@staticmethod
def _m_split_all_reduce_max(tp_group, buf: torch.Tensor) -> None:
# GroupCoordinator.all_reduce is SUM-only, so MAX goes straight to the
# communicator. Mirrors _broadcast_indexer_topk_from_rank0_impl: PyNCCL
# under capture when available, process group otherwise.
tmp = buf if buf.is_contiguous() else buf.contiguous()
if (
tmp.device.type == "cuda"
and torch.cuda.is_current_stream_capturing()
and tp_group.pynccl_comm is not None
):
with tp_group.pynccl_comm.change_state(enable=True):
tp_group.pynccl_comm.all_reduce(tmp, op=dist.ReduceOp.MAX)
else:
dist.all_reduce(tmp, op=dist.ReduceOp.MAX, group=tp_group.device_group)
if tmp is not buf:
buf.copy_(tmp)

def _forward_cuda_k_only(
self,
x: torch.Tensor,
Expand Down