From 5e0fcf147fc09b3cf6590f9681542d57d08e98dd Mon Sep 17 00:00:00 2001 From: apinge Date: Tue, 28 Apr 2026 11:25:28 +0000 Subject: [PATCH 1/3] add pagged attention nhd for aiter_backend --- .../srt/layers/attention/aiter_backend.py | 62 ++++++++++++++----- 1 file changed, 48 insertions(+), 14 deletions(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index a52fbfab9961..ee6290df7fa6 100644 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -5,6 +5,7 @@ """ import logging +import os from dataclasses import dataclass from enum import Enum, auto from typing import TYPE_CHECKING, Optional @@ -37,6 +38,7 @@ mla_prefill_ps_asm_fwd, mla_reduce_v1, paged_attention_ragged, + paged_attention_ragged_nhd, ) from aiter.mla import mla_decode_fwd, mla_prefill_fwd except ImportError: @@ -94,10 +96,8 @@ class ForwardMetadata: max_extend_len: Optional[int] = None -global_workspace_buffer = None - -_AITER_PARTITION_SIZE_ROCM = 256 +_AITER_PARTITION_SIZE_ROCM = 256 # 256 512 1024 for paged_attention_ragged_nhd class AiterAttnBackend(AttentionBackend): def __init__( @@ -187,17 +187,28 @@ def __init__( self.max_num_partitions = ( self.max_context_len + _AITER_PARTITION_SIZE_ROCM - 1 ) // _AITER_PARTITION_SIZE_ROCM - nbyes_per_qo_elem = torch.finfo(torch.float32).bits // 8 if not self.use_mla: - self.workspace_buffer = torch.empty( - (max_bs * self.num_head * self.max_num_partitions * self.head_dim) - * nbyes_per_qo_elem - + 2 * (max_bs * self.num_head * self.max_num_partitions) * 4, - dtype=torch.uint8, - device=self.device, - ) + # self.workspace_buffer = torch.empty( + # (max_bs * self.num_head * self.max_num_partitions * self.head_dim) + # * nbyes_per_qo_elem + # + 2 * (max_bs * self.num_head * self.max_num_partitions) * 4, + # dtype=torch.uint8, + # device=self.device, + # ).contiguous() + # aiter pa_ragged_nhd: three tensors (no single uint8 blob + host pointer math). + _nhp = max_bs * self.num_head * self.max_num_partitions + self.pa_nhd_exp_sums = torch.empty( + (_nhp,), dtype=torch.float32, device=self.device + ).contiguous() + self.pa_nhd_max_logits = torch.empty( + (_nhp,), dtype=torch.float32, device=self.device + ).contiguous() + self.pa_nhd_tmp_out = torch.empty( + (_nhp * self.head_dim,), dtype=torch.bfloat16, device=self.device + ).contiguous() + self.scale = float(1.0 / (self.head_dim**0.5)) self.k_scale = self.v_scale = torch.tensor([1.0], dtype=torch.float32).to( @@ -1699,6 +1710,7 @@ def forward_decode( ) else: o = torch.empty_like(q, dtype=self.input_dtype) + #o_nhd = torch.empty_like(q, dtype=self.input_dtype) # use for debug if save_kv_cache: forward_batch.token_to_kv_pool.set_kv_buffer( @@ -1770,10 +1782,11 @@ def forward_decode( k_cache = k_cache.to(dtype) v_cache = v_cache.to(dtype) - - paged_attention_ragged( + paged_attention_ragged_nhd( o.view(-1, layer.tp_q_head_num, layer.qk_head_dim), - self.workspace_buffer, + self.pa_nhd_exp_sums, + self.pa_nhd_max_logits, + self.pa_nhd_tmp_out, q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), k_cache.view(-1, 1, layer.tp_k_head_num, layer.qk_head_dim), v_cache.view(-1, 1, layer.tp_v_head_num, layer.v_head_dim), @@ -1792,6 +1805,27 @@ def forward_decode( None, _AITER_PARTITION_SIZE_ROCM, ) + # paged_attention_ragged( + # o.view(-1, layer.tp_q_head_num, layer.qk_head_dim), + # self.workspace_buffer, + # q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), + # k_cache.view(-1, 1, layer.tp_k_head_num, layer.qk_head_dim), + # v_cache.view(-1, 1, layer.tp_v_head_num, layer.v_head_dim), + # self.scale, + # self.forward_metadata.kv_indptr, + # self.forward_metadata.kv_indices, + # self.kv_last_page_len, + # 1, + # self.max_num_partitions, + # None, + # "auto", + # "NHD", + # self.logits_soft_cap, + # self.k_scale, + # self.v_scale, + # None, + # _AITER_PARTITION_SIZE_ROCM, + # ) return o From 3a38ce994382e0a5bcc11174e75dd8079f0e5355 Mon Sep 17 00:00:00 2001 From: apinge Date: Tue, 28 Apr 2026 11:29:00 +0000 Subject: [PATCH 2/3] revert unused modification --- python/sglang/srt/layers/attention/aiter_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index ee6290df7fa6..268b2c18a670 100644 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -5,7 +5,6 @@ """ import logging -import os from dataclasses import dataclass from enum import Enum, auto from typing import TYPE_CHECKING, Optional @@ -96,6 +95,7 @@ class ForwardMetadata: max_extend_len: Optional[int] = None +global_workspace_buffer = None _AITER_PARTITION_SIZE_ROCM = 256 # 256 512 1024 for paged_attention_ragged_nhd From f742fe9dee0e59b7787ad3a56f22408e9ece4e79 Mon Sep 17 00:00:00 2001 From: apinge Date: Tue, 28 Apr 2026 11:30:07 +0000 Subject: [PATCH 3/3] remove comment --- python/sglang/srt/layers/attention/aiter_backend.py | 1 - 1 file changed, 1 deletion(-) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 268b2c18a670..dbcfd96730d0 100644 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -1710,7 +1710,6 @@ def forward_decode( ) else: o = torch.empty_like(q, dtype=self.input_dtype) - #o_nhd = torch.empty_like(q, dtype=self.input_dtype) # use for debug if save_kv_cache: forward_batch.token_to_kv_pool.set_kv_buffer(