Fix ROCm fused KV and KDA paths (#31688)

Co-authored-by: wangwenchen0407 <wangwenchen@meta.com>
This commit is contained in:
Lianmin Zheng
2026-07-19 15:02:35 -07:00
committed by GitHub
co-authored by wangwenchen0407
parent 688a6d23f1
commit 555267ed05
2 changed files with 16 additions and 4 deletions
@@ -429,6 +429,8 @@ def chunk_kda_fwd_kernel_inter_solve_fused(
tl.store(p_Akkd11, b_Akk_d1.to(Akkd.dtype.element_ty), boundary_check=(0, 1))
tl.store(p_Akkd22, b_Akk_d2.to(Akkd.dtype.element_ty), boundary_check=(0, 1))
tl.store(p_Akkd33, b_Akk_d3.to(Akkd.dtype.element_ty), boundary_check=(0, 1))
# Forward substitution reloads these global tiles across warps below.
tl.debug_barrier()
b_Ai00 = b_Akk_d0
b_Ai11 = b_Akk_d1
+14 -4
View File
@@ -25,6 +25,7 @@ import triton
import triton.language as tl
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
from sglang.jit_kernel.rope import FusedSetKVBufferArg
from sglang.srt.environ import envs
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled
@@ -306,8 +307,6 @@ def create_fused_set_kv_buffer_arg(
layer: RadixAttention,
forward_batch: ForwardBatch,
):
from sglang.jit_kernel.rope import FusedSetKVBufferArg
layer_id = layer.layer_id
token_to_kv_pool = get_token_to_kv_pool()
@@ -315,6 +314,7 @@ def create_fused_set_kv_buffer_arg(
v_buffer = token_to_kv_pool.get_value_buffer(layer_id)
if not _is_hip:
# CUDA path.
assert layer.k_scale is None and layer.v_scale is None, "scale not supported"
return FusedSetKVBufferArg(
value=value,
@@ -323,10 +323,20 @@ def create_fused_set_kv_buffer_arg(
cache_loc=forward_batch.out_cache_loc,
)
else:
# ROCm path.
page_size = token_to_kv_pool.page_size
# A non-hybrid pool has no full->SWA remap: SWA and full layers
# share one slot space indexed directly by out_cache_loc (as --disable-hybrid-swa-memory
# gives). Leaving swa_slot_mapping=None makes the fused store write at
# out_cache_loc, matching the CUDA path (which never fuses the hybrid pool).
full_to_swa = (
token_to_kv_pool.full_to_swa_index_mapping
if isinstance(token_to_kv_pool, SWAKVPool)
else None
)
slot_mapping_swa = (
token_to_kv_pool.full_to_swa_index_mapping.long()
if layer.sliding_window_size > 0
full_to_swa.long()
if layer.sliding_window_size > 0 and full_to_swa is not None
else None
)
# SHUFFLE 5D pools (k_buffer.ndim == 5) consumed natively by