Skip to content
Merged
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
61 changes: 41 additions & 20 deletions python/sglang/srt/layers/attention/dsa_backend.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from __future__ import annotations

import logging
from contextlib import nullcontext
from dataclasses import dataclass
from typing import (
TYPE_CHECKING,
Expand All @@ -15,6 +16,7 @@
import torch

from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
from sglang.srt.constants import GPU_MEMORY_TYPE_CUDA_GRAPH
from sglang.srt.runtime_context import get_parallel

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -70,6 +72,7 @@
is_sm100_supported,
print_warning_once,
)
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter

# Opt-in (default off): route the fp8 sparse-MLA prefill path through the Triton
# per-query flash kernel instead of TileLang. Validated on gfx950 (GLM-5.1 @
Expand Down Expand Up @@ -365,6 +368,7 @@ def __init__(
)
self.dsa_index_topk = get_dsa_index_topk(model_runner.model_config.hf_config)
self.max_context_len = model_runner.model_config.context_len
self.enable_memory_saver = model_runner.server_args.enable_memory_saver
self.num_q_heads = (
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
)
Expand Down Expand Up @@ -1129,6 +1133,16 @@ def _cal_indexer_k_start_end(
token_to_batch_idx = dsa_cp_round_robin_split_data(token_to_batch_idx)
return (ks, ke), token_to_batch_idx

def _cuda_graph_memory_region(self):
"""Region whose tensors are released while the engine is paused."""
adapter = TorchMemorySaverAdapter.create(
enable=self.enable_memory_saver
and envs.SGLANG_MEMORY_SAVER_CUDA_GRAPH.get()
)
if not adapter.enabled:
return nullcontext()
return adapter.region(tag=GPU_MEMORY_TYPE_CUDA_GRAPH)

def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
"""Initialize CUDA graph state for the attention backend.

Expand Down Expand Up @@ -1164,6 +1178,31 @@ def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
)

max_ctx_len = self.req_to_token.shape[1]
# The wide page_table (max_num_tokens x max_ctx_len int32) and the
# flashmla metadata are rewritten in full before every replay, so they
# can live in the pausable cuda-graph region. The others are written
# once at init and would come back zeroed after a resume.
with self._cuda_graph_memory_region():
page_table = (
None
if self.dsa_drop_wide_page_table
else torch.zeros(
max_num_tokens,
max_ctx_len,
dtype=torch.int32,
device=self.device,
)
)
flashmla_metadata = (
self._compute_flashmla_metadata(
cache_seqlens=torch.ones(
max_num_tokens, dtype=torch.int32, device=self.device
),
seq_len_q=1,
)
if self.dsa_decode_impl == "flashmla_kv"
else None
)
self.decode_cuda_graph_metadata: Dict = {
"cache_seqlens": torch.ones(
max_num_tokens, dtype=torch.int32, device=self.device
Expand All @@ -1190,26 +1229,8 @@ def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
if self.dsa_drop_wide_page_table
else None
),
"page_table": (
None
if self.dsa_drop_wide_page_table
else torch.zeros(
max_num_tokens,
max_ctx_len,
dtype=torch.int32,
device=self.device,
)
),
"flashmla_metadata": (
self._compute_flashmla_metadata(
cache_seqlens=torch.ones(
max_num_tokens, dtype=torch.int32, device=self.device
),
seq_len_q=1,
)
if self.dsa_decode_impl == "flashmla_kv"
else None
),
"page_table": page_table,
"flashmla_metadata": flashmla_metadata,
}

def _build_forward_metadata_cuda_graph(
Expand Down
Loading