Fix ROCm fused KV and KDA paths (#31688)
Co-authored-by: wangwenchen0407 <wangwenchen@meta.com>
This commit is contained in:
co-authored by
wangwenchen0407
parent
688a6d23f1
commit
555267ed05
@@ -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_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_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))
|
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_Ai00 = b_Akk_d0
|
||||||
b_Ai11 = b_Akk_d1
|
b_Ai11 = b_Akk_d1
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ import triton
|
|||||||
import triton.language as tl
|
import triton.language as tl
|
||||||
|
|
||||||
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
|
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.environ import envs
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.layers.utils.cp_utils import is_prefill_context_parallel_enabled
|
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,
|
layer: RadixAttention,
|
||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
):
|
):
|
||||||
from sglang.jit_kernel.rope import FusedSetKVBufferArg
|
|
||||||
|
|
||||||
layer_id = layer.layer_id
|
layer_id = layer.layer_id
|
||||||
token_to_kv_pool = get_token_to_kv_pool()
|
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)
|
v_buffer = token_to_kv_pool.get_value_buffer(layer_id)
|
||||||
|
|
||||||
if not _is_hip:
|
if not _is_hip:
|
||||||
|
# CUDA path.
|
||||||
assert layer.k_scale is None and layer.v_scale is None, "scale not supported"
|
assert layer.k_scale is None and layer.v_scale is None, "scale not supported"
|
||||||
return FusedSetKVBufferArg(
|
return FusedSetKVBufferArg(
|
||||||
value=value,
|
value=value,
|
||||||
@@ -323,10 +323,20 @@ def create_fused_set_kv_buffer_arg(
|
|||||||
cache_loc=forward_batch.out_cache_loc,
|
cache_loc=forward_batch.out_cache_loc,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
|
# ROCm path.
|
||||||
page_size = token_to_kv_pool.page_size
|
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 = (
|
slot_mapping_swa = (
|
||||||
token_to_kv_pool.full_to_swa_index_mapping.long()
|
full_to_swa.long()
|
||||||
if layer.sliding_window_size > 0
|
if layer.sliding_window_size > 0 and full_to_swa is not None
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
# SHUFFLE 5D pools (k_buffer.ndim == 5) consumed natively by
|
# SHUFFLE 5D pools (k_buffer.ndim == 5) consumed natively by
|
||||||
|
|||||||
Reference in New Issue
Block a user