diff --git a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py index d6c499df0..42dbc1589 100644 --- a/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py +++ b/python/sglang/srt/layers/attention/nsa/index_buf_accessor.py @@ -4,6 +4,8 @@ import torch import triton import triton.language as tl +from sglang.srt.utils import is_hip + if TYPE_CHECKING: from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool @@ -347,12 +349,17 @@ def _set_k_and_s_triton( raise ValueError( f"index_k_scale must be 1D or 2D, got shape {index_k_scale.shape}" ) - - assert buf_numel_per_page == 64 * (128 + 4) + if is_hip(): + assert buf_numel_per_page == 1 * (128 + 4) + else: + assert buf_numel_per_page == 64 * (128 + 4) assert num_tokens_to_write == num_tokens_to_write_ == num_tokens_to_write__ assert index_head_dim == 128 assert scale_dim == 1 - assert page_size == 64 + if is_hip(): + assert page_size == 1 + else: + assert page_size == 64 assert buf.dtype == torch.uint8 assert loc.dtype == torch.int64, f"{loc.dtype=}" # can be int32 diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 9ed629967..b80d3bdc9 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -12,14 +12,16 @@ from sglang.srt.layers.utils import MultiPlatformOp from sglang.srt.utils import add_prefix, ceil_align, is_cuda, is_hip, is_npu global _use_multi_stream - -if is_cuda(): +_is_cuda = is_cuda() +_is_hip = is_hip() +_is_npu = is_npu() +if _is_cuda: try: import deep_gemm except ImportError as e: deep_gemm = e -if is_npu(): +if _is_npu: import custom_ops # noqa: F401 import torch_npu from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream @@ -42,7 +44,8 @@ from sglang.srt.server_args import get_global_server_args if TYPE_CHECKING: from sglang.srt.mem_cache.memory_pool import NSATokenToKVPool -DUAL_STREAM_TOKEN_THRESHOLD = 1024 if is_cuda() else 0 + +DUAL_STREAM_TOKEN_THRESHOLD = 1024 if _is_cuda else 0 class BaseIndexerMetadata(ABC): @@ -59,6 +62,13 @@ class BaseIndexerMetadata(ABC): The page size of the table is 64. """ + @abstractmethod + def get_page_table_1(self) -> torch.Tensor: + """ + Return: (batch_size, num_blocks) int32, page table. + The page size of the table is 1. + """ + @abstractmethod def get_seqlens_expanded(self) -> torch.Tensor: """ @@ -101,7 +111,11 @@ class BaseIndexerMetadata(ABC): def rotate_activation(x: torch.Tensor) -> torch.Tensor: assert x.dtype == torch.bfloat16 - from sgl_kernel import hadamard_transform + # from sgl_kernel import hadamard_transform + if _is_hip: + from fast_hadamard_transform import hadamard_transform + else: + from sgl_kernel import hadamard_transform hidden_size = x.size(-1) assert ( @@ -145,7 +159,7 @@ class Indexer(MultiPlatformOp): else: self.cp_size = None self.cp_rank = None - if is_cuda(): + if _is_cuda: self.sm_count = deep_gemm.get_num_sms() self.half_device_sm_count = ceil_align(self.sm_count // 2, 8) pp_size = get_global_server_args().pp_size @@ -205,13 +219,13 @@ class Indexer(MultiPlatformOp): else: yield - @torch.compile(dynamic=True) + @torch.compile(dynamic=True) if not _is_hip else lambda f: f def _project_and_scale_head_gates(self, x: torch.Tensor): weights, _ = self.weights_proj(x.float()) weights = weights * self.n_heads**-0.5 return weights - @torch.compile(dynamic=True) + @torch.compile(dynamic=True) if not _is_hip else lambda f: f def _get_logits_head_gate(self, x: torch.Tensor, q_scale: torch.Tensor): weights, _ = self.weights_proj(x.float()) weights = weights * self.n_heads**-0.5 @@ -323,10 +337,13 @@ class Indexer(MultiPlatformOp): page_size = forward_batch.token_to_kv_pool.page_size # NOTE(dark): blocksize = 64 is hardcoded in deep_gemm - assert page_size == 64, "only support page size 64" - - # NOTE(dark): this support extend/decode/decode+graph - block_tables = metadata.get_page_table_64() + if _is_hip: + assert page_size == 1, "only support page size 1" + block_tables = metadata.get_page_table_1() + else: + assert page_size == 64, "only support page size 64" + # NOTE(dark): this support extend/decode/decode+graph + block_tables = metadata.get_page_table_64() max_seq_len = block_tables.shape[1] * page_size kv_cache_fp8 = forward_batch.token_to_kv_pool.get_index_k_with_scale_buffer( @@ -344,32 +361,64 @@ class Indexer(MultiPlatformOp): # Reuse pre-computed schedule metadata if available (from init_forward_metadata), # otherwise fall back to computing it here. schedule_metadata = getattr(metadata, "paged_mqa_schedule_metadata", None) - if schedule_metadata is None: - schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata( - seqlens_32, blocksize, self.sm_count - ) + if _is_cuda: + if schedule_metadata is None: + schedule_metadata = deep_gemm.get_paged_mqa_logits_metadata( + seqlens_32, blocksize, self.sm_count + ) assert len(q_fp8.shape) == 3 q_fp8 = q_fp8.unsqueeze(1) # the next_n dim is 1 now assert len(kv_cache_fp8.shape) == 2 - block_kv = 64 + block_kv = 1 if _is_hip else 64 num_heads_kv = 1 head_dim_with_sf = 132 - kv_cache_fp8 = kv_cache_fp8.view( - kv_cache_fp8.shape[0], block_kv, num_heads_kv, head_dim_with_sf - ) + if _is_hip: + kv_cache_fp8 = kv_cache_fp8.view( + -1, block_kv, num_heads_kv, head_dim_with_sf + ) + else: + kv_cache_fp8 = kv_cache_fp8.view( + kv_cache_fp8.shape[0], block_kv, num_heads_kv, head_dim_with_sf + ) assert len(weights.shape) == 3 weights = weights.squeeze(2) - logits = deep_gemm.fp8_paged_mqa_logits( - q_fp8, - kv_cache_fp8, - weights, - seqlens_32, - block_tables, - schedule_metadata, - max_seq_len, - clean_logits=False, - ) + + if _is_hip: + from aiter.ops.triton.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits + + batch_size, next_n, heads, _ = q_fp8.shape + logits = torch.full( + (batch_size * next_n, max_seq_len), + float("-inf"), + device=q_fp8.device, + dtype=torch.float32, + ) + deepgemm_fp8_paged_mqa_logits( + q_fp8, + kv_cache_fp8, + weights, + logits, + seqlens_32, + block_tables, + max_seq_len, + Preshuffle=False, + KVBlockSize=block_kv, + ChunkK=128, + TotalCuCount=256, + WavePerEU=5, + ) + else: + logits = deep_gemm.fp8_paged_mqa_logits( + q_fp8, + kv_cache_fp8, + weights, + seqlens_32, + block_tables, + schedule_metadata, + max_seq_len, + clean_logits=False, + ) # NOTE(dark): logits should be cleaned in topk_transform topk_result = metadata.topk_transform(logits, self.index_topk) @@ -408,13 +457,20 @@ class Indexer(MultiPlatformOp): assert forward_batch.forward_mode.is_extend_without_speculative() page_size = forward_batch.token_to_kv_pool.page_size - assert page_size == 64, "only support page size 64" + if _is_hip: + assert page_size == 1, "only support page size 1" + else: + assert page_size == 64, "only support page size 64" + assert len(weights.shape) == 3 weights = weights.squeeze(-1) k_fp8_list = [] k_scale_list = [] - block_tables = metadata.get_page_table_64() + if _is_hip: + block_tables = metadata.get_page_table_1() + else: + block_tables = metadata.get_page_table_64() assert ( forward_batch.seq_lens_cpu is not None @@ -459,14 +515,22 @@ class Indexer(MultiPlatformOp): if not need_chunk: assert q_fp8[:q_offset].shape[0] != 0 with self._with_real_sm_count(): - logits = deep_gemm.fp8_mqa_logits( - q_fp8[:q_offset], - kv_fp8, - weights[:q_offset], - ks, - ke, - clean_logits=False, - ) + if _is_hip: + from aiter.ops.triton.fp8_mqa_logits import fp8_mqa_logits + + kv, scale = kv_fp8 + logits = fp8_mqa_logits( + q_fp8[:q_offset], kv, scale, weights[:q_offset], ks, ke + ) + else: + logits = deep_gemm.fp8_mqa_logits( + q_fp8[:q_offset], + kv_fp8, + weights[:q_offset], + ks, + ke, + clean_logits=False, + ) assert logits.shape[0] == len(seq_lens_expanded) assert logits.shape[1] == k_offset @@ -496,14 +560,27 @@ class Indexer(MultiPlatformOp): end = min(start + max_rows, q_offset) with self._with_real_sm_count(): - logits_chunk = deep_gemm.fp8_mqa_logits( - q_fp8[start:end], - kv_fp8, - weights[start:end], - ks[start:end], - ke[start:end], - clean_logits=False, - ) + if _is_hip: + from aiter.ops.triton.fp8_mqa_logits import fp8_mqa_logits + + kv, scale = kv_fp8 + logits = fp8_mqa_logits( + q_fp8[start:end], + kv_fp8, + scale, + weights[start:end], + ks[start:end], + ke[start:end], + ) + else: + logits_chunk = deep_gemm.fp8_mqa_logits( + q_fp8[start:end], + kv_fp8, + weights[start:end], + ks[start:end], + ke[start:end], + clean_logits=False, + ) lengths_chunk = seq_lens_expanded[start:end] @@ -548,6 +625,7 @@ class Indexer(MultiPlatformOp): return_indices: bool = True, ) -> Optional[torch.Tensor]: assert forward_batch.forward_mode.is_extend_without_speculative() + x_meta = x[0] if isinstance(x, tuple) else x # Fast path: only compute and store k cache, skip all q and weights ops key = self._get_k_bf16(x, positions, enable_dual_stream) @@ -573,7 +651,7 @@ class Indexer(MultiPlatformOp): seq_lens_expanded.shape[0], self.index_topk, dtype=torch.float32, - device=x.device, + device=x_meta.device, ) return metadata.topk_transform(dummy_logits, self.index_topk) @@ -734,7 +812,7 @@ class Indexer(MultiPlatformOp): topk: int, layer_id: int, ) -> Optional[torch.Tensor]: - if not is_npu(): + if not _is_npu: from sglang.srt.layers.attention.nsa.tilelang_kernel import fp8_index page_size = forward_batch.token_to_kv_pool.page_size @@ -818,14 +896,18 @@ class Indexer(MultiPlatformOp): layer_id: int, return_indices: bool = True, ) -> Optional[torch.Tensor]: - if is_hip(): + if _is_hip: from sglang.srt.layers.attention.nsa.tilelang_kernel import act_quant - elif not is_npu(): + elif not _is_npu: from sglang.srt.layers.attention.nsa.triton_kernel import act_quant if TYPE_CHECKING: assert isinstance(forward_batch.token_to_kv_pool, NSATokenToKVPool) + # When upstream uses fused FP8 RMSNorm+quant, activations may be passed as + # a tuple like (x_fp8, x_scale[, y]). Use `x_meta` for shape/device queries. + x_meta = x[0] if isinstance(x, tuple) else x + metadata = forward_batch.attn_backend.get_indexer_metadata( layer_id, forward_batch ) @@ -891,7 +973,38 @@ class Indexer(MultiPlatformOp): q_fp8, q_scale = act_quant(query, self.block_size, self.scale_fmt) k_fp8, k_scale = act_quant(key, self.block_size, self.scale_fmt) - weights = self._get_logits_head_gate(x, q_scale) + # `_get_logits_head_gate` expects a Tensor. For tuple activations, dequantize + # to a float tensor here (callsite), keeping `_get_logits_head_gate` backend-agnostic. + if isinstance(x, tuple): + assert len(x) in ( + 2, + 3, + ), "For tuple input, only (x, x_s) or (x, x_s, y) formats are accepted" + x_q, x_s = x[0], x[1] + if ( + x_s is not None + and x_q.dim() == 2 + and x_s.dim() == 2 + and x_q.shape[0] == x_s.shape[0] + ): + m, n = x_q.shape + ng = x_s.shape[1] + if ng > 0 and n % ng == 0: + group = n // ng + x_for_gate = ( + x_q.to(torch.float32) + .view(m, ng, group) + .mul_(x_s.to(torch.float32).unsqueeze(-1)) + .view(m, n) + ) + else: + x_for_gate = x_q.to(torch.float32) + else: + x_for_gate = x_q.to(torch.float32) + else: + x_for_gate = x + + weights = self._get_logits_head_gate(x_for_gate, q_scale) # k_fp8: (seq_len, head_dim) fp8_e4m3fn # k_buffer: (num_total_tokens + page_size, head_dim) fp8_e4m3fn @@ -906,7 +1019,7 @@ class Indexer(MultiPlatformOp): index_k_scale=k_scale, ) - if is_cuda(): + if _is_cuda or _is_hip: assert forward_batch.seq_lens_cpu is not None if len(forward_batch.seq_lens_cpu) == 0: # this seems b/c max-pad, no worries? @@ -915,7 +1028,10 @@ class Indexer(MultiPlatformOp): # "HACK: seq_lens empty but x not empty, hackily return all-invalid topk_result" # ) return torch.full( - (x.shape[0], self.index_topk), -1, dtype=torch.int, device="cuda" + (x_meta.shape[0], self.index_topk), + -1, + dtype=torch.int, + device=x_meta.device, ) if ( diff --git a/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py b/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py index 05266ee72..c32d74fb9 100644 --- a/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py +++ b/python/sglang/srt/layers/attention/nsa/tilelang_kernel.py @@ -147,6 +147,8 @@ def fp8_index_kernel(h: int, d: int, clear_accum=True): T.copy(k_s[i_b, i1_n * blk_n1 + i2_n * blk_n2], k_s_frag) logits = T.alloc_fragment((blk_n2, h), FP32) + if not clear_accum: + T.fill(logits, 0) T.gemm( k_smem, q_smem, diff --git a/python/sglang/srt/layers/attention/nsa_backend.py b/python/sglang/srt/layers/attention/nsa_backend.py index b464f3f8d..eb7d8cd0a 100644 --- a/python/sglang/srt/layers/attention/nsa_backend.py +++ b/python/sglang/srt/layers/attention/nsa_backend.py @@ -174,6 +174,9 @@ class NSAIndexerMetadata(BaseIndexerMetadata): def get_page_table_64(self) -> torch.Tensor: return self.attn_metadata.real_page_table + def get_page_table_1(self) -> torch.Tensor: + return self.attn_metadata.page_table_1 + def get_seqlens_expanded(self) -> torch.Tensor: return self.attn_metadata.nsa_seqlens_expanded diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index 8549fbe69..378421fc6 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -52,7 +52,7 @@ from sglang.srt.mem_cache.utils import ( set_mla_kv_buffer_triton, set_mla_kv_scale_buffer_triton, ) -from sglang.srt.utils import is_cuda, is_npu, next_power_of_2 +from sglang.srt.utils import is_cuda, is_hip, is_npu, next_power_of_2 from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter @@ -68,6 +68,7 @@ logger = logging.getLogger(__name__) GB = 1024 * 1024 * 1024 _is_cuda = is_cuda() _is_npu = is_npu() +_is_hip = is_hip() def get_tensor_size_bytes(t: Union[torch.Tensor, List[torch.Tensor]]): @@ -1724,7 +1725,10 @@ class NSATokenToKVPool(MLATokenToKVPool): # num head == 1 and head dim == 128 for index_k in NSA assert index_head_dim == 128 - assert self.page_size == 64 + if _is_hip: + assert self.page_size == 1 + else: + assert self.page_size == 64 with ( torch.cuda.use_mem_pool(self.custom_mem_pool) if self.custom_mem_pool diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 75746bfe7..3aa447f78 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -1523,10 +1523,34 @@ class DeepseekV2AttentionMLA(nn.Module): # NSA Indexer: cache quantized keys, auto-skip topk for sequences <= nsa_index_topk if self.use_nsa: - q_lora = self.q_a_layernorm(q) - q = self.q_b_proj(q_lora)[0].view( - -1, self.num_local_heads, self.qk_head_dim - ) + # NSA requires unquantized q_lora for the indexer. When q_b_proj is FP8 + # on gfx95, we can still use fused RMSNorm+FP8 quant, but MUST request + # the unquantized output for q_lora; otherwise q_lora becomes the (fp8,scale) + # tuple. + if ( + _use_aiter_gfx95 + and self.q_b_proj.weight.dtype == torch.float8_e4m3fn + ): + q_quanted, q_lora, _, _ = fused_rms_fp8_group_quant( + q, + self.q_a_layernorm.weight, + self.q_a_layernorm.variance_epsilon, + None, + None, + None, + group_size=128, + dtype_quant=torch.float8_e4m3fn, + res1=None, + output_unquantized_inp1=True, + ) + q = self.q_b_proj(q_quanted)[0].view( + -1, self.num_local_heads, self.qk_head_dim + ) + else: + q_lora = self.q_a_layernorm(q) + q = self.q_b_proj(q_lora)[0].view( + -1, self.num_local_heads, self.qk_head_dim + ) _ = self.indexer( x=hidden_states, q_lora=q_lora, @@ -1703,23 +1727,38 @@ class DeepseekV2AttentionMLA(nn.Module): self.kv_a_layernorm.variance_epsilon, ) else: + q_lora = None if ( _use_aiter_gfx95 and self.q_b_proj.weight.dtype == torch.float8_e4m3fn ): - - q, _, k_nope, _ = fused_rms_fp8_group_quant( - q, - self.q_a_layernorm.weight, - self.q_a_layernorm.variance_epsilon, - k_nope, - self.kv_a_layernorm.weight, - self.kv_a_layernorm.variance_epsilon, - group_size=128, - dtype_quant=torch.float8_e4m3fn, - res1=None, - output_unquantized_inp1=False, - ) + if self.use_nsa: + q_quanted, q_lora, k_nope, _ = fused_rms_fp8_group_quant( + q, + self.q_a_layernorm.weight, + self.q_a_layernorm.variance_epsilon, + k_nope, + self.kv_a_layernorm.weight, + self.kv_a_layernorm.variance_epsilon, + group_size=128, + dtype_quant=torch.float8_e4m3fn, + res1=None, + output_unquantized_inp1=True, + ) + q = q_quanted + else: + q, _, k_nope, _ = fused_rms_fp8_group_quant( + q, + self.q_a_layernorm.weight, + self.q_a_layernorm.variance_epsilon, + k_nope, + self.kv_a_layernorm.weight, + self.kv_a_layernorm.variance_epsilon, + group_size=128, + dtype_quant=torch.float8_e4m3fn, + res1=None, + output_unquantized_inp1=False, + ) else: q = self.q_a_layernorm(q) @@ -1727,7 +1766,8 @@ class DeepseekV2AttentionMLA(nn.Module): # q_lora needed by indexer if self.use_nsa: - q_lora = q + if q_lora is None: + q_lora = q # overlap q_b_proj and indexer during decode if ( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 6802f2011..fffed38fa 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1090,7 +1090,7 @@ class ServerArgs: self.attention_backend = "nsa" logger.info("Use nsa attention backend for DeepSeek with DSA.") - if not is_npu(): # CUDA GPU + if not is_npu(): # CUDA or ROCm GPU if self.enable_nsa_prefill_context_parallel: logger.warning( f"Context parallel feature is still under experiment. It has only been verified on Hopper platform." @@ -1126,8 +1126,15 @@ class ServerArgs: f"attn_tp_size={self.tp_size}, attention weights will be sharded across {self.tp_size} ranks." ) - self.page_size = 64 - logger.warning("Setting page size to 64 for DeepSeek DSA.") + if is_hip(): + self.page_size = 1 + logger.warning( + "Setting page size to 1 for DeepSeek DSA on ROCm." + ) + else: + # For CUDA GPU + self.page_size = 64 + logger.warning("Setting page size to 64 for DeepSeek DSA.") # For Hopper, we support both bf16 and fp8 kv cache; for Blackwell, we support fp8 only currently import torch