From a683c6e6d3e40e35c61e35cc22364a7c45950deb Mon Sep 17 00:00:00 2001 From: fsx950223 Date: Fri, 24 Jul 2026 08:25:47 +0000 Subject: [PATCH 1/7] refactor(pa): Migrate metadata kernels to tile APIs Signed-off-by: fsx950223 Co-authored-by: Cursor --- kernels/attention/pa_decode_fp8.py | 47 +- kernels/attention/pa_metadata.py | 1654 ++++++++++------------------ tests/kernels/test_pa.py | 35 + 3 files changed, 598 insertions(+), 1138 deletions(-) diff --git a/kernels/attention/pa_decode_fp8.py b/kernels/attention/pa_decode_fp8.py index b1867529e..5bbb46c98 100644 --- a/kernels/attention/pa_decode_fp8.py +++ b/kernels/attention/pa_decode_fp8.py @@ -30,16 +30,6 @@ _PA_METADATA_GRID_OVERSUB = 3 MFMA_N = 16 -_PACKED_FP8_QUERY_DTYPES = tuple( - dtype - for dtype in ( - torch.uint8, - getattr(torch, "float8_e4m3fnuz", None), - getattr(torch, "float8_e4m3fn", None), - ) - if dtype is not None -) - # ===================================================================== # Launch API — Persistent Scheduling mode @@ -100,10 +90,7 @@ def get_pa_metadata( partition_indptr = torch.zeros(batch_size + 1, dtype=torch.int32, device=dev) partition_indptr[1:] = torch.cumsum(_parts_per_batch, dim=0).to(torch.int32) - block_size = key_cache.shape[-2] if len(key_cache.shape) == 5 else key_cache.shape[-2] - ( - (work_meta_data_size, work_meta_data_type), (work_indptr_size, work_indptr_type), (work_info_set_size, work_info_set_type), (reduce_indptr_size, reduce_indptr_type), @@ -111,7 +98,6 @@ def get_pa_metadata( (reduce_partial_map_size, reduce_partial_map_type), ) = get_pa_metadata_info_v1(batch_size, num_kv_heads, num_cu=num_sm) - work_metadata_ptrs = torch.empty(work_meta_data_size, dtype=work_meta_data_type, device=dev) work_indptr = torch.empty(work_indptr_size, dtype=work_indptr_type, device=dev) work_info = torch.empty(work_info_set_size, dtype=work_info_set_type, device=dev) reduce_indptr = torch.empty(reduce_indptr_size, dtype=reduce_indptr_type, device=dev) @@ -120,24 +106,18 @@ def get_pa_metadata( get_pa_metadata_v1( seqlens_qo_indptr, - kv_indptr, context_lengths, - num_query_heads // num_kv_heads, - num_kv_heads, - True, - work_metadata_ptrs, work_indptr, work_info, reduce_indptr, reduce_final_map, reduce_partial_map, + query_group_size=num_query_heads // num_kv_heads, + num_kv_heads=num_kv_heads, kv_granularity=partition_size, - block_size=block_size, - max_seqlen_qo=query_length, - uni_seqlen_qo=query_length, - fast_mode=True, - max_split_per_batch=-1, + query_length=query_length, num_cu=num_sm, + stream=torch.cuda.current_stream(dev), ) # The FlyDSL get_pa_metadata_v1 produces the reduce_* maps natively @@ -216,15 +196,11 @@ def _prepare_scale_tensor( def _get_query_input_dtype(query: torch.Tensor) -> str: - if query.dtype in _PACKED_FP8_QUERY_DTYPES: - return "packed_fp8" if query.dtype == torch.bfloat16: return "bf16" if query.dtype == torch.float16: return "f16" - raise ValueError( - f"Unsupported query dtype for pa_decode_ps_launch: {query.dtype}. " "Expected packed FP8/uint8, bf16, or f16." - ) + raise ValueError(f"Unsupported query dtype for pa_decode_ps_launch: {query.dtype}. Expected bf16 or f16.") def _get_output_dtype_str(output: torch.Tensor) -> str: @@ -269,9 +245,8 @@ def get_recommended_splits( return max(4, min(n, 8)) -# Small block_size (16/64) is routed through the load-balanced worklist -# (metadata) path: `compile_pa_decode_metadata` gathers 256//block_size physical -# pages per 256-token partition, for both per-tensor and per-token KV quant. +# Small block sizes use the standalone tile kernel; the metadata decode path +# below is reserved for 1024-token physical pages. _PA_DECODE_PS_SMALL_BLOCK_SIZES = (16, 64) @@ -322,12 +297,6 @@ def pa_decode_ps_launch( device=dev, is_graph_capturing=is_graph_capturing, ) - if query_input_dtype == "packed_fp8": - raise ValueError( - "`pa_decode_ps_launch` no longer accepts host query_scale and only supports " - "bf16/f16 query inputs with kernel-internal query scale computation." - ) - # Detect per-token vs per-tensor quantization from scale tensor # dimensionality: a >1-D scale tensor carries one scale per (block, head, # token), which enables the per-token K/V path in the metadata kernel. @@ -601,7 +570,7 @@ def pa_decode_ps_launch( max_seqlen_q=query_length, final_output=output, num_query_heads=num_query_heads, - head_size=int(query.shape[-1]), + head_dim=int(query.shape[-1]), stream=s, ) diff --git a/kernels/attention/pa_metadata.py b/kernels/attention/pa_metadata.py index 7a4b76162..32bfec696 100644 --- a/kernels/attention/pa_metadata.py +++ b/kernels/attention/pa_metadata.py @@ -1,37 +1,15 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2025 FlyDSL Project Contributors -"""FlyDSL implementation of aiter's ``get_pa_metadata_v1`` worklist scheduler. - -This replaces the ``aiter.ops.attention.get_pa_metadata_v1`` C++/CUDA dependency -(``module_pa_metadata.so``) with a FlyDSL device kernel, so paged-attention -decode (``kernels/pa_decode_fp8.py``) can build its CU worklist without aiter. - -Scope — PA-decode-specialized port of -``aiter/csrc/kernels/mla/metadata/v1_2_pa_device.cuh::kn_get_pa_metadata_v1_2``. -The following invariants always hold for the PA decode use and are baked in: - -* ``kQoSplits == False`` — ``packed_qo_len = query_length * gqa`` is small - (<= ~32) so it never exceeds ``kPackedQoLenPerWg=128`` ⇒ ``num_qo_tiles = 1``, - ``qo_tile_size = query_length``. -* uniform qo length across batches (``uni_seqlen_qo = query_length``). -* causal, non-sparse (``topk = -1``), ``qk_batch_ratio = 1``, - ``num_splits = num_cu`` (``max_split_per_batch = -1``). - -All six outputs are produced as a faithful drop-in for the C++ kernel: -``work_metadata_ptrs``, ``work_indptr``, ``work_info`` (8 fields), -``reduce_indptr``, ``reduce_final_map`` and ``reduce_partial_map`` — each -verified element-for-element against aiter. The caller consumes them directly -(no post-hoc expansion). - -work_info layout (8 x int32 per work), matching ``PaWorkInfo``: +"""Paged-attention worklist scheduler and persistent block-1024 decode kernel. + +The scheduler is specialized for uniform query lengths, causal non-sparse +attention, one query tile, and ``num_splits = num_cu``. It produces +``work_indptr``, ``work_info``, and the three reduction maps without aiter. + +``work_info`` layout (8 x int32 per work): [0] batch_idx [1] partial_qo_loc(-1 if no split) [2] qo_start [3] qo_end [4] kv_start [5] kv_end [6] kv_offset(=0) [7] q_head_range = (qhead_end << 16) | (qhead_start & 0xFFFF) - -The kernel is launched single-thread (grid=block=(1,1,1)); the scheduler is a -serial algorithm (warp reductions / lane-parallel fills in the original collapse -to serial loops). It runs once per shape and the result is cached, so single- -thread is the correct, simplest model. """ import functools @@ -41,18 +19,16 @@ import flydsl.compiler as flyc import flydsl.expr as fx -from flydsl._mlir.dialects import vector -from flydsl.expr import arith, as_ir_value, const_expr, gpu, range_constexpr, rocdl +from flydsl.expr import arith, const_expr, gpu, range_constexpr, rocdl +from flydsl.expr import math as fmath from flydsl.expr.typing import Int32, T from flydsl.runtime.device import get_rocm_arch -from kernels.attention.pa_common import _compute_block_base_dw_i64, _prefetch_q_chunks -from kernels.common import buffer_ops, dpp_utils +from kernels.attention.pa_common import _compute_block_base_dw_i64 +from kernels.common import dpp_utils from kernels.common.kernels_common import get_warp_size from kernels.common.tensor_shim import _run_compiled from kernels.common.utils import ( exp2_f32_fast, - global_load_i64x2, - global_ptr_from_addr, is_pow2, rcp_f32, udiv_const, @@ -61,14 +37,40 @@ ) _WORK_INFO_FIELDS = 8 +_FLAT_BUFFER_ELEMENTS = 1 << 30 +_PA_OUTPUT_DTYPES = { + "f32": fx.Float32, + "f16": fx.Float16, + "bf16": fx.BFloat16, +} + + +def _global_pointer_from_addr(addr, dtype, *, alignment: int): + ptr_type = fx.PointerType.get( + elem_ty=dtype.ir_type, + address_space=fx.AddressSpace.Global, + alignment=alignment, + ) + return fx.inttoptr(ptr_type, addr) + + +def _copy_load(source, offset, copy_atom, register): + fx.copy(copy_atom, fx.slice(source, (None, fx.Int32(offset))), register) + return fx.memref_load_vec(register) + + +def _load_global_16b(global_ptr, byte_offset, copy_atom, reg): + src = fx.make_view(global_ptr + byte_offset, fx.make_layout(16, 1)) + fx.copy(copy_atom, src, reg) + return fx.memref_load_vec(reg).bitcast(fx.Int64) def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = None): - """Buffer sizes/dtypes, matching ``aiter.get_pa_metadata_info_v1``. + """Return buffer sizes and dtypes for the PA worklist tensors. Returns (shape, dtype) tuples for: - work_metadata_ptrs, work_indptr, work_info_set, - reduce_indptr, reduce_final_map, reduce_partial_map. + work_indptr, work_info_set, reduce_indptr, + reduce_final_map, reduce_partial_map. ``num_cu`` overrides the worklist bin count (default = device CU count); pass a multiple of the CU count to oversubscribe the persistent grid. @@ -83,7 +85,6 @@ def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = max_split_tiles = min(batch_size + cu_num - 1, (cu_num - 1) * 2) return ( - ((2,), torch.uint64), # work_metadata_ptrs ((cu_num + 1,), torch.int32), # work_indptr ((max_work, _WORK_INFO_FIELDS), torch.int32), # work_info_set ((tile_cnt + 1,), torch.int32), # reduce_indptr @@ -92,7 +93,6 @@ def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = ) -# ── PA decode geometry + helpers (shared math lives in kernels/utils.py) ── KV_BLOCK_SIZE = 1024 # physical page size (matches SP3 kBlockSize) KV_COMPUTE_BLOCK = 256 # tile size (matches SP3 kTileKV) @@ -140,6 +140,8 @@ def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = def _load_k_flat( k_global_ptr, + k_copy_atom, + k_reg, k_block_base_dw_i64, tile_token_offset_i32, k_tok_thread_base, @@ -156,7 +158,7 @@ def _load_k_flat( kbo_dw = kbo * c_tok_stride_dw for qkhe in range_constexpr(qkhe_loop): ka_dw = k_block_base_dw_i64 + fx.Int64(kbo_dw + k_he_off_dw[qkhe]) - k2 = global_load_i64x2(k_global_ptr, ka_dw * fx.Int64(4)) + k2 = _load_global_16b(k_global_ptr, ka_dw * fx.Int64(4), k_copy_atom, k_reg) k2_words = fx.Vector(k2) k_flat.append(k2_words[0]) k_flat.append(k2_words[1]) @@ -164,71 +166,18 @@ def _load_k_flat( return k_flat -def _build_pa_thread_invariants( +def _build_pa_k_thread_invariants( warp_id, lane16id, rowid, *, - trans_v, - per_token_kv, qkhe_loop: int = 2, - vhe_loop: int = 2, ): - c_tokens_per_warp = fx.Int32(TOKENS_PER_WARP) - c_mfma_n = fx.Int32(MFMA_N) - k_tok_thread_base = warp_id * c_tokens_per_warp + lane16id + k_tok_thread_base = warp_id * fx.Int32(TOKENS_PER_WARP) + lane16id c_tok_stride_dw = fx.Int32(FP8_ELEMS_16B // 4) c_he_stride_dw = fx.Int32(KV_BLOCK_SIZE * FP8_ELEMS_16B // 4) k_he_off_dw = [rowid * c_he_stride_dw + fx.Int32(qkhe * 4) * c_he_stride_dw for qkhe in range(qkhe_loop)] - - vhead_elems = [fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * c_mfma_n + lane16id for vhe in range(vhe_loop)] - v_tok_thread_off = [fx.Int32(vt * TOKENS_PER_WARP) + rowid * c_mfma_n for vt in range(VTLOOP)] - if const_expr(trans_v): - vhead_elem_dw = [vhead_elems[vhe] * fx.Int32(FP8_ELEMS_16B // 4) for vhe in range(vhe_loop)] - else: - vhead_elem_dw = [vhead_elems[vhe] * fx.Int32(KV_BLOCK_SIZE // 4) for vhe in range(vhe_loop)] - - kv_tok_thread_base = warp_id * c_tokens_per_warp + rowid * 4 - rowid_8x8 = rowid >> fx.Int32(1) - offset_in_slot = rowid & fx.Int32(1) - prob_row_i32 = PROB_ROW_STRIDE_BYTES // 4 - prob_row_i64 = PROB_ROW_STRIDE_BYTES // 8 - prob_wr_thread_base = ( - warp_id * fx.Int32(4 * MFMA_N * prob_row_i32) - + lane16id * fx.Int32(prob_row_i32) - + rowid_8x8 * fx.Int32(2) - + offset_in_slot - ) - pv_prob_read_base = rowid * fx.Int32(MFMA_N * prob_row_i64) + lane16id * fx.Int32(prob_row_i64) - - sm_lane_wave_base = lane16id * fx.Int32(NUM_WARPS) - sm_max_off = sm_lane_wave_base + warp_id - sm_sum_off = fx.Int32(NUM_WARPS * MFMA_N) + sm_lane_wave_base + warp_id - sm_rd_max_offs = [sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] - sm_rd_sum_offs = [fx.Int32(NUM_WARPS * MFMA_N) + sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] - - sm_vmax_wr_off = None - sm_vmax_rd_offs = None - if const_expr(per_token_kv): - sm_vmax_wr_off = fx.Int32(2 * NUM_WARPS * MFMA_N) + sm_lane_wave_base + warp_id - sm_vmax_rd_offs = [fx.Int32(2 * NUM_WARPS * MFMA_N) + sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] - - return ( - k_tok_thread_base, - c_tok_stride_dw, - k_he_off_dw, - v_tok_thread_off, - vhead_elem_dw, - kv_tok_thread_base, - prob_wr_thread_base, - pv_prob_read_base, - sm_max_off, - sm_sum_off, - sm_rd_max_offs, - sm_rd_sum_offs, - sm_vmax_wr_off, - sm_vmax_rd_offs, - ) + return k_tok_thread_base, c_tok_stride_dw, k_he_off_dw def _compute_mtp_group_state( @@ -279,18 +228,15 @@ def _finish_q_fragments( rowid, local_qhead_idx, *, - head_size: int, + head_dim: int, qkhe_loop: int, q_lanes_per_head: int, ): - c_head_dw = fx.Int32(head_size // 4) - lds_q_base = local_qhead_idx * c_head_dw + lane16id * 2 + lds_q_base = local_qhead_idx * fx.Int32(head_dim // 4) + lane16id * 2 abs_mask = fx.Vector.filled(4, 0x7FFFFFFF, fx.Int32) - c_zero_f = fx.Float32(0.0) - c_one_f = fx.Float32(1.0) q_f32_chunks = [] - local_max = c_zero_f + local_max = fx.Float32(0.0) for q_src in q_chunks: q_f32 = fx.Vector(q_src).to(fx.Float32) q_f32_chunks.append(q_f32) @@ -304,9 +250,9 @@ def _finish_q_fragments( local_max = fx.maxnumf(local_max, dpp_utils.dpp_xor_f32(local_max, sh)) query_scale_lane = fx.Float32( arith.select( - local_max > c_zero_f, + local_max > fx.Float32(0.0), local_max * fx.Float32(1.0 / FP8_MAX).ir_value(), - c_one_f, + fx.Float32(1.0), ) ) inv_query_scale = rcp_f32(query_scale_lane) @@ -337,7 +283,7 @@ def _finish_q_fragments( ].ir_value() for qkhe in range_constexpr(qkhe_loop): for qkr in range_constexpr(2): - lds_rd = lane16id * fx.Int32(head_size // 8) + fx.Int32(qkhe * 8) + rowid * fx.Int32(2) + fx.Int32(qkr) + lds_rd = lane16id * fx.Int32(head_dim // 8) + fx.Int32(qkhe * 8) + rowid * fx.Int32(2) + fx.Int32(qkr) q_v1 = fx.ptr_load( fx.recast_iter(fx.Int64, logits_base) + (lds_rd), result_type=fx.Vector.make_type(1, fx.Int64) ) @@ -346,7 +292,9 @@ def _finish_q_fragments( def _prefetch_mtp_group_query( - q_rsrc, + q_tiles, + q_copy_atom, + q_reg, batch_idx, kv_h, stride_q_seq, @@ -357,7 +305,6 @@ def _prefetch_mtp_group_query( mtp_group_idx, query_length, query_group_size, - query_load_is_bf16, q_lanes_per_head, ): qi_val, qhi_pos, qi_for_q, local_qhead_idx_for_q = _compute_mtp_group_state( @@ -367,18 +314,16 @@ def _prefetch_mtp_group_query( query_length=query_length, query_group_size=query_group_size, ) - q_row = batch_idx * arith.constant(query_length, type=T.i32) + qi_for_q - q_base = ( - q_row * stride_q_seq - + (kv_h * arith.constant(query_group_size, type=T.i32) + local_qhead_idx_for_q) * stride_q_head - ) - q_chunks = _prefetch_q_chunks( - q_rsrc, - q_base, - lane16id, - query_load_is_bf16=query_load_is_bf16, - q_lanes_per_head=q_lanes_per_head, - ) + q_row = batch_idx * fx.Int32(query_length) + qi_for_q + q_base = q_row * stride_q_seq + (kv_h * fx.Int32(query_group_size) + local_qhead_idx_for_q) * stride_q_head + q_load_lane = lane16id + if const_expr(q_lanes_per_head < MFMA_N): + q_load_lane = (lane16id < fx.Int32(q_lanes_per_head)).select(lane16id, fx.Int32(0)) + q_elem = q_base + q_load_lane * fx.Int32(Q_ELEMS_PER_LANE) + q_tile = q_elem // fx.Int32(4) + q_chunks = [] + for qwi in range_constexpr(Q_CHUNKS_PER_LANE): + q_chunks.append(_copy_load(q_tiles, q_tile + fx.Int32(qwi), q_copy_atom, q_reg)) return qi_val, qhi_pos, q_chunks @@ -390,7 +335,7 @@ def _finish_mtp_group_q_fragments( rowid, local_qhead_idx, *, - head_size: int, + head_dim: int, qkhe_loop: int, q_lanes_per_head: int, ): @@ -402,7 +347,7 @@ def _finish_mtp_group_q_fragments( lane16id, rowid, local_qhead_idx, - head_size=head_size, + head_dim=head_dim, qkhe_loop=qkhe_loop, q_lanes_per_head=q_lanes_per_head, ) @@ -413,64 +358,79 @@ def _normalize_pa_output(running_sum, outs, zero_f): one_f = fx.Float32(1.0).ir_value() safe_sum = arith.select(running_sum > zero_f, running_sum, one_f) inv_sum = rcp_f32(safe_sum) - inv_sum_vec = vector.broadcast(T.f32x4, inv_sum) - return [out * inv_sum_vec for out in outs] + return [out * fx.Float32(inv_sum) for out in outs] @flyc.jit def _make_pa_phase_helpers( *, trans_v, - per_token_q, per_token_kv, - needs_mask, - query_length, kv_h, v_global_ptr, - ks_rsrc, - vs_rsrc, + v_copy_atom, + v_reg, + ks_tiles, + vs_tiles, + scale_copy_atom, + scale_reg, logits_base, softmax_base, scale_base, stride_ks_block, stride_ks_head, softmax_scale_base, - softmax_q_scale, k_scale_val, - scale, v_scale_val, warp_id, lane16id, rowid, - k_tok_thread_base, - v_tok_thread_off, - vhead_elem_dw, - kv_tok_thread_base, - prob_wr_thread_base, - pv_prob_read_base, - sm_max_off, - sm_sum_off, - sm_rd_max_offs, - sm_rd_sum_offs, - sm_vmax_wr_off, - sm_vmax_rd_offs, - c_w, - neg_inf, - zero_f, - cache_scale_vecs=False, - head_size: int = 128, - qkhe_loop: int = 2, - vhe_loop: int = 2, + head_dim: int = 128, ): - apply_causal_mask = needs_mask or query_length > 1 + qkhe_loop = head_dim // QKHE_PER_FETCH + vhe_loop = head_dim // MFMA_N // NUM_WARPS + + vhead_elems = [ + fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + lane16id for vhe in range(vhe_loop) + ] + v_tok_thread_off = [fx.Int32(vt * TOKENS_PER_WARP) + rowid * fx.Int32(MFMA_N) for vt in range(VTLOOP)] + if const_expr(trans_v): + vhead_elem_dw = [vhead_elems[vhe] * fx.Int32(FP8_ELEMS_16B // 4) for vhe in range(vhe_loop)] + else: + vhead_elem_dw = [vhead_elems[vhe] * fx.Int32(KV_BLOCK_SIZE // 4) for vhe in range(vhe_loop)] + + kv_tok_thread_base = warp_id * fx.Int32(TOKENS_PER_WARP) + rowid * fx.Int32(4) + rowid_8x8 = rowid >> fx.Int32(1) + offset_in_slot = rowid & fx.Int32(1) + prob_row_i32 = PROB_ROW_STRIDE_BYTES // 4 + prob_row_i64 = PROB_ROW_STRIDE_BYTES // 8 + prob_wr_thread_base = ( + warp_id * fx.Int32(4 * MFMA_N * prob_row_i32) + + lane16id * fx.Int32(prob_row_i32) + + rowid_8x8 * fx.Int32(2) + + offset_in_slot + ) + pv_prob_read_base = rowid * fx.Int32(MFMA_N * prob_row_i64) + lane16id * fx.Int32(prob_row_i64) + + sm_lane_wave_base = lane16id * fx.Int32(NUM_WARPS) + sm_max_off = sm_lane_wave_base + warp_id + sm_sum_off = fx.Int32(NUM_WARPS * MFMA_N) + sm_lane_wave_base + warp_id + sm_rd_max_offs = [sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] + sm_rd_sum_offs = [fx.Int32(NUM_WARPS * MFMA_N) + sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] + + sm_vmax_wr_off = None + sm_vmax_rd_offs = None + if const_expr(per_token_kv): + sm_vmax_wr_off = fx.Int32(2 * NUM_WARPS * MFMA_N) + sm_lane_wave_base + warp_id + sm_vmax_rd_offs = [fx.Int32(2 * NUM_WARPS * MFMA_N) + sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] + + neg_inf = fx.Float32(float("-inf")) + zero_f = fx.Float32(0.0) + pv_prob_i64_elems = [] for vt in range_constexpr(VTLOOP): for j in range_constexpr(2): - p_elem = ( - arith.constant(vt * 4 * MFMA_N * (PROB_ROW_STRIDE_BYTES // 8), type=T.i32) - + pv_prob_read_base - + arith.constant(j, type=T.i32) - ) + p_elem = fx.Int32(vt * 4 * MFMA_N * (PROB_ROW_STRIDE_BYTES // 8)) + pv_prob_read_base + fx.Int32(j) pv_prob_i64_elems.append(p_elem) def _load_kv_scale_scalars(tile_token_offset_i32, phys_block): @@ -478,33 +438,20 @@ def _load_kv_scale_scalars(tile_token_offset_i32, phys_block): scale_block_base = phys_block * stride_ks_block + kv_h * stride_ks_head scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id scale_global_token = tile_token_offset_i32 + scale_stage_token - k_scale_scalar = buffer_ops.buffer_load( - ks_rsrc, - scale_block_base + scale_global_token, - vec_width=1, - dtype=fx.Float32, - ) - v_scale_scalar = buffer_ops.buffer_load( - vs_rsrc, - scale_block_base + scale_global_token, - vec_width=1, - dtype=fx.Float32, - ) + scale_offset = fx.Int32(scale_block_base + scale_global_token) + k_scale_scalar = _copy_load(ks_tiles, scale_offset, scale_copy_atom, scale_reg)[0] + v_scale_scalar = _copy_load(vs_tiles, scale_offset, scale_copy_atom, scale_reg)[0] return k_scale_scalar, v_scale_scalar - return None + return () def _load_v_and_scales( v_block_base_dw, tile_token_offset_i32, - *, - phys_block, - preloaded_scale_scalars=None, + scale_scalars, ): if const_expr(per_token_kv): scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id - if const_expr(preloaded_scale_scalars is None): - preloaded_scale_scalars = _load_kv_scale_scalars(tile_token_offset_i32, phys_block) - k_scale_scalar, v_scale_scalar = preloaded_scale_scalars + k_scale_scalar, v_scale_scalar = scale_scalars fx.ptr_store( fx.Vector.from_elements([k_scale_scalar], dtype=fx.Float32), scale_base + scale_stage_token, @@ -522,146 +469,93 @@ def _load_v_and_scales( v_token_in_block = tile_token_offset_i32 + v_tok_thread_off[vt] if const_expr(trans_v): vt_group = v_token_in_block >> fx.Int32(4) - va_dw_delta = ( - vt_group * arith.constant(head_size * FP8_ELEMS_16B // 4, type=T.i32) + vhead_elem_dw[vhe] - ) + va_dw_delta = vt_group * fx.Int32(head_dim * FP8_ELEMS_16B // 4) + vhead_elem_dw[vhe] else: va_dw_delta = vhead_elem_dw[vhe] + (v_token_in_block >> fx.Int32(2)) va_byte = (v_block_base_dw + fx.Int64(va_dw_delta)) * fx.Int64(4) - v_i64x2 = global_load_i64x2(v_global_ptr, va_byte) + v_i64x2 = _load_global_16b(v_global_ptr, va_byte, v_copy_atom, v_reg) vhe_data.append(v_i64x2) v_results.append(vhe_data) if const_expr(per_token_kv): gpu.barrier() - if const_expr(cache_scale_vecs): - k_scale_vecs = [] - v_scale_vecs = [] - for td in range_constexpr(TLOOP): - scale_row_base = kv_tok_thread_base + fx.Int32(td * MFMA_N) - k_scale_vecs.append( - fx.ptr_load(scale_base + (scale_row_base), result_type=fx.Vector.make_type(4, fx.Float32)) - ) - v_scale_vecs.append( - fx.ptr_load( - scale_base + (fx.Int32(LDS_SCALE_V_OFFSET) + scale_row_base), - result_type=fx.Vector.make_type(4, fx.Float32), - ) + k_scale_vecs = [] + v_scale_vecs = [] + for td in range_constexpr(TLOOP): + scale_row_base = kv_tok_thread_base + fx.Int32(td * MFMA_N) + k_scale_vecs.append( + fx.ptr_load(scale_base + (scale_row_base), result_type=fx.Vector.make_type(4, fx.Float32)) + ) + v_scale_vecs.append( + fx.ptr_load( + scale_base + (fx.Int32(LDS_SCALE_V_OFFSET) + scale_row_base), + result_type=fx.Vector.make_type(4, fx.Float32), ) - return v_results, k_scale_vecs, v_scale_vecs + ) + return v_results, k_scale_vecs, v_scale_vecs return v_results - def _scale_row_base(td: int): - return kv_tok_thread_base + fx.Int32(td * MFMA_N) - - def _load_k_scale_vec(td: int): - return fx.ptr_load(scale_base + (_scale_row_base(td)), result_type=fx.Vector.make_type(4, fx.Float32)) - - def _load_v_scale_vec(td: int): - return fx.ptr_load( - scale_base + (fx.Int32(LDS_SCALE_V_OFFSET) + _scale_row_base(td)), - result_type=fx.Vector.make_type(4, fx.Float32), - ) - - def _get_k_scale_vec(td: int, k_scale_vecs=None): - if const_expr(cache_scale_vecs): - return k_scale_vecs[td] - return _load_k_scale_vec(td) - - def _get_v_scale_vec(td: int, v_scale_vecs=None): - if const_expr(cache_scale_vecs): - return v_scale_vecs[td] - return _load_v_scale_vec(td) - def _store_vmax_warp(partition_start, *, seq_end=None, v_scale_vecs=None): if const_expr(per_token_kv): kv_tok_base = partition_start + kv_tok_thread_base if const_expr(seq_end is not None) else None v_max_warp = zero_f for td in range_constexpr(TLOOP): - vs = _get_v_scale_vec(td, v_scale_vecs) - for i in range_constexpr(4): - if const_expr(kv_tok_base is not None): - kv_tok = kv_tok_base + arith.constant(td * MFMA_N + i, type=T.i32) - vs_i = vector.extract(as_ir_value(vs), static_position=[i], dynamic_position=[]) - vs_i = arith.select(kv_tok < seq_end, vs_i, zero_f) - vs = vector.insert(vs_i, vs, static_position=[i], dynamic_position=[]) - v_max_warp = fx.maxnumf(v_max_warp, fx.Vector(vs).reduce("max")) + vs = fx.Vector(v_scale_vecs[td]) + if const_expr(kv_tok_base is not None): + vs = (_token_vec_i32(kv_tok_base, td) < seq_end).select(vs, zero_f) + v_max_warp = fx.maxnumf(v_max_warp, vs.reduce("max")) for sh in [32, 16]: - v_max_warp = fx.maxnumf(v_max_warp, v_max_warp.shuffle_xor(arith.constant(sh, type=T.i32), c_w)) + v_max_warp = fx.maxnumf(v_max_warp, v_max_warp.shuffle_xor(fx.Int32(sh), fx.Int32(WARP_SIZE))) fx.ptr_store( fx.Vector.from_elements([v_max_warp], dtype=fx.Float32), softmax_base + sm_vmax_wr_off, ) def _token_vec_i32(kv_tok_base, td: int): - kv_tok_td_base = kv_tok_base + arith.constant(td * MFMA_N, type=T.i32) + kv_tok_td_base = kv_tok_base + fx.Int32(td * MFMA_N) return fx.Vector.from_elements( - [kv_tok_td_base + arith.constant(i, type=T.i32) for i in range_constexpr(4)], + [kv_tok_td_base + fx.Int32(i) for i in range_constexpr(4)], dtype=fx.Int32, ) - def _apply_token_mask_vec(logit_vec, td: int, kv_tok_base, causal_bound, false_value): - tok_vec = _token_vec_i32(kv_tok_base, td) - if const_expr(apply_causal_mask): - in_range = tok_vec < causal_bound - return arith.select(in_range, logit_vec, vector.broadcast(T.f32x4, arith.unwrap(false_value))) - return logit_vec - def _qk_and_intra_softmax( k_ops, partition_start, q_frags, causal_bound, - query_scale_lane=None, + query_scale_lane, *, preloaded_scales=None, ): - if const_expr(preloaded_scales is not None): - if const_expr(cache_scale_vecs and per_token_kv): - k_scale_vecs, v_scale_vecs = preloaded_scales + if const_expr(per_token_kv): + k_scale_vecs, v_scale_vecs = preloaded_scales - query_scale_vec = None - if const_expr(per_token_q): - query_scale_vec = vector.broadcast(T.f32x4, query_scale_lane * softmax_scale_base) + query_scale = fx.Float32(query_scale_lane * softmax_scale_base) d_out = [] for td in range_constexpr(TLOOP): - acc = arith.constant_vector(0.0, T.f32x4) + acc = fx.Vector.filled(4, 0.0, fx.Float32) for k_step in range_constexpr(qkhe_loop * 2): acc = rocdl.mfma_f32_16x16x32_fp8_fp8(T.f32x4, [k_ops[td][k_step], q_frags[k_step], acc, 0, 0, 0]) if const_expr(per_token_kv): - if const_expr(cache_scale_vecs and per_token_kv): - k_scale_vec = _get_k_scale_vec(td, k_scale_vecs) - else: - k_scale_vec = _get_k_scale_vec(td) - scale_vec = ( - k_scale_vec * query_scale_vec - if const_expr(per_token_q) - else k_scale_vec * vector.broadcast(T.f32x4, softmax_q_scale) - ) - d_out.append(acc * scale_vec) + d_out.append(fx.Vector(acc) * (k_scale_vecs[td] * query_scale)) else: - if const_expr(per_token_q): - d_out.append(acc * (query_scale_vec * vector.broadcast(T.f32x4, k_scale_val))) - else: - d_out.append(acc * vector.broadcast(T.f32x4, scale)) + d_out.append(fx.Vector(acc) * (query_scale * k_scale_val)) - kv_tok_base = partition_start + kv_tok_thread_base if const_expr(apply_causal_mask) else None + kv_tok_base = partition_start + kv_tok_thread_base qk_max = neg_inf for td in range_constexpr(TLOOP): - logits_vec = d_out[td] - if const_expr(kv_tok_base is not None): - logits_vec = _apply_token_mask_vec(logits_vec, td, kv_tok_base, causal_bound, neg_inf) - d_out[td] = logits_vec + logits_vec = (_token_vec_i32(kv_tok_base, td) < causal_bound).select(fx.Vector(d_out[td]), neg_inf) + d_out[td] = logits_vec qk_max = fx.maxnumf(qk_max, fx.Vector(logits_vec).reduce("max")) for sh in [32, 16]: - qk_max = fx.maxnumf(qk_max, qk_max.shuffle_xor(arith.constant(sh, type=T.i32), c_w)) + qk_max = fx.maxnumf(qk_max, qk_max.shuffle_xor(fx.Int32(sh), fx.Int32(WARP_SIZE))) fx.ptr_store( fx.Vector.from_elements([qk_max], dtype=fx.Float32), softmax_base + sm_max_off, ) - if const_expr(cache_scale_vecs and per_token_kv): + if const_expr(per_token_kv): return d_out, v_scale_vecs return d_out @@ -673,27 +567,24 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs): partition_max = fx.maxnumf(partition_max, max_vec[w]) new_rmax = fx.maxnumf(rmax, partition_max) - safe_eff_max = arith.select(partition_max > neg_inf, new_rmax, zero_f) if const_expr(needs_mask) else new_rmax + safe_eff_max = arith.select(partition_max > neg_inf, new_rmax, zero_f) local_exp_sum = zero_f for td in range_constexpr(TLOOP): - diff_vec = fx.Vector(d_out[td]) - vector.broadcast(T.f32x4, arith.unwrap(safe_eff_max)) - p_vec = exp2_f32_fast(diff_vec * vector.broadcast(T.f32x4, arith.unwrap(fx.Float32(LOG2E)))) + diff_vec = fx.Vector(d_out[td]) - fx.Float32(safe_eff_max) + p_vec = exp2_f32_fast(diff_vec * fx.Float32(LOG2E)) local_exp_sum = local_exp_sum + fx.Vector(p_vec).reduce("add") d_out[td] = p_vec for sh in [32, 16]: - local_exp_sum = local_exp_sum + local_exp_sum.shuffle_xor(arith.constant(sh, type=T.i32), c_w) + local_exp_sum = local_exp_sum + local_exp_sum.shuffle_xor(fx.Int32(sh), fx.Int32(WARP_SIZE)) fx.ptr_store( fx.Vector.from_elements([local_exp_sum], dtype=fx.Float32), softmax_base + sm_sum_off, ) - if const_expr(needs_mask): - accum_scale = arith.select( - rmax > neg_inf, - exp2_f32_fast((rmax - new_rmax) * fx.Float32(LOG2E).ir_value()), - zero_f, - ) - else: - accum_scale = exp2_f32_fast((rmax - new_rmax) * fx.Float32(LOG2E).ir_value()) + accum_scale = arith.select( + rmax > neg_inf, + exp2_f32_fast((rmax - new_rmax) * fx.Float32(LOG2E).ir_value()), + zero_f, + ) gpu.barrier() sum_vec = fx.ptr_load(softmax_base + (sm_rd_sum_offs[0]), result_type=fx.Vector.make_type(4, fx.Float32)) @@ -705,9 +596,8 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs): accum_sum = arith.mulf(arith.unwrap(accum_scale), arith.unwrap(rsum), fastmath=arith.FastMathFlags.contract) rsum = arith.addf(accum_sum, arith.unwrap(partition_sum), fastmath=arith.FastMathFlags.contract) rmax = new_rmax - accum_scale_vec = vector.broadcast(T.f32x4, arith.unwrap(accum_scale)) for vhe in range_constexpr(vhe_loop): - outs[vhe] = outs[vhe] * accum_scale_vec + outs[vhe] = outs[vhe] * fx.Float32(accum_scale) if const_expr(per_token_kv): v_max_global = zero_f @@ -719,32 +609,25 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs): v_max_safe_scaled = v_max_scaled + fx.Float32(1e-8 / FP8_MAX).ir_value() norm_factor = rcp_f32(v_max_safe_scaled) v_correction = v_max_scaled - _vec_norm_p = arith.unwrap(norm_factor) for td in range_constexpr(TLOOP): - d_out[td] = d_out[td] * (_get_v_scale_vec(td, v_scale_vecs) * vector.broadcast(T.f32x4, _vec_norm_p)) + d_out[td] = d_out[td] * (v_scale_vecs[td] * fx.Float32(norm_factor)) else: v_correction = v_scale_val for td in range_constexpr(TLOOP): pv = fx.Vector(d_out[td]) - lo = rocdl.cvt_pk_fp8_f32(T.i32, pv[0], pv[1], arith.constant(0, type=T.i32), False) + lo = rocdl.cvt_pk_fp8_f32(T.i32, pv[0], pv[1], fx.Int32(0), False) pk = rocdl.cvt_pk_fp8_f32(T.i32, pv[2], pv[3], lo, True) - elem_base = prob_wr_thread_base + arith.constant(td * MFMA_N * (PROB_ROW_STRIDE_BYTES // 4), type=T.i32) + elem_base = prob_wr_thread_base + fx.Int32(td * MFMA_N * (PROB_ROW_STRIDE_BYTES // 4)) pk_vec = fx.Vector.from_elements([pk], dtype=fx.Int32) fx.ptr_store(pk_vec, logits_base + elem_base) return rmax, rsum, outs, v_correction def _pv_mfma(v_ops, outs, v_correction): - v_correction = fx.Float32(v_correction).ir_value() fm_contract = arith.FastMathFlags.contract - v_correction_vec = vector.broadcast(T.f32x4, v_correction) - - # ── Batch-load all P_i64 from LDS upfront ── - # `p_i64` depends only on (vt, j), NOT on vhe, so the previous - # per-vhe inner LDS load was redundant: VHELOOP × VTLOOP*2 reads - # of the same VTLOOP*2 LDS slots. Issue all VTLOOP*2 ds_read_b64 - # ops once at the start so the compiler pipelines them — lgkmcnt - # drains during the address arithmetic before the MFMA chain. + v_correction_vec = fx.Vector.filled(4, fx.Float32(v_correction), fx.Float32) + + # P depends only on (vt, j); load it once before every VHE MFMA chain. p_i64_all = [] for vt in range_constexpr(VTLOOP): for j in range_constexpr(2): @@ -757,7 +640,7 @@ def _pv_mfma(v_ops, outs, v_correction): ) for vhe in range_constexpr(vhe_loop): - tmp_out = arith.constant_vector(0.0, T.f32x4) + tmp_out = fx.Vector.filled(4, 0.0, fx.Float32) for vt in range_constexpr(VTLOOP): v_i64x2 = fx.Vector(v_ops[vt][vhe]) for j in range_constexpr(2): @@ -789,112 +672,6 @@ def _pv_mfma(v_ops, outs, v_correction): ) -_PA_DECODE_PS_SMALL_BLOCK_SIZES = (16, 64) - - -@flyc.jit -def _pa_small_block_load_k_flat( - k_global_ptr, - kv_h_i32, - stride_k_block_i32, - stride_k_head_i32, - lane16id_i32, - rowid_i32, - *, - block_size: int, - phys_blocks, - qkhe_loop: int = 2, -): - """Load K data for one warp's 64-token slice of a 256-token partition. - - Returns ``k_flat`` (a list of ``TLOOP * qkhe_loop * 2`` i64 scalars) compatible - with ``unflatten_k`` and downstream MFMA invocations. - """ - c_he_stride_dw = fx.Int32(block_size * FP8_ELEMS_16B // 4) - c_tok_stride_dw = fx.Int32(FP8_ELEMS_16B // 4) - k_he_off_dw = [rowid_i32 * c_he_stride_dw + fx.Int32(qkhe * 4) * c_he_stride_dw for qkhe in range(qkhe_loop)] - k_head_off = kv_h_i32 * stride_k_head_i32 - - k_flat = [] - if const_expr(block_size == 64): - # Each warp owns exactly one physical block (64 tokens). - phys_block = phys_blocks - k_block_base_dw = _compute_block_base_dw_i64(phys_block, stride_k_block_i32, k_head_off) - for td in range_constexpr(TLOOP): - within_block_token = fx.Int32(td * MFMA_N) + lane16id_i32 - kbo_dw = within_block_token * c_tok_stride_dw - for qkhe in range_constexpr(qkhe_loop): - ka_dw = k_block_base_dw + fx.Int64(kbo_dw + k_he_off_dw[qkhe]) - k2 = global_load_i64x2(k_global_ptr, ka_dw * fx.Int64(4)) - k2_words = fx.Vector(k2) - k_flat.append(k2_words[0]) - k_flat.append(k2_words[1]) - else: - # block_size == 16: each warp spans 4 blocks (one MFMA tile per block). - within_block_token = lane16id_i32 - kbo_dw = within_block_token * c_tok_stride_dw - for td in range_constexpr(TLOOP): - phys_block = phys_blocks[td] - k_block_base_dw = _compute_block_base_dw_i64(phys_block, stride_k_block_i32, k_head_off) - for qkhe in range_constexpr(qkhe_loop): - ka_dw = k_block_base_dw + fx.Int64(kbo_dw + k_he_off_dw[qkhe]) - k2 = global_load_i64x2(k_global_ptr, ka_dw * fx.Int64(4)) - rocdl.sched_barrier(rocdl.mask_vmem_rd) - k2_words = fx.Vector(k2) - k_flat.append(k2_words[0]) - k_flat.append(k2_words[1]) - return k_flat - - -@flyc.jit -def _pa_small_block_load_v_trans( - v_global_ptr, - kv_h_i32, - stride_v_block_i32, - stride_v_head_i32, - warp_id_i32, - lane16id_i32, - rowid_i32, - v_phys_blocks, - *, - block_size: int, - head_size: int = 128, - vhe_loop: int = 2, -): - """Load V tiles for one CTA's 256-token partition (``trans_v=True``). - - Returns ``v_results[vt][vhe]`` (i64x2) indexed exactly as the reference - ``_load_v_and_scales`` so it can be passed as ``preloaded_v_and_scales``. - """ - v_head_off = kv_h_i32 * stride_v_head_i32 - vhead_elems = [ - fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id_i32 * fx.Int32(MFMA_N) + lane16id_i32 for vhe in range(vhe_loop) - ] - vhead_elem_dw = [vhead_elems[vhe] * fx.Int32(FP8_ELEMS_16B // 4) for vhe in range(vhe_loop)] - c_subblock_dw = fx.Int32(head_size * FP8_ELEMS_16B // 4) - - v_results = [] - for vt in range_constexpr(VTLOOP): - phys_block = v_phys_blocks[vt] - if const_expr(block_size == 64): - # vt selects the physical block (4 blocks per partition); rowid - # selects the 16-token sub-block within that physical block. - sub_block_idx = rowid_i32 - else: - # block_size == 16: (vt * 4 + rowid) selects the block; only one - # 16-token sub-block per physical block, so sub_block_idx == 0. - sub_block_idx = fx.Int32(0) - v_block_base_dw = _compute_block_base_dw_i64(phys_block, stride_v_block_i32, v_head_off) - vhe_data = [] - for vhe in range_constexpr(vhe_loop): - va_dw_delta = sub_block_idx * c_subblock_dw + vhead_elem_dw[vhe] - va_byte = (v_block_base_dw + fx.Int64(va_dw_delta)) * fx.Int64(4) - v_i64x2 = global_load_i64x2(v_global_ptr, va_byte) - vhe_data.append(v_i64x2) - v_results.append(vhe_data) - return v_results - - @functools.lru_cache(maxsize=256) def compile_pa_metadata_v1( *, @@ -905,22 +682,10 @@ def compile_pa_metadata_v1( query_length: int, warp_size: int, ): - """Compile the FlyDSL worklist scheduler for a fixed device/shape config. - - ``num_batches`` stays a runtime kernel argument; everything else is baked. - - Launched as a single warp (``grid=block=(warp_size,1,1)``), matching aiter's - ``<<>>``: Phase-1 (per-batch block counts + sum) is - warp-parallel (lane-strided + warp reduce), while the serial CU x batch - scheduler runs uniformly on all lanes — exactly as the original, whose - work_indptr / work_info writes are not lane-guarded (benign same-value - races). The original's lane-divided / lane-0-guarded work is only for the - reduce_* maps, which the caller recomputes and we therefore omit. - """ + """Compile the single-wave worklist scheduler for a fixed shape config.""" assert is_pow2(kv_granularity), "kv_granularity must be power of 2" assert num_cu % num_heads_k == 0, "num_cu must be divisible by num_heads_k" num_splits_per_khead = num_cu // num_heads_k - # warp-reduce shuffle offsets: warp_size/2, ..., 1 _shuffle_offsets = [] _o = warp_size // 2 while _o >= 1: @@ -930,7 +695,6 @@ def compile_pa_metadata_v1( @flyc.kernel(known_block_size=(warp_size, 1, 1)) def pa_metadata_v1_kernel( seqlens_qo_indptr_ptr: fx.Tensor, # [num_batches + 1] i32 (cumulative qo seqlens) - pages_kv_indptr_ptr: fx.Tensor, # [num_batches + 1] i32 (cumulative pages) context_lens_ptr: fx.Tensor, # [num_batches] i32 work_indptr_ptr: fx.Tensor, # [num_cu + 1] i32 (output) work_info_ptr: fx.Tensor, # [max_work * 8] i32 (output, flattened) @@ -939,20 +703,28 @@ def pa_metadata_v1_kernel( reduce_partial_map_ptr: fx.Tensor, # [max_split_tiles] i32 (output) num_batches: Int32, ): - i32 = T.i32 - sq_rsrc = buffer_ops.create_buffer_resource(seqlens_qo_indptr_ptr, max_size=True) - ctx_rsrc = buffer_ops.create_buffer_resource(context_lens_ptr, max_size=True) - wi_rsrc = buffer_ops.create_buffer_resource(work_indptr_ptr, max_size=True) - winfo_rsrc = buffer_ops.create_buffer_resource(work_info_ptr, max_size=True) - rip_rsrc = buffer_ops.create_buffer_resource(reduce_indptr_ptr, max_size=True) - rfm_rsrc = buffer_ops.create_buffer_resource(reduce_final_map_ptr, max_size=True) - rpm_rsrc = buffer_ops.create_buffer_resource(reduce_partial_map_ptr, max_size=True) - # pages_kv_indptr_ptr is accepted for signature compat but unused: the - # work unit is a `partition_size`-token partition, so both the load - # balance and the kv ranges are counted as ceil(context_len/kv_gran) - # partitions (kv_granularity == partition_size). This matches the - # original kernel's kLdsBatchInfo=true path where curr_kv_pages is the - # per-batch num_blocks, not the pages_kv_indptr delta. + scalar_layout = fx.make_layout(1, 1) + work_layout = fx.make_layout(4, 1) + reduce_final_layout = fx.make_layout(2, 1) + + def _divide_buffer(tensor, tile_layout): + return fx.logical_divide(fx.rocdl.make_buffer_tensor(tensor), tile_layout) + + sq = _divide_buffer(seqlens_qo_indptr_ptr, scalar_layout) + context_lens = _divide_buffer(context_lens_ptr, scalar_layout) + work_indptr = _divide_buffer(work_indptr_ptr, scalar_layout) + work_info = _divide_buffer(work_info_ptr, work_layout) + reduce_indptr = _divide_buffer(reduce_indptr_ptr, scalar_layout) + reduce_final_map = _divide_buffer(reduce_final_map_ptr, reduce_final_layout) + reduce_partial_map = _divide_buffer(reduce_partial_map_ptr, scalar_layout) + + copy_i32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Int32) + copy_i32x2 = fx.make_copy_atom(fx.rocdl.BufferCopy64b(), fx.Int32) + copy_i32x4 = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Int32) + load_reg = fx.make_rmem_tensor(1, fx.Int32) + store_reg = fx.make_rmem_tensor(1, fx.Int32) + store_reg2 = fx.make_rmem_tensor(2, fx.Int32) + store_reg4 = fx.make_rmem_tensor(4, fx.Int32) c0 = fx.Int32(0) c1 = fx.Int32(1) @@ -964,33 +736,18 @@ def pa_metadata_v1_kernel( c_kvg = fx.Int32(kv_granularity) lane = fx.Int32(gpu.thread_id("x")) - def _load(rsrc, off): - return fx.Int32(buffer_ops.buffer_load(rsrc, fx.Int32(off).ir_value(), vec_width=1, dtype=i32)) - def _num_part(batch_idx): - # number of partition_size-token partitions for this batch = - # ceil(context_len[batch_idx] / kv_granularity) - ctxv = _load(ctx_rsrc, batch_idx) + ctxv = _copy_load(context_lens, batch_idx, copy_i32, load_reg)[0] return fx.Int32(arith.ceildivui(ctxv.ir_value(), c_kvg.ir_value())) - def _store(rsrc, off, val): - # NOTE: no masked stores — masked buffer_store sets OOB offset - # (0x7FFFFFFF) expecting HW bounds-check to drop it, but our - # max_size resources disable bounds-checking, so a masked store - # faults. All stores here are unconditional + overwrite-safe. - buffer_ops.buffer_store(fx.Int32(val).ir_value(), rsrc, fx.Int32(off).ir_value()) + def _store(divided, off, val): + fx.memref_store_vec(fx.Vector.from_elements([fx.Int32(val)], dtype=fx.Int32), store_reg) + fx.copy(copy_i32, store_reg, fx.slice(divided, (None, off))) - def _sel(cond_b, a, b): - return fx.Int32(arith.select(cond_b.ir_value(), fx.Int32(a).ir_value(), fx.Int32(b).ir_value())) + _store(work_indptr, 0, 0) + _store(reduce_indptr, 0, 0) - # work_indptr[0] = 0 ; reduce_indptr[0] = 0 - _store(wi_rsrc, 0, 0) - _store(rip_rsrc, 0, 0) - - # ---- Phase 1: sum_blocks = Sum_b ceil(context_lens[b] / kv_granularity) ---- - # (causal + uniform + tiny qo => effective_kv == context_lens[b]) - # warp-parallel: each lane sums batches {lane, lane+ws, lane+2ws, ...}, - # then a warp reduce-add gives the total in every lane (matches aiter). + # Lane-strided partition count followed by a wave reduction. b = lane sum_blocks = c0 while b < c_nb: @@ -998,18 +755,15 @@ def _sel(cond_b, a, b): b = b + c_ws sum_blocks = sum_blocks + nblk for sh in _shuffle_offsets: - sum_blocks = sum_blocks + sum_blocks.shuffle_xor(arith.constant(sh, type=i32), c_ws.ir_value()) + sum_blocks = sum_blocks + sum_blocks.shuffle_xor(fx.Int32(sh), c_ws) average = fx.Int32(arith.divui(sum_blocks.ir_value(), c_nspk.ir_value())) reminder = fx.Int32(arith.remui(sum_blocks.ir_value(), c_nspk.ir_value())) def _remain_for_cid(cid_val): - # remain = average + (1 if (cid % num_splits_per_khead) < reminder else 0) mod = fx.Int32(arith.remui(cid_val.ir_value(), c_nspk.ir_value())) - return average + _sel(mod < reminder, 1, 0) + return average + (mod < reminder).select(c1, c0) - # ---- Phase 2: per khead, flattened CU x batch scheduler ---- - # cid and num_works persist across kheads; the rest reset per khead. cid = c0 num_works = c0 @@ -1021,11 +775,6 @@ def _remain_for_cid(cid_val): kvend0 = _num_part(c0) # partitions in batch 0 (cumulative kv end) remain0 = _remain_for_cid(cid) - # State (11 i32), loop-carried through the scf.while emitted from the - # Python `while` below: - # 0 cid, 1 batch, 2 kvblk, 3 nsplit, 4 num_works, 5 pidx, - # 6 kvbeg, 7 kvend, 8 remain, 9 last_reduce_indptr, 10 global_reduce_tile_idx - # cid + num_works persist across kheads; lri + grt reset per khead. cid_ = cid batch_ = c0 kvblk_ = c0 @@ -1043,146 +792,129 @@ def _remain_for_cid(cid_val): remain_kv = pages - kvblk_ do_finish = remain_ >= remain_kv # fx bool - # qo_start/qo_end from seqlens_qo_indptr (the QoState array path, - # matching C++). batch_ < num_batches in the loop, so batch_+1 <= - # num_batches indexes the valid last element (no OOB). For uniform - # qo (sq = arange*query_length) this equals query_length*batch. - qo_start = _load(sq_rsrc, batch_) - qo_end = _load(sq_rsrc, batch_ + c1) + qo_start = _copy_load(sq, batch_, copy_i32, load_reg)[0] + qo_end = _copy_load(sq, batch_ + c1, copy_i32, load_reg)[0] kv_start = kvbeg_ + kvblk_ # same for both branches - # ---- finish branch (CU completes this batch) ---- f_kv_end = kvend_ # min(kv_start + remain_kv, kvend_) == kvend_ nsplit_pos = nsplit_ > c0 - f_ploc = _sel(nsplit_pos, pidx_, -1) - f_pidx2 = _sel(nsplit_pos, pidx_ + c_qlen, pidx_) + f_ploc = nsplit_pos.select(pidx_, fx.Int32(-1)) + f_pidx2 = nsplit_pos.select(pidx_ + c_qlen, pidx_) f_nworks2 = nworks_ + c1 f_remain2 = remain_ - remain_kv f_batch2 = batch_ + c1 - # next batch kv window (in partition units). max_size buffer rsrc - # disables HW bounds checking, so clamp the index before loading - # context_lens (f_batch2 can equal num_batches → OOB on last batch; - # the result is unused after the loop exits anyway). + # Select evaluates both values, so clamp the speculative load. nb_in_range = f_batch2 < c_nb - safe_idx = _sel(nb_in_range, f_batch2, 0) - f_new_pages = _sel(nb_in_range, _num_part(safe_idx), 0) + safe_idx = nb_in_range.select(f_batch2, c0) + f_new_pages = nb_in_range.select(_num_part(safe_idx), c0) f_kvbeg2 = kvend_ f_kvend2 = kvend_ + f_new_pages - # ---- split branch (CU does a partial; close cid) ---- s_emit = remain_ > c0 s_kv_end_raw = kv_start + remain_ - s_kv_end = _sel(s_kv_end_raw < kvend_, s_kv_end_raw, kvend_) - s_nworks2 = _sel(s_emit, nworks_ + c1, nworks_) - s_pidx2 = _sel(s_emit, pidx_ + c_qlen, pidx_) - s_kvblk2 = _sel(s_emit, kvblk_ + remain_, kvblk_) - s_nsplit2 = _sel(s_emit, nsplit_ + c1, nsplit_) + s_kv_end = (s_kv_end_raw < kvend_).select(s_kv_end_raw, kvend_) + s_nworks2 = s_emit.select(nworks_ + c1, nworks_) + s_pidx2 = s_emit.select(pidx_ + c_qlen, pidx_) + s_kvblk2 = s_emit.select(kvblk_ + remain_, kvblk_) + s_nsplit2 = s_emit.select(nsplit_ + c1, nsplit_) s_cid2 = cid_ + c1 s_remain2 = _remain_for_cid(s_cid2) - # ---- emit work entry at slot nworks_ ---- - # Overwrite-safe: if this step does not emit, nworks_ is unchanged - # so the slot is reused by the next emit (or lies beyond valid_work). - w_ploc = _sel(do_finish, f_ploc, pidx_) - w_kv_end = _sel(do_finish, f_kv_end, s_kv_end) - base = nworks_ * fx.Int32(_WORK_INFO_FIELDS) - _store(winfo_rsrc, base + fx.Int32(0), batch_) - _store(winfo_rsrc, base + fx.Int32(1), w_ploc) - _store(winfo_rsrc, base + fx.Int32(2), qo_start) - _store(winfo_rsrc, base + fx.Int32(3), qo_end) - _store(winfo_rsrc, base + fx.Int32(4), kv_start) - _store(winfo_rsrc, base + fx.Int32(5), w_kv_end) - _store(winfo_rsrc, base + fx.Int32(6), c0) - _store(winfo_rsrc, base + fx.Int32(7), fx.Int32(qhr_const)) - - # ---- reduce maps: only when finishing a SPLIT batch (nsplit>0) ---- - # This batch was split across (nsplit_+1) CUs and this CU finishes - # it, forming one reduce group. Faithful to the C++ kernel - # (kQoSplits=False path). + w_ploc = do_finish.select(f_ploc, pidx_) + w_kv_end = do_finish.select(f_kv_end, s_kv_end) + work_values = [ + batch_, + w_ploc, + qo_start, + qo_end, + kv_start, + w_kv_end, + c0, + fx.Int32(qhr_const), + ] + for half in range_constexpr(2): + start = half * 4 + fx.memref_store_vec( + fx.Vector.from_elements( + [work_values[start + field] for field in range_constexpr(4)], + dtype=fx.Int32, + ), + store_reg4, + ) + fx.copy( + copy_i32x4, + store_reg4, + fx.slice(work_info, (None, nworks_ * fx.Int32(2) + fx.Int32(half))), + ) + do_reduce = do_finish & nsplit_pos num_splits = nsplit_ + c1 - # reduce_indptr[grt+1] = lri + num_splits ; reduce_final_map[grt] = - # (qo_start, qo_end). Unconditional + overwrite-safe (same argument - # as work_indptr: grt only advances on do_reduce, so non-do_reduce - # writes to grt+1 are overwritten by the next do_reduce or the tail; - # reduce_final_map[grt*2..] beyond the final grt is never read). - _store(rip_rsrc, grt_ + c1, lri_ + num_splits) - _store(rfm_rsrc, grt_ * fx.Int32(2), qo_start) - _store(rfm_rsrc, grt_ * fx.Int32(2) + c1, qo_end) - # reduce_partial_map[lri + s] = pidx - (nsplit - s)*qlen, s in [0,num_splits) - # nested loop runs num_splits times when do_reduce, else 0 times. - rcount = _sel(do_reduce, num_splits, 0) + # Non-reduce writes are overwritten before their slots become visible. + _store(reduce_indptr, grt_ + c1, lri_ + num_splits) + fx.memref_store_vec( + fx.Vector.from_elements([qo_start, qo_end], dtype=fx.Int32), + store_reg2, + ) + fx.copy(copy_i32x2, store_reg2, fx.slice(reduce_final_map, (None, grt_))) + rcount = do_reduce.select(num_splits, c0) sidx = c0 while sidx < rcount: val = pidx_ - (nsplit_ - sidx) * c_qlen - _store(rpm_rsrc, lri_ + sidx, val) + _store(reduce_partial_map, lri_ + sidx, val) sidx = sidx + c1 - n_lri = lri_ + _sel(do_reduce, num_splits, 0) - n_grt = grt_ + _sel(do_reduce, 1, 0) - - # ---- new state via select(do_finish, finish, split) ---- - n_cid = _sel(do_finish, cid_, s_cid2) - n_batch = _sel(do_finish, f_batch2, batch_) - n_kvblk = _sel(do_finish, 0, s_kvblk2) - n_nsplit = _sel(do_finish, 0, s_nsplit2) - n_nworks = _sel(do_finish, f_nworks2, s_nworks2) - n_pidx = _sel(do_finish, f_pidx2, s_pidx2) - n_kvbeg = _sel(do_finish, f_kvbeg2, kvbeg_) - n_kvend = _sel(do_finish, f_kvend2, kvend_) - n_remain = _sel(do_finish, f_remain2, s_remain2) - - # ---- work_indptr[cid_+1] = n_nworks (unconditional, overwrite-safe) ---- - # In the finish branch cid_ does not advance, so repeated writes to - # work_indptr[cid_+1] keep updating until the cid closes; the last - # write (before cid advances in the split branch, or loop exit) holds - # the correct running num_works for that cid. cid_+1 <= num_cu (loop - # guard cid_ < num_cu) so the index is always in-bounds. - _store(wi_rsrc, cid_ + c1, n_nworks) - - # reassign loop-carried state (becomes the scf.while yield) - cid_ = n_cid - batch_ = n_batch - kvblk_ = n_kvblk - nsplit_ = n_nsplit - nworks_ = n_nworks - pidx_ = n_pidx - kvbeg_ = n_kvbeg - kvend_ = n_kvend - remain_ = n_remain - lri_ = n_lri - grt_ = n_grt + next_num_works = do_finish.select(f_nworks2, s_nworks2) + + # The last same-value-race write before advancing cid is authoritative. + _store(work_indptr, cid_ + c1, next_num_works) + + ( + cid_, + batch_, + kvblk_, + nsplit_, + nworks_, + pidx_, + kvbeg_, + kvend_, + remain_, + lri_, + grt_, + ) = ( + do_finish.select(cid_, s_cid2), + do_finish.select(f_batch2, batch_), + do_finish.select(c0, s_kvblk2), + do_finish.select(c0, s_nsplit2), + next_num_works, + do_finish.select(f_pidx2, s_pidx2), + do_finish.select(f_kvbeg2, kvbeg_), + do_finish.select(f_kvend2, kvend_), + do_finish.select(f_remain2, s_remain2), + lri_ + do_reduce.select(num_splits, c0), + grt_ + do_reduce.select(c1, c0), + ) cid = cid_ num_works = nworks_ last_reduce_indptr = lri_ global_reduce_tile_idx = grt_ - # ---- post-khead close: advance cid past the last processed cid so - # the next khead (and the tail) start fresh. The loop already wrote - # work_indptr[last_cid+1]=num_works on its final iteration, so no - # store is needed here — only the cid advance. in_range = cid < c_numcu - cid = _sel(in_range, cid + c1, cid) + cid = in_range.select(cid + c1, cid) - # ---- tail: work_indptr[i] = num_works for i in [cid, num_cu] ---- it_t = cid while it_t <= c_numcu: - _store(wi_rsrc, it_t, num_works) + _store(work_indptr, it_t, num_works) it_t = it_t + c1 - # ---- tail: reduce_indptr[i] = last_reduce_indptr for i in [grt, num_batches] ---- - # (reduce_indptr has num_batches+1 entries; uses the final khead's grt/lri, - # matching the original which resets grt per khead and fills the tail once.) c_rip_size = c_nb + c1 # reduce_indptr length = num_batches + 1 it_r = global_reduce_tile_idx while it_r < c_rip_size: - _store(rip_rsrc, it_r, last_reduce_indptr) + _store(reduce_indptr, it_r, last_reduce_indptr) it_r = it_r + c1 @flyc.jit def launch_pa_metadata_v1( seqlens_qo_indptr: fx.Tensor, - pages_kv_indptr: fx.Tensor, context_lens: fx.Tensor, work_indptr: fx.Tensor, work_info: fx.Tensor, @@ -1194,7 +926,6 @@ def launch_pa_metadata_v1( ): pa_metadata_v1_kernel( seqlens_qo_indptr, - pages_kv_indptr, context_lens, work_indptr, work_info, @@ -1209,72 +940,46 @@ def launch_pa_metadata_v1( def get_pa_metadata_v1( seqlens_qo_indptr: torch.Tensor, - pages_kv_indptr: torch.Tensor, context_lens: torch.Tensor, - num_heads_per_head_k: int, - num_heads_k: int, - is_causal: bool, - work_metadata_ptrs: torch.Tensor, work_indptr: torch.Tensor, work_info: torch.Tensor, reduce_indptr: torch.Tensor, reduce_final_map: torch.Tensor, reduce_partial_map: torch.Tensor, - kv_granularity: int = 16, - block_size: int = 16, - max_seqlen_qo: int = -1, - uni_seqlen_qo: int = -1, - fast_mode: bool = True, - topk: int = -1, - max_split_per_batch: int = -1, + *, + query_group_size: int, + num_kv_heads: int, + kv_granularity: int, + query_length: int, num_cu: int = None, + stream=None, ) -> None: - """Drop-in replacement for ``aiter.ops.attention.get_pa_metadata_v1``. - - PA-decode-specialized: requires causal, non-sparse, uniform qo. Fills - ``work_indptr`` and ``work_info`` in-place; ``reduce_*`` are left untouched - (the caller recomputes them from work_indptr/work_info). + """Build the PA-decode worklist and reduction maps. ``num_cu`` overrides the worklist bin count (default = device CU count); pass a multiple of the CU count to oversubscribe the persistent grid. """ - assert is_causal, "FlyDSL pa_metadata only supports causal" - assert topk == -1, "FlyDSL pa_metadata does not support sparse (topk)" - assert uni_seqlen_qo >= 1, "FlyDSL pa_metadata requires uniform qo length" - - dev = pages_kv_indptr.device + dev = context_lens.device if num_cu is None: num_cu = torch.cuda.get_device_properties(dev).multi_processor_count num_batches = context_lens.shape[0] - query_length = uni_seqlen_qo warp_size = get_warp_size(get_rocm_arch()) compiled = compile_pa_metadata_v1( num_cu=num_cu, - num_heads_k=num_heads_k, - gqa=num_heads_per_head_k, + num_heads_k=num_kv_heads, + gqa=query_group_size, kv_granularity=kv_granularity, query_length=query_length, warp_size=warp_size, ) - # work_metadata_ptrs[0/1] = device addresses of work_indptr / work_info, - # matching the C++ kernel (which writes them in-kernel via reinterpret_cast). - # These are exactly the tensors' data_ptr() values, so writing them host-side - # produces identical bytes. - if work_metadata_ptrs is not None and work_metadata_ptrs.numel() >= 2: - work_metadata_ptrs[0] = work_indptr.data_ptr() - work_metadata_ptrs[1] = work_info.data_ptr() - - # work_info [max_work, 8] and reduce_final_map [num_batches, 2] are written - # flattened by the kernel. work_info_flat = work_info.view(-1) reduce_final_map_flat = reduce_final_map.view(-1) _run_compiled( compiled["launch"], seqlens_qo_indptr, - pages_kv_indptr, context_lens, work_indptr, work_info_flat, @@ -1282,22 +987,18 @@ def get_pa_metadata_v1( reduce_final_map_flat, reduce_partial_map, num_batches, - fx.Stream(None), + stream, ) -# ===================================================================== -# compile_pa_decode_metadata — Persistent Scheduling PA decode kernel -# ===================================================================== @functools.lru_cache(maxsize=256) def compile_pa_decode_metadata( softmax_scale=None, trans_v=False, - needs_mask=True, query_group_size=16, per_token_kv=False, query_length: int = 1, - query_input_dtype: str = "packed_fp8", + query_input_dtype: str = "bf16", head_dim: int = 128, block_size: int = None, output_dtype_str: str = "bf16", @@ -1309,82 +1010,48 @@ def compile_pa_decode_metadata( The worklist is load-balanced at ``KV_COMPUTE_BLOCK`` (256-token) **partition** granularity (see ``get_pa_metadata``): ``work_info.kv_start/kv_end`` are - cumulative partition indices. Each work item is decoded as a range of - 256-token partitions; for ``block_size < 256`` each partition gathers - ``256 // block_size`` physical pages, and for ``block_size > 256`` (1024) each - partition is a 256-token sub-tile of one physical page. ``partial_qo_loc`` - (``work_info[1]``) ``< 0`` writes the final output directly to ``out``; + cumulative partition indices. Only 1024-token physical pages are supported; + each page contains four 256-token partitions. ``partial_qo_loc`` + (``work_info[1]``) ``< 0`` writes the final output directly to ``out``, while ``>= 0`` writes a partial slot that ``pa_reduce_v1`` later combines. """ if block_size is None: block_size = KV_BLOCK_SIZE + if block_size != KV_BLOCK_SIZE: + raise ValueError(f"compile_pa_decode_metadata only supports block_size={KV_BLOCK_SIZE}, got {block_size}") if head_dim % QKHE_PER_FETCH != 0 or head_dim % (MFMA_N * NUM_WARPS) != 0 or head_dim % Q_ELEMS_PER_LANE != 0: raise ValueError(f"Unsupported head_dim={head_dim}; must be a multiple of {MFMA_N * NUM_WARPS}.") - _HEAD = head_dim _QKHELOOP = head_dim // QKHE_PER_FETCH _VHELOOP = head_dim // MFMA_N // NUM_WARPS _Q_LANES_PER_HEAD = head_dim // Q_ELEMS_PER_LANE _N_K_h = TLOOP * _QKHELOOP * 2 - query_packed_fp8 = query_input_dtype == "packed_fp8" - query_load_is_bf16 = query_input_dtype == "bf16" - query_scale_in_kernel = not query_packed_fp8 - cache_scale_vecs = True - if const_expr(query_packed_fp8): - raise ValueError( - "`compile_pa_decode_metadata` only supports bf16/f16 queries with kernel-internal query scale." - ) + if query_input_dtype not in ("bf16", "f16"): + raise ValueError(f"`compile_pa_decode_metadata` only supports bf16/f16 queries, got {query_input_dtype!r}") + _QUERY_DTYPE = fx.BFloat16 if query_input_dtype == "bf16" else fx.Float16 + _OUTPUT_DTYPE = _PA_OUTPUT_DTYPES[output_dtype_str] if softmax_scale is None: softmax_scale = 1.0 / (head_dim**0.5) _softmax_scale = float(softmax_scale) - _block_size = int(block_size) - # A partition is KV_COMPUTE_BLOCK (256) tokens. For small blocks each - # partition gathers ``_blocks_per_partition`` physical pages; for block_size - # >= 256 each physical page holds ``_parts_per_block`` partitions (sub-tiles). - _is_small_block = _block_size < KV_COMPUTE_BLOCK - _blocks_per_partition = KV_COMPUTE_BLOCK // _block_size if _is_small_block else 1 - _parts_per_block = _block_size // KV_COMPUTE_BLOCK if not _is_small_block else 1 - if _is_small_block: - if _block_size not in _PA_DECODE_PS_SMALL_BLOCK_SIZES: - raise ValueError( - f"compile_pa_decode_metadata: unsupported small block_size={_block_size}; " - f"expected one of {_PA_DECODE_PS_SMALL_BLOCK_SIZES} or >= {KV_COMPUTE_BLOCK}." - ) - if per_token_kv: - raise NotImplementedError( - "compile_pa_decode_metadata: per_token_kv=True is not supported for " - "small block_size 16/64; small blocks use compile_pa_decode_ps." - ) - if not trans_v: - raise NotImplementedError( - "compile_pa_decode_metadata: trans_v=False is not supported for small block_size 16/64." - ) + parts_per_block = KV_BLOCK_SIZE // KV_COMPUTE_BLOCK - # LDS allocation - # Extra LDS for cross-warp v_scale_max reduction (per_token_kv only): - # NUM_WARPS floats per lane16id slot, aligned to same layout as softmax data. + # Per-token mode adds one cross-warp vmax slot per lane. LDS_VMAX_BYTES = NUM_WARPS * MFMA_N * 4 if const_expr(per_token_kv) else 0 # 256 or 0 LDS_SOFTMAX_TOTAL = LDS_SOFTMAX_BYTES + LDS_VMAX_BYTES LDS_SCALE_TOTAL = LDS_SCALE_BYTES if const_expr(per_token_kv) else 0 - logits_off = 0 softmax_off = LDS_LOGITS_BYTES scale_off = softmax_off + LDS_SOFTMAX_TOTAL - # Phys-block staging LDS for the small-block path (cross-warp visibility of - # the per-warp page indices so V can read all blocks of a partition). - bt_off = scale_off + LDS_SCALE_TOTAL - _LDS_TOTAL_BYTES = bt_off + (NUM_WARPS * TLOOP * 4 if _is_small_block else 0) + _LDS_TOTAL_BYTES = scale_off + LDS_SCALE_TOTAL @fx.struct class SharedStorage: buf: fx.Array[fx.Int32, _LDS_TOTAL_BYTES // 4, 16] - # ── @flyc.kernel ───────────────────────────────────────────────── @flyc.kernel(known_block_size=(BLOCK_THREADS, 1, 1)) def pa_decode_metadata_kenrel( - # Raw-pointer kernargs: bare i64 data_ptr() (strides are explicit args). - out_ptr: fx.Int64, # output [batch, num_q_heads, head_size] + out_ptr: fx.Int64, # output [batch, num_q_heads, head_dim] partial_out_ptr: fx.Int64, # partial output [num_partials, 1, nhead, head_dim] fp32 partial_lse_ptr: fx.Int64, # partial LSE [num_partials, 1, nhead, 1] fp32 - query_ptr: fx.Int64, # queries [batch, num_q_heads, head_size] + query_ptr: fx.Int64, # queries [batch, num_q_heads, head_dim] key_cache_ptr: fx.Int64, # key cache value_cache_ptr: fx.Int64, # value cache context_lengths_ptr: fx.Int64, # [batch] int32 @@ -1407,149 +1074,114 @@ def pa_decode_metadata_kenrel( stride_pl_partial: Int32, # stride for partial_lse partial dim (nhead) stride_ks_block: Int32, # key_scale stride for block dim (num_kv_heads * KV_BLOCK_SIZE); 0 for per-tensor stride_ks_head: Int32, # key_scale stride for head dim (KV_BLOCK_SIZE); 0 for per-tensor - stride_po_ql: Int32, # stride for partial_output query-length dim (num_query_heads * head_size) + stride_po_ql: Int32, # stride for partial_output query-length dim (num_query_heads * head_dim) stride_pl_ql: Int32, # stride for partial_lse query-length dim (num_query_heads) ): - tid = gpu.thread_idx.x - cu_id = gpu.block_idx.x # CU index (0..num_sm-1) - - # ── Thread decomposition ── - lane16id = tid & arith.constant(15, type=T.i32) - rowid = (tid >> arith.constant(4, type=T.i32)) & arith.constant(3, type=T.i32) - warp_id = tid >> arith.constant(6, type=T.i32) - - # ── Buffer resources ── - q_rsrc = buffer_ops.create_buffer_resource_from_addr(query_ptr) - out_rsrc = buffer_ops.create_buffer_resource_from_addr(out_ptr) - k_global_ptr = global_ptr_from_addr(key_cache_ptr) - v_global_ptr = global_ptr_from_addr(value_cache_ptr) - po_rsrc = buffer_ops.create_buffer_resource_from_addr(partial_out_ptr) - pl_rsrc = buffer_ops.create_buffer_resource_from_addr(partial_lse_ptr) - cl_rsrc = buffer_ops.create_buffer_resource_from_addr(context_lengths_ptr) - wi_rsrc = buffer_ops.create_buffer_resource_from_addr(work_indptr_ptr) - winfo_rsrc = buffer_ops.create_buffer_resource_from_addr(work_info_ptr) - kpi_rsrc = buffer_ops.create_buffer_resource_from_addr(kv_page_indices_ptr) - kvindptr_rsrc = buffer_ops.create_buffer_resource_from_addr(kv_indptr_ptr) - pip_rsrc = buffer_ops.create_buffer_resource_from_addr(partition_indptr_ptr) - ks_rsrc = buffer_ops.create_buffer_resource_from_addr(key_scale_ptr) - vs_rsrc = buffer_ops.create_buffer_resource_from_addr(value_scale_ptr) - - q_scale_val = arith.constant(1.0, type=T.f32) + tid = fx.Int32(gpu.thread_id("x")) + cu_id = fx.Int32(gpu.block_id("x")) # CU index (0..num_sm-1) + + lane16id = tid & fx.Int32(15) + rowid = (tid >> fx.Int32(4)) & fx.Int32(3) + warp_id = tid >> fx.Int32(6) + + def _divide_addr(addr, dtype, width): + ptr = _global_pointer_from_addr(addr, dtype, alignment=width * dtype.width // 8) + flat = fx.make_view(ptr, fx.make_layout(_FLAT_BUFFER_ELEMENTS, 1)) + return fx.logical_divide( + fx.rocdl.make_buffer_tensor(flat), + fx.make_layout(width, 1), + ) + + q_tiles = _divide_addr(query_ptr, _QUERY_DTYPE, 4) + out_tiles = _divide_addr(out_ptr, _OUTPUT_DTYPE, 4) + partial_out_tiles = _divide_addr(partial_out_ptr, fx.Float32, 4) + partial_lse = _divide_addr(partial_lse_ptr, fx.Float32, 1) + context_lengths = _divide_addr(context_lengths_ptr, fx.Int32, 1) + work_indptr = _divide_addr(work_indptr_ptr, fx.Int32, 1) + work_info = _divide_addr(work_info_ptr, fx.Int32, 4) + kv_page_indices = _divide_addr(kv_page_indices_ptr, fx.Int32, 1) + kv_indptr = _divide_addr(kv_indptr_ptr, fx.Int32, 1) + partition_indptr = _divide_addr(partition_indptr_ptr, fx.Int32, 1) + key_scales = _divide_addr(key_scale_ptr, fx.Float32, 1) + value_scales = _divide_addr(value_scale_ptr, fx.Float32, 1) + + copy_i32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Int32) + copy_i32x4 = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Int32) + copy_f32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Float32) + copy_q = fx.make_copy_atom(fx.rocdl.BufferCopy64b(), _QUERY_DTYPE) + copy_out = fx.make_copy_atom( + fx.rocdl.BufferCopy128b() if _OUTPUT_DTYPE.width == 32 else fx.rocdl.BufferCopy64b(), + _OUTPUT_DTYPE, + ) + copy_f32x4 = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Float32) + + i32_reg = fx.make_rmem_tensor(1, fx.Int32) + i32x4_reg = fx.make_rmem_tensor(4, fx.Int32) + scale_reg = fx.make_rmem_tensor(1, fx.Float32) + q_reg = fx.make_rmem_tensor(4, _QUERY_DTYPE) + out_reg = fx.make_rmem_tensor(4, _OUTPUT_DTYPE) + partial_out_reg = fx.make_rmem_tensor(4, fx.Float32) + partial_lse_reg = fx.make_rmem_tensor(1, fx.Float32) + + k_global_ptr = _global_pointer_from_addr(key_cache_ptr, fx.Uint8, alignment=16) + v_global_ptr = _global_pointer_from_addr(value_cache_ptr, fx.Uint8, alignment=16) + global_copy_16b = fx.make_copy_atom(fx.UniversalCopy128b(), fx.Uint8) + global_reg_16b = fx.make_rmem_tensor(16, fx.Uint8) + if const_expr(per_token_kv): - k_scale_val = arith.constant(1.0, type=T.f32) - v_scale_val = arith.constant(1.0, type=T.f32) + k_scale_val = fx.Float32(1.0) + v_scale_val = fx.Float32(1.0) else: - k_scale_val = buffer_ops.buffer_load(ks_rsrc, arith.constant(0, type=T.i32), vec_width=1) - v_scale_val = buffer_ops.buffer_load(vs_rsrc, arith.constant(0, type=T.i32), vec_width=1) + k_scale_val = _copy_load(key_scales, 0, copy_f32, scale_reg)[0] + v_scale_val = _copy_load(value_scales, 0, copy_f32, scale_reg)[0] lds = fx.SharedAllocator().allocate(SharedStorage).peek() - buf_base = lds.buf.ptr - logits_base = fx.add_offset(buf_base, fx.Int32(logits_off // 4)) - softmax_base = fx.recast_iter(fx.Float32, fx.add_offset(buf_base, fx.Int32(softmax_off // 4))) + logits_base = lds.buf.ptr + softmax_base = fx.recast_iter(fx.Float32, fx.add_offset(lds.buf.ptr, fx.Int32(softmax_off // 4))) scale_base = None if const_expr(per_token_kv): - scale_base = fx.recast_iter(fx.Float32, fx.add_offset(buf_base, fx.Int32(scale_off // 4))) - bt_base = None - if const_expr(_is_small_block): - bt_base = fx.add_offset(buf_base, fx.Int32(bt_off // 4)) - - # ── Constants ── - c_kb = stride_k_block - c_kh = stride_k_head - c_vb = stride_v_block - c_vh = stride_v_head - - _softmax_scale_const = arith.constant(_softmax_scale, type=T.f32) - _softmax_q_scale = _softmax_scale_const * q_scale_val - _scale = _softmax_q_scale * k_scale_val # per-tensor only; per-token uses per-token k_scale - c_w = arith.constant(WARP_SIZE, type=T.i32) - NEG_INF = arith.constant(float("-inf"), type=T.f32) - ZERO_F = arith.constant(0.0, type=T.f32) - c_cps = arith.constant(KV_COMPUTE_BLOCK, type=T.i32) # 256-token partition - c_one = arith.constant(1, type=T.i32) - - local_qhead_idx = warp_id * arith.constant(4, type=T.i32) + rowid - ( - _k_tok_thread_base, - _c_tok_stride_dw, - _k_he_off_dw, - _v_tok_thread_off, - _vhead_elem_dw, - _kv_tok_thread_base, - _prob_wr_thread_base, - _pv_prob_read_base, - _sm_max_off, - _sm_sum_off, - _sm_rd_max_offs, - _sm_rd_sum_offs, - _sm_vmax_wr_off, - _sm_vmax_rd_offs, - ) = _build_pa_thread_invariants( + scale_base = fx.recast_iter(fx.Float32, fx.add_offset(lds.buf.ptr, fx.Int32(scale_off // 4))) + + _softmax_scale_const = fx.Float32(_softmax_scale) + NEG_INF = fx.Float32(float("-inf")) + ZERO_F = fx.Float32(0.0) + c_cps = fx.Int32(KV_COMPUTE_BLOCK) # 256-token partition + c_one = fx.Int32(1) + + local_qhead_idx = warp_id * fx.Int32(4) + rowid + _k_tok_thread_base, _c_tok_stride_dw, _k_he_off_dw = _build_pa_k_thread_invariants( warp_id, lane16id, rowid, - trans_v=trans_v, - per_token_kv=per_token_kv, qkhe_loop=_QKHELOOP, - vhe_loop=_VHELOOP, ) - # ── Work loop bounds ── - # wi[cu_id] and wi[cu_id+1] are adjacent int32; load both in one vec2 load. - work_bounds = buffer_ops.buffer_load(wi_rsrc, cu_id, vec_width=2, dtype=T.i32) - work_start = fx.Vector(work_bounds)[0] - work_end = fx.Vector(work_bounds)[1] - - # Outer work loop — each work item = one (batch, kv_head_range, kv_page_range) - _work_start_idx = fx.Index(arith.unwrap(work_start)) - _work_end_idx = fx.Index(arith.unwrap(work_end)) - _work_step = arith.index(1) - - for _wi in range(_work_start_idx, _work_end_idx, _work_step): - work_idx = arith.index_cast(T.i32, _wi) - - # ── Load work_info[work_idx] — 8 × int32, as 2 × vec4 loads ── - # info_base is a multiple of 8, so both dwordx4 loads are naturally - # aligned (info_base @ 32 B, info_base+4 @ 16 B). Fields 3 and 6 are - # currently unused and simply not extracted. - info_base = work_idx * arith.constant(8, type=T.i32) - wi_lo = buffer_ops.buffer_load(winfo_rsrc, info_base, vec_width=4, dtype=T.i32) - wi_hi = buffer_ops.buffer_load( - winfo_rsrc, info_base + arith.constant(4, type=T.i32), vec_width=4, dtype=T.i32 - ) - wi_lo_v = fx.Vector(wi_lo) + work_start = _copy_load(work_indptr, cu_id, copy_i32, i32_reg)[0] + work_end = _copy_load(work_indptr, cu_id + fx.Int32(1), copy_i32, i32_reg)[0] + + for work_idx in range(work_start, work_end, fx.Int32(1)): + # Two naturally aligned dwordx4 loads cover the eight work fields. + info_tile = work_idx * fx.Int32(2) + wi_lo_v = _copy_load(work_info, info_tile, copy_i32x4, i32x4_reg) batch_idx = wi_lo_v[0] partial_idx = wi_lo_v[1] qo_start = wi_lo_v[2] - wi_hi_v = fx.Vector(wi_hi) + wi_hi_v = _copy_load(work_info, info_tile + fx.Int32(1), copy_i32x4, i32x4_reg) kv_start = wi_hi_v[0] kv_end = wi_hi_v[1] q_head_range = wi_hi_v[3] - # work_info.kv_start/kv_end are cumulative partition indices (256-token - # units, summed across batches). partition_indptr[batch] gives the - # cumulative-partition base for this sequence (→ local partition index), - # kv_indptr[batch] gives the physical-page base into kv_page_indices. - kv_part_base = buffer_ops.buffer_load(pip_rsrc, batch_idx, vec_width=1, dtype=T.i32) - # kv_indptr[batch] / kv_indptr[batch+1] in one dwordx2 load: this - # sequence's physical-page base and end in the flat kv_page_indices - # array. kv_page_end clamps small-block page-gather reads so the last - # (partial) partition never reads past the sequence. - _kvind2 = buffer_ops.buffer_load(kvindptr_rsrc, batch_idx, vec_width=2, dtype=T.i32) - _kvind2_v = fx.Vector(_kvind2) - kv_page_base = _kvind2_v[0] - kv_page_end = _kvind2_v[1] + # Convert cumulative partition indices to sequence-local page indices. + kv_part_base = _copy_load(partition_indptr, batch_idx, copy_i32, i32_reg)[0] + kv_page_base = _copy_load(kv_indptr, batch_idx, copy_i32, i32_reg)[0] local_part_start = kv_start - kv_part_base - # Derive kv_head from q_head_range - q_head_start = q_head_range & arith.constant(0xFFFF, type=T.i32) + q_head_start = q_head_range & fx.Int32(0xFFFF) kv_h = udiv_const(q_head_start, query_group_size) - # Context length for this sequence - context_len = buffer_ops.buffer_load(cl_rsrc, batch_idx, vec_width=1, dtype=T.i32) - # Head offsets for K and V cache - _k_head_off = kv_h * c_kh - _v_head_off = kv_h * c_vh + context_len = _copy_load(context_lengths, batch_idx, copy_i32, i32_reg)[0] + _k_head_off = kv_h * stride_k_head + _v_head_off = kv_h * stride_v_head ( _load_kv_scale_scalars, @@ -1560,132 +1192,45 @@ def pa_decode_metadata_kenrel( _pv_mfma, ) = _make_pa_phase_helpers( trans_v=trans_v, - per_token_q=query_scale_in_kernel, per_token_kv=per_token_kv, - needs_mask=needs_mask, - query_length=query_length, kv_h=kv_h, v_global_ptr=v_global_ptr, - ks_rsrc=ks_rsrc, - vs_rsrc=vs_rsrc, + v_copy_atom=global_copy_16b, + v_reg=global_reg_16b, + ks_tiles=key_scales, + vs_tiles=value_scales, + scale_copy_atom=copy_f32, + scale_reg=scale_reg, logits_base=logits_base, softmax_base=softmax_base, scale_base=scale_base, stride_ks_block=stride_ks_block, stride_ks_head=stride_ks_head, softmax_scale_base=_softmax_scale_const, - softmax_q_scale=_softmax_q_scale, k_scale_val=k_scale_val, - scale=_scale, v_scale_val=v_scale_val, warp_id=warp_id, lane16id=lane16id, rowid=rowid, - k_tok_thread_base=_k_tok_thread_base, - v_tok_thread_off=_v_tok_thread_off, - vhead_elem_dw=_vhead_elem_dw, - kv_tok_thread_base=_kv_tok_thread_base, - prob_wr_thread_base=_prob_wr_thread_base, - pv_prob_read_base=_pv_prob_read_base, - sm_max_off=_sm_max_off, - sm_sum_off=_sm_sum_off, - sm_rd_max_offs=_sm_rd_max_offs, - sm_rd_sum_offs=_sm_rd_sum_offs, - sm_vmax_wr_off=_sm_vmax_wr_off, - sm_vmax_rd_offs=_sm_vmax_rd_offs, - c_w=c_w, - neg_inf=NEG_INF, - zero_f=ZERO_F, - cache_scale_vecs=cache_scale_vecs, - head_size=_HEAD, - qkhe_loop=_QKHELOOP, - vhe_loop=_VHELOOP, + head_dim=head_dim, ) - # Inner KV loop — one CTA processes one 256-token sub-tile across all - # 1024-token physical blocks. MTP groups loop is nested INSIDE so K/V - # load once per physical block and are reused; Q is hoisted out (loaded - # once per work item, kept in registers). + # Q stays in registers while K/V are reused across all MTP groups. def _unwrap(v): return v.ir_value() if hasattr(v, "ir_value") else v - c_ql = arith.constant(query_length, type=T.i32) - c_zero_i32 = arith.constant(0, type=T.i32) - c_bpp = arith.constant(_blocks_per_partition, type=T.i32) + c_ql = fx.Int32(query_length) + c_zero_i32 = fx.Int32(0) - # Output target: partial_qo_loc (work_info[1]) < 0 → write the final - # output directly; >= 0 → write a partial slot (combined later by - # pa_reduce_v1). The partial buffer reserves the first `query_length` - # rows (pa_reduce_v1 runs on partial_output[query_length:]), so the - # partial row base is `partial_idx + query_length`. qo_start - # (work_info[2]) is the final-output row base for direct works. + # Negative partial_idx writes final output; split rows reserve the first QL slots. _is_direct = partial_idx < c_zero_i32 _po_row_base = partial_idx + c_ql - # Loop over the work item's partitions: [kv_start, kv_end) cumulative - # partition indices → num_parts local 256-token partitions, in reverse - # (sink-prone partition 0 processed last for online-softmax stability). num_parts_in_work = kv_end - kv_start last_part_idx_val = num_parts_in_work - c_one - _loop_start_g = arith.index(0) - _loop_stop_g = fx.Index(arith.unwrap(num_parts_in_work)) - _loop_step_g = arith.index(1) _mtp_groups = math.ceil(query_length * query_group_size / 16) - # ── Small-block (16/64) physical-page gather helpers ── - # A 256-token partition spans `_blocks_per_partition` physical pages. - # Each warp loads its own K page(s); the per-warp page indices are - # staged to LDS so every warp can read all pages for the V load. - # (Only used when `_is_small_block`; for block_size >= 256 a partition - # is a 256-token sub-tile of a single physical page.) - _kpi_last = kv_page_end - c_one # last in-bounds page index for this seq - - def _meta_load_phys_clamped(elem_idx): - # Clamp to the sequence's page range so the last (partial) partition - # never reads past the flat kv_page_indices window (kpi_rsrc has no - # HW bounds check). Out-of-range lanes map to tokens >= context_len, - # which softmax masks to 0, so the clamped block's content is unused. - safe = arith.select(elem_idx < kv_page_end, elem_idx, _kpi_last) - return buffer_ops.buffer_load(kpi_rsrc, safe, vec_width=1, dtype=T.i32) - - def _meta_stage_phys(local_part): - page_base = kv_page_base + local_part * c_bpp - if const_expr(_block_size == 64): - return _meta_load_phys_clamped(page_base + warp_id) - wbase = page_base + warp_id * arith.constant(TLOOP, type=T.i32) - elems = [_meta_load_phys_clamped(wbase + arith.constant(td, type=T.i32)) for td in range(TLOOP)] - return fx.Vector.from_elements(elems, dtype=fx.Int32) - - def _meta_store_phys_to_lds(phys_vec): - if (lane16id | rowid) == c_zero_i32: - if const_expr(_block_size == 64): - fx.ptr_store( - fx.Vector.from_elements([phys_vec], dtype=fx.Int32), - bt_base + warp_id, - ) - else: - fx.ptr_store(phys_vec, bt_base + warp_id * arith.constant(TLOOP, type=T.i32)) - - def _meta_load_v_phys_from_lds(): - v_phys_blocks = [] - if const_expr(_block_size == 64): - phys_block_vec = fx.ptr_load(bt_base + (0), result_type=fx.Vector.make_type(VTLOOP, fx.Int32)) - for vt in range_constexpr(VTLOOP): - v_phys_blocks.append(phys_block_vec[vt]) - else: - for vt in range_constexpr(VTLOOP): - bt_lds_off = arith.constant(vt * TLOOP, type=T.i32) + rowid - v_phys_blocks.append( - fx.ptr_load(bt_base + (bt_lds_off), result_type=fx.Vector.make_type(1, fx.Int32))[0] - ) - return v_phys_blocks - - # ── Pre-load Q for every MTP group ONCE per work item. Each - # group's q_frags / qi / qhi / qscale stay in registers across - # the entire KV loop, so we pay the Q-load cost (global → LDS → - # registers) exactly once per work-item regardless of how many - # blocks the work item spans. q_frags_per_mtp = [] qi_per_mtp = [] qhi_per_mtp = [] @@ -1694,7 +1239,9 @@ def _meta_load_v_phys_from_lds(): if const_expr(_mtp_g > 0): gpu.barrier() mtp_prefetch = _prefetch_mtp_group_query( - q_rsrc, + q_tiles, + copy_q, + q_reg, batch_idx, kv_h, stride_q_seq, @@ -1704,7 +1251,6 @@ def _meta_load_v_phys_from_lds(): mtp_group_idx=_mtp_g, query_length=query_length, query_group_size=query_group_size, - query_load_is_bf16=query_load_is_bf16, q_lanes_per_head=_Q_LANES_PER_HEAD, ) _qi, _qhi, _qfrags, _qscale = _finish_mtp_group_q_fragments( @@ -1714,7 +1260,7 @@ def _meta_load_v_phys_from_lds(): lane16id, rowid, local_qhead_idx, - head_size=_HEAD, + head_dim=head_dim, qkhe_loop=_QKHELOOP, q_lanes_per_head=_Q_LANES_PER_HEAD, ) @@ -1724,134 +1270,88 @@ def _meta_load_v_phys_from_lds(): qscale_per_mtp.append(_qscale) gpu.barrier() - # MTP causal bound per group (depends only on qi, computed once). causal_bound_per_mtp = [ - context_len + arith.constant(1 - query_length, type=T.i32) + qi_per_mtp[_mtp_g] - for _mtp_g in range(_mtp_groups) + context_len + fx.Int32(1 - query_length) + qi_per_mtp[_mtp_g] for _mtp_g in range(_mtp_groups) ] - # ── K init: load the reverse-start (last) partition's K (loop-carried) ── local_last_part = local_part_start + last_part_idx_val - if const_expr(_is_small_block): - _first_phys_blocks = _meta_stage_phys(local_last_part) - k_flat0 = _pa_small_block_load_k_flat( - k_global_ptr, - kv_h, - c_kb, - c_kh, - lane16id, - rowid, - block_size=_block_size, - phys_blocks=_first_phys_blocks, - qkhe_loop=_QKHELOOP, - ) - scale_scalars0 = None - else: - _first_phys_block = buffer_ops.buffer_load( - kpi_rsrc, kv_page_base + udiv_const(local_last_part, _parts_per_block), vec_width=1, dtype=T.i32 - ) - _first_tile_tok = urem_const(local_last_part, _parts_per_block) * c_cps - first_k_base = _compute_block_base_dw_i64(_first_phys_block, c_kb, _k_head_off) - scale_scalars0 = _load_kv_scale_scalars(_first_tile_tok, _first_phys_block) - k_flat0 = _load_k_flat( - k_global_ptr, - first_k_base, - _first_tile_tok, - _k_tok_thread_base, - _c_tok_stride_dw, - _k_he_off_dw, - qkhe_loop=_QKHELOOP, - ) + first_phys_block = _copy_load( + kv_page_indices, + kv_page_base + udiv_const(local_last_part, parts_per_block), + copy_i32, + i32_reg, + )[0] + first_tile_tok = urem_const(local_last_part, parts_per_block) * c_cps + first_k_base = _compute_block_base_dw_i64(first_phys_block, stride_k_block, _k_head_off) + scale_scalars0 = _load_kv_scale_scalars(first_tile_tok, first_phys_block) + k_flat0 = _load_k_flat( + k_global_ptr, + global_copy_16b, + global_reg_16b, + first_k_base, + first_tile_tok, + _k_tok_thread_base, + _c_tok_stride_dw, + _k_he_off_dw, + qkhe_loop=_QKHELOOP, + ) - # Multi-MTP state packing: (rmax, rsum, outs...) per MTP group, - # + _N_K_h K values, + 2 scale scalars (per_token_kv only). state_width = 2 + _VHELOOP - def _pack_states_kv(states, k_flat, scale_scalars=None): - flat = [] - for st in states: - rmax, rsum = st[0], st[1] - outs = [st[2 + vhe] for vhe in range_constexpr(_VHELOOP)] - flat.extend([_unwrap(rmax), _unwrap(rsum)]) - flat.extend(_unwrap(out) for out in outs) + def _pack_states_kv(states, k_flat, scale_scalars): + flat = [_unwrap(value) for state in states for value in state] flat.extend(_unwrap(v) for v in k_flat) - if const_expr(cache_scale_vecs and per_token_kv): - flat.extend(_unwrap(v) for v in scale_scalars) + flat.extend(_unwrap(v) for v in scale_scalars) return flat def _unpack_states_kv(flat): - base = state_width * _mtp_groups - states = [tuple(flat[state_width * i + j] for j in range(state_width)) for i in range(_mtp_groups)] - k_flat = list(flat[base : base + _N_K_h]) - if const_expr(cache_scale_vecs and per_token_kv): - scale_scalars = tuple(flat[base + _N_K_h : base + _N_K_h + 2]) - else: - scale_scalars = None - return states, k_flat, scale_scalars + state_end = state_width * _mtp_groups + k_end = state_end + _N_K_h + states = [tuple(flat[state_width * i : state_width * (i + 1)]) for i in range(_mtp_groups)] + return states, list(flat[state_end:k_end]), tuple(flat[k_end:]) init_states = [ - tuple([NEG_INF, ZERO_F] + [arith.constant_vector(0.0, T.f32x4) for _ in range_constexpr(_VHELOOP)]) + tuple([NEG_INF, ZERO_F] + [fx.Vector.filled(4, 0.0, fx.Float32) for _ in range_constexpr(_VHELOOP)]) for _ in range(_mtp_groups) ] - # KV outer loop over physical blocks (MTP processing nested inside). for ib, state in range( - _loop_start_g, - _loop_stop_g, - _loop_step_g, + fx.Int32(0), + num_parts_in_work, + fx.Int32(1), init=_pack_states_kv(init_states, k_flat0, scale_scalars0), ): cur_states, k_flat, scale_scalars = _unpack_states_kv(state) - # Reverse iteration: scf.for walks ib forward (0..N-1); remap to - # the local partition index lp = N-1..0 so the sink-prone first - # partition is processed last. - rel_part = last_part_idx_val - arith.index_cast(T.i32, ib) + # Process partition zero last for online-softmax stability. + rel_part = last_part_idx_val - ib lp = local_part_start + rel_part next_rel = rel_part - c_one - next_rel_clamped = arith.select(next_rel >= c_zero_i32, next_rel, c_zero_i32) + next_rel_clamped = (next_rel >= c_zero_i32).select(next_rel, c_zero_i32) next_lp = local_part_start + next_rel_clamped k_ops = unflatten_k(k_flat, qkhe_loop=_QKHELOOP) partition_start = lp * c_cps # within-sequence token offset of this 256-tile - # Load V (and per-token scales if applicable) ONCE per partition; - # reused across all MTP groups below. - if const_expr(_is_small_block): - _meta_store_phys_to_lds(_meta_stage_phys(lp)) - gpu.barrier() - v_ops = _pa_small_block_load_v_trans( - v_global_ptr, - kv_h, - c_vb, - c_vh, - warp_id, - lane16id, - rowid, - _meta_load_v_phys_from_lds(), - block_size=_block_size, - head_size=_HEAD, - vhe_loop=_VHELOOP, + phys_block = _copy_load( + kv_page_indices, + kv_page_base + udiv_const(lp, parts_per_block), + copy_i32, + i32_reg, + )[0] + tile_token_offset = urem_const(lp, parts_per_block) * c_cps + v_base = _compute_block_base_dw_i64(phys_block, stride_v_block, _v_head_off) + if const_expr(per_token_kv): + v_ops, k_scale_vecs, v_scale_vecs = _load_v_and_scales( + v_base, + tile_token_offset, + scale_scalars, ) else: - phys_block = buffer_ops.buffer_load( - kpi_rsrc, kv_page_base + udiv_const(lp, _parts_per_block), vec_width=1, dtype=T.i32 + v_ops = _load_v_and_scales( + v_base, + tile_token_offset, + scale_scalars, ) - tile_token_offset = urem_const(lp, _parts_per_block) * c_cps - v_base = _compute_block_base_dw_i64(phys_block, c_vb, _v_head_off) - if const_expr(cache_scale_vecs and per_token_kv): - v_ops, k_scale_vecs, v_scale_vecs = _load_v_and_scales( - v_base, - tile_token_offset, - phys_block=phys_block, - preloaded_scale_scalars=scale_scalars, - ) - else: - v_ops = _load_v_and_scales( - v_base, - tile_token_offset, - phys_block=phys_block, - preloaded_scale_scalars=scale_scalars, - ) new_states = [] for _mtp_g in range_constexpr(_mtp_groups): if const_expr(_mtp_g > 0): @@ -1860,7 +1360,7 @@ def _unpack_states_kv(flat): rmax, rsum = state[0], state[1] outs = [state[2 + vhe] for vhe in range_constexpr(_VHELOOP)] - if const_expr(cache_scale_vecs and per_token_kv): + if const_expr(per_token_kv): d_out, v_scales = _qk_and_intra_softmax( k_ops, partition_start, @@ -1879,10 +1379,7 @@ def _unpack_states_kv(flat): ) v_scales = None - # Bugfix: per_token_kv path needs v_max staged to LDS so - # _cross_warp_softmax_and_prob_pack can read it for - # norm_factor. Without this write the read sees stale/ - # uninitialized LDS and produces NaN. + # Stage vmax before the cross-wave normalization reads it. if const_expr(per_token_kv): _store_vmax_warp(partition_start, seq_end=context_len, v_scale_vecs=v_scales) @@ -1894,53 +1391,30 @@ def _unpack_states_kv(flat): outs = _pv_mfma(v_ops, outs, v_correction) new_states.append(tuple([rmax, rsum] + outs)) - # Prefetch next partition's K (once per iter, after all MTP groups) - if const_expr(_is_small_block): - k_next_flat = _pa_small_block_load_k_flat( - k_global_ptr, - kv_h, - c_kb, - c_kh, - lane16id, - rowid, - block_size=_block_size, - phys_blocks=_meta_stage_phys(next_lp), - qkhe_loop=_QKHELOOP, - ) - next_scale_scalars = None - else: - next_phys_block = buffer_ops.buffer_load( - kpi_rsrc, kv_page_base + udiv_const(next_lp, _parts_per_block), vec_width=1, dtype=T.i32 - ) - next_tile_tok = urem_const(next_lp, _parts_per_block) * c_cps - next_k_base = _compute_block_base_dw_i64(next_phys_block, c_kb, _k_head_off) - next_scale_scalars = _load_kv_scale_scalars(next_tile_tok, next_phys_block) - k_next_flat = _load_k_flat( - k_global_ptr, - next_k_base, - next_tile_tok, - _k_tok_thread_base, - _c_tok_stride_dw, - _k_he_off_dw, - qkhe_loop=_QKHELOOP, - ) + next_phys_block = _copy_load( + kv_page_indices, + kv_page_base + udiv_const(next_lp, parts_per_block), + copy_i32, + i32_reg, + )[0] + next_tile_tok = urem_const(next_lp, parts_per_block) * c_cps + next_k_base = _compute_block_base_dw_i64(next_phys_block, stride_k_block, _k_head_off) + next_scale_scalars = _load_kv_scale_scalars(next_tile_tok, next_phys_block) + k_next_flat = _load_k_flat( + k_global_ptr, + global_copy_16b, + global_reg_16b, + next_k_base, + next_tile_tok, + _k_tok_thread_base, + _c_tok_stride_dw, + _k_he_off_dw, + qkhe_loop=_QKHELOOP, + ) results = yield _pack_states_kv(new_states, k_next_flat, next_scale_scalars) - # ── Normalize + store one slot per MTP group ── - # partial_qo_loc (work_info[1]) < 0 → write the fully-normalized output - # directly to `out` at row qo_start+qi; >= 0 → write a partial slot - # (+LSE) at row partial_idx+query_length+qi for pa_reduce_v1. final_states, _, _ = _unpack_states_kv(results) - from flydsl._mlir.dialects import math as _mlir_math - - def _store_out_vec(vec_f32x4, elem_off): - if const_expr(output_dtype_str == "f32"): - buffer_ops.buffer_store(vec_f32x4, out_rsrc, elem_off) - elif const_expr(output_dtype_str == "f16"): - buffer_ops.buffer_store(fx.Vector(vec_f32x4).to(fx.Float16), out_rsrc, elem_off) - else: - buffer_ops.buffer_store(fx.Vector(vec_f32x4).to(fx.BFloat16), out_rsrc, elem_off) for _mtp_g in range_constexpr(_mtp_groups): final_state = final_states[_mtp_g] @@ -1952,42 +1426,34 @@ def _store_out_vec(vec_f32x4, elem_off): outelems_norm = _normalize_pa_output(running_sum, outs, ZERO_F) qi_val_mg = qi_per_mtp[_mtp_g] qhi_pos_mg = qhi_per_mtp[_mtp_g] - qhead = kv_h * arith.constant(query_group_size, type=T.i32) + qhi_pos_mg + qhead = kv_h * fx.Int32(query_group_size) + qhi_pos_mg if _is_direct: out_row = qo_start + qi_val_mg for vhe in range_constexpr(_VHELOOP): - hs_base = ( - arith.constant(vhe * NUM_WARPS * MFMA_N, type=T.i32) - + warp_id * arith.constant(MFMA_N, type=T.i32) - + rowid * arith.constant(4, type=T.i32) - ) + hs_base = fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + rowid * fx.Int32(4) out_off = out_row * stride_out_seq + qhead * stride_out_head + hs_base - _store_out_vec(outelems_norm[vhe], out_off) + fx.memref_store_vec(outelems_norm[vhe].to(_OUTPUT_DTYPE), out_reg) + fx.copy(copy_out, out_reg, fx.slice(out_tiles, (None, out_off // fx.Int32(4)))) else: _po_row = _po_row_base + qi_val_mg for vhe in range_constexpr(_VHELOOP): - hs_base = ( - arith.constant(vhe * NUM_WARPS * MFMA_N, type=T.i32) - + warp_id * arith.constant(MFMA_N, type=T.i32) - + rowid * arith.constant(4, type=T.i32) - ) - po_off = _po_row * stride_po_ql + qhead * arith.constant(_HEAD, type=T.i32) + hs_base - buffer_ops.buffer_store( - outelems_norm[vhe], po_rsrc, po_off * arith.constant(4, type=T.i32), offset_is_bytes=True + hs_base = fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + rowid * fx.Int32(4) + po_off = _po_row * stride_po_ql + qhead * fx.Int32(head_dim) + hs_base + fx.memref_store_vec(outelems_norm[vhe], partial_out_reg) + fx.copy( + copy_f32x4, + partial_out_reg, + fx.slice(partial_out_tiles, (None, po_off // fx.Int32(4))), ) - # LSE (split partials only) - safe_sum_lse = arith.select(running_sum > ZERO_F, running_sum, arith.constant(1.0, type=T.f32)) - log_sum = _mlir_math.log(safe_sum_lse, fastmath=arith.FastMathFlags.fast) + safe_sum_lse = (running_sum > ZERO_F).select(running_sum, fx.Float32(1.0)) + log_sum = fmath.log(safe_sum_lse, fastmath=arith.FastMathFlags.fast) lse_val = running_max + log_sum pl_off = _po_row * stride_pl_ql + qhead - lse_as_i32 = arith.bitcast(T.i32, arith.unwrap(lse_val)) - buffer_ops.buffer_store( - lse_as_i32, pl_rsrc, pl_off * arith.constant(4, type=T.i32), offset_is_bytes=True - ) + fx.memref_store_vec(fx.Vector.from_elements([lse_val], dtype=fx.Float32), partial_lse_reg) + fx.copy(copy_f32, partial_lse_reg, fx.slice(partial_lse, (None, pl_off))) - # ── @flyc.jit launch wrapper ───────────────────────────────────── @flyc.jit def launch_pa_decode_metadata( out: fx.Int64, @@ -2050,43 +1516,26 @@ def launch_pa_decode_metadata( s_ks_head, s_po_ql, s_pl_ql, - # value_attrs=_mfma_agpr_value_attrs(), ).launch(grid=(num_sm, 1, 1), block=(BLOCK_THREADS, 1, 1), stream=stream) - # launch_pa_decode_metadata.compile_hints["llvm_options"] = PA_MFMA_AGPR_LLVM_OPTIONS - return { "launch": launch_pa_decode_metadata, "kernel": pa_decode_metadata_kenrel, } -# Deterministic split-partition reduce; drop-in replacement for the racy aiter -# ``pa_reduce_v1`` / ``mla_reduce_v1`` on the PS decode path (root cause of the -# flaky ``test_pa`` NaN). Each output element is combined by exactly one thread, -# serially via online softmax (out = Σ_s w_s·o_s / Σ_s w_s, w_s = exp(lse_s − max)). -# The per-split LSE is a per-(row, head) scalar shared by all head_size lanes, so -# the combine is thread-independent — no atomics/cross-thread reduction → bit-exact. -# -# Inputs (from compile_pa_decode_metadata / get_pa_metadata; caller passes the -# ``[query_length:]`` slices, mirroring the old pa_reduce_v1 call): -# partial_output f32 : row [num_query_heads, head_size], stride stride_po_row; -# each split already normalized by its local denom. -# partial_lse f32 : row [num_query_heads], stride stride_pl_row; = max + log(sum). -# reduce_indptr [batch+1]: group g splits = slots [indptr[g], indptr[g+1]). -# reduce_final_map [batch,2]: group g → final output row range. -# reduce_partial_map: slot → partial row base (in the sliced buffer). -# Direct (non-split) outputs have empty groups (indptr delta 0) and are skipped. +# One thread per output element serially combines splits with online softmax. @functools.lru_cache(maxsize=64) def compile_pa_metadata_reduce( *, query_length: int, num_query_heads: int, - head_size: int, + head_dim: int, output_dtype_str: str, ): - block_threads = head_size - assert 0 < block_threads <= 1024, "head_size must fit in one workgroup" + block_threads = head_dim + assert 0 < block_threads <= 1024, "head_dim must fit in one workgroup" + output_dtype = _PA_OUTPUT_DTYPES[output_dtype_str] @flyc.kernel(known_block_size=(block_threads, 1, 1)) def pa_metadata_reduce_kernel( @@ -2098,7 +1547,7 @@ def pa_metadata_reduce_kernel( reduce_partial_map_ptr: fx.Tensor, stride_out_seq: Int32, stride_out_head: Int32, - stride_po_row: Int32, # num_query_heads * head_size + stride_po_row: Int32, # num_query_heads * head_dim stride_pl_row: Int32, # num_query_heads ): tid = fx.Int32(gpu.thread_id("x")) @@ -2106,64 +1555,71 @@ def pa_metadata_reduce_kernel( qhead = fx.Int32(gpu.block_id("y")) qi = fx.Int32(gpu.block_id("z")) - out_rsrc = buffer_ops.create_buffer_resource(final_output_ptr, max_size=True) - po_rsrc = buffer_ops.create_buffer_resource(partial_output_ptr, max_size=True) - pl_rsrc = buffer_ops.create_buffer_resource(partial_lse_ptr, max_size=True) - rip_rsrc = buffer_ops.create_buffer_resource(reduce_indptr_ptr, max_size=True) - rfm_rsrc = buffer_ops.create_buffer_resource(reduce_final_map_ptr, max_size=True) - rpm_rsrc = buffer_ops.create_buffer_resource(reduce_partial_map_ptr, max_size=True) - - c_zero_f = fx.Float32(0.0) - c_one_f = fx.Float32(1.0) - c_neg_inf = fx.Float32(float("-inf")) - c_log2e = fx.Float32(LOG2E) - c_head = fx.Int32(head_size) + scalar_layout = fx.make_layout(1, 1) + flat_layout = fx.make_layout(_FLAT_BUFFER_ELEMENTS, 1) + + def _flat_divide(tensor): + buffer_tensor = fx.rocdl.make_buffer_tensor(tensor) + flat = fx.make_view(fx.get_iter(buffer_tensor), flat_layout) + return fx.logical_divide(flat, scalar_layout) + + final_output = _flat_divide(final_output_ptr) + partial_output = _flat_divide(partial_output_ptr) + partial_lse = _flat_divide(partial_lse_ptr) + reduce_indptr = _flat_divide(reduce_indptr_ptr) + reduce_final_map = _flat_divide(reduce_final_map_ptr) + reduce_partial_map = _flat_divide(reduce_partial_map_ptr) + + copy_i32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Int32) + copy_f32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Float32) + copy_output = fx.make_copy_atom( + fx.rocdl.BufferCopy32b() if output_dtype.width == 32 else fx.rocdl.BufferCopy16b(), + output_dtype, + ) + reg_i32 = fx.make_rmem_tensor(1, fx.Int32) + reg_f32 = fx.make_rmem_tensor(1, fx.Float32) + reg_output = fx.make_rmem_tensor(1, output_dtype) - lo = buffer_ops.buffer_load(rip_rsrc, g, vec_width=1, dtype=T.i32) - hi = buffer_ops.buffer_load(rip_rsrc, g + fx.Int32(1), vec_width=1, dtype=T.i32) + lo = _copy_load(reduce_indptr, g, copy_i32, reg_i32)[0] + hi = _copy_load(reduce_indptr, g + fx.Int32(1), copy_i32, reg_i32)[0] - # Empty group (direct / padding) → nothing to combine; leave row untouched. if hi > lo: - qo_start = buffer_ops.buffer_load(rfm_rsrc, g * fx.Int32(2), vec_width=1, dtype=T.i32) + qo_start = _copy_load(reduce_final_map, g * fx.Int32(2), copy_i32, reg_i32)[0] out_row = qo_start + qi - # Online-softmax combine over the group's splits (dynamic count). - # State = (running_max, running_denom, running_acc), per-thread. - for _s, st in fx.range( - fx.Index(arith.unwrap(lo)), - fx.Index(arith.unwrap(hi)), - 1, - init=[c_neg_inf, c_zero_f, c_zero_f], + for slot, st in range( + lo, + hi, + fx.Int32(1), + init=[fx.Float32(float("-inf")), fx.Float32(0.0), fx.Float32(0.0)], ): m = fx.Float32(st[0]) denom = fx.Float32(st[1]) acc = fx.Float32(st[2]) - slot = fx.Int32(_s) - prow = buffer_ops.buffer_load(rpm_rsrc, slot, vec_width=1, dtype=T.i32) + qi - lse = buffer_ops.buffer_load(pl_rsrc, prow * stride_pl_row + qhead, vec_width=1, dtype=T.f32) - v = buffer_ops.buffer_load( - po_rsrc, prow * stride_po_row + qhead * c_head + tid, vec_width=1, dtype=T.f32 - ) + prow = _copy_load(reduce_partial_map, slot, copy_i32, reg_i32)[0] + qi + lse = _copy_load(partial_lse, prow * stride_pl_row + qhead, copy_f32, reg_f32)[0] + v = _copy_load( + partial_output, + prow * stride_po_row + qhead * fx.Int32(head_dim) + tid, + copy_f32, + reg_f32, + )[0] m_new = m.maximumf(lse) - scale_old = exp2_f32_fast((m - m_new) * c_log2e) - w = exp2_f32_fast((lse - m_new) * c_log2e) + scale_old = exp2_f32_fast((m - m_new) * fx.Float32(LOG2E)) + w = exp2_f32_fast((lse - m_new) * fx.Float32(LOG2E)) denom_new = denom * scale_old + w acc_new = acc * scale_old + v * w results = yield [m_new, denom_new, acc_new] denom_f = fx.Float32(results[1]) acc_f = fx.Float32(results[2]) - safe_denom = arith.select(denom_f > c_zero_f, denom_f, c_one_f) + safe_denom = (denom_f > fx.Float32(0.0)).select(denom_f, fx.Float32(1.0)) out_acc = acc_f * rcp_f32(safe_denom) out_off = out_row * stride_out_seq + qhead * stride_out_head + tid - if const_expr(output_dtype_str == "f32"): - out_val = out_acc - elif const_expr(output_dtype_str == "f16"): - out_val = out_acc.to(fx.Float16) - else: - out_val = out_acc.to(fx.BFloat16) - buffer_ops.buffer_store(out_val, out_rsrc, out_off) + output_val = fx.Float32(out_acc).to(output_dtype) + fx.memref_store_vec(fx.Vector.from_elements([output_val], dtype=output_dtype), reg_output) + fx.copy(copy_output, reg_output, fx.slice(final_output, (None, out_off))) @flyc.jit def launch_pa_metadata_reduce( @@ -2217,7 +1673,7 @@ def pa_metadata_reduce( max_seqlen_q: int, final_output: torch.Tensor, num_query_heads: int, - head_size: int, + head_dim: int, stream=None, ) -> None: """Deterministic drop-in replacement for ``aiter.pa_reduce_v1`` on the PS path. @@ -2228,13 +1684,13 @@ def pa_metadata_reduce( in-kernel, so passing ``reduce_indptr.numel() - 1`` groups needs no host sync. """ num_groups = reduce_indptr.numel() - 1 - stride_po_row = num_query_heads * head_size + stride_po_row = num_query_heads * head_dim stride_pl_row = num_query_heads out_dtype_str = _PA_PS_REDUCE_DTYPE_STR[final_output.dtype] compiled = compile_pa_metadata_reduce( query_length=int(max_seqlen_q), num_query_heads=int(num_query_heads), - head_size=int(head_size), + head_dim=int(head_dim), output_dtype_str=out_dtype_str, ) _run_compiled( @@ -2250,5 +1706,5 @@ def pa_metadata_reduce( stride_po_row, stride_pl_row, num_groups, - stream if stream is not None else fx.Stream(None), + stream, ) diff --git a/tests/kernels/test_pa.py b/tests/kernels/test_pa.py index da100c3e6..04179b328 100644 --- a/tests/kernels/test_pa.py +++ b/tests/kernels/test_pa.py @@ -1116,6 +1116,41 @@ def test_normal_accuracy( ) +@pytest.mark.parametrize(("query_length", "quant_mode"), [(1, "per_tensor"), (4, "per_token")]) +def test_metadata_accuracy(query_length: int, quant_mode: str) -> None: + """Exercise the block-1024 persistent worklist decode and split reducer.""" + run_pa_decode_ps_test( + context_length=1027, + batch_size=3, + num_heads=(16, 1), + head_size=128, + block_size=1024, + compute_type=dtypes.d_dtypes["fp8"], + query_length=query_length, + quant_mode=quant_mode, + context_partition_size=256, + trans_v=True, + kv_varlen=False, + sliding_window=0, + ) + + +@pytest.mark.parametrize("block_size", [16, 64, 256, 2048]) +def test_metadata_rejects_non_1024_block_size(block_size: int) -> None: + from kernels.attention.pa_metadata import compile_pa_decode_metadata + + with pytest.raises(ValueError, match="only supports block_size=1024"): + compile_pa_decode_metadata(block_size=block_size) + + +@pytest.mark.parametrize("query_input_dtype", ["fp8", "f32"]) +def test_metadata_rejects_unsupported_query_dtype(query_input_dtype: str) -> None: + from kernels.attention.pa_metadata import compile_pa_decode_metadata + + with pytest.raises(ValueError, match="only supports bf16/f16 queries"): + compile_pa_decode_metadata(query_input_dtype=query_input_dtype) + + @pytest.mark.parametrize("compute_type", ["fp8"]) @pytest.mark.parametrize("context_partition_size", [256]) @pytest.mark.parametrize("head_size", [128]) From 91699ec7972618ea4761c8d7db58a8d43845498c Mon Sep 17 00:00:00 2001 From: fsx950223 Date: Fri, 24 Jul 2026 08:37:19 +0000 Subject: [PATCH 2/7] refactor(pa): Simplify metadata kernel helpers Signed-off-by: fsx950223 Co-authored-by: Cursor --- kernels/attention/pa_metadata.py | 132 +++++++++++-------------------- 1 file changed, 44 insertions(+), 88 deletions(-) diff --git a/kernels/attention/pa_metadata.py b/kernels/attention/pa_metadata.py index 32bfec696..36c9ed852 100644 --- a/kernels/attention/pa_metadata.py +++ b/kernels/attention/pa_metadata.py @@ -78,17 +78,15 @@ def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = if num_cu is None: gpu = torch.cuda.current_device() num_cu = torch.cuda.get_device_properties(gpu).multi_processor_count - cu_num = num_cu - tile_cnt = batch_size - max_work = (tile_cnt + cu_num - 1) * num_head_k - max_split_tiles = min(batch_size + cu_num - 1, (cu_num - 1) * 2) + max_work = (batch_size + num_cu - 1) * num_head_k + max_split_tiles = min(batch_size + num_cu - 1, (num_cu - 1) * 2) return ( - ((cu_num + 1,), torch.int32), # work_indptr + ((num_cu + 1,), torch.int32), # work_indptr ((max_work, _WORK_INFO_FIELDS), torch.int32), # work_info_set - ((tile_cnt + 1,), torch.int32), # reduce_indptr - ((tile_cnt, 2), torch.int32), # reduce_final_map + ((batch_size + 1,), torch.int32), # reduce_indptr + ((batch_size, 2), torch.int32), # reduce_final_map ((max_split_tiles,), torch.int32), # reduce_partial_map ) @@ -229,9 +227,9 @@ def _finish_q_fragments( local_qhead_idx, *, head_dim: int, - qkhe_loop: int, - q_lanes_per_head: int, ): + qkhe_loop = head_dim // QKHE_PER_FETCH + q_lanes_per_head = head_dim // Q_ELEMS_PER_LANE lds_q_base = local_qhead_idx * fx.Int32(head_dim // 4) + lane16id * 2 abs_mask = fx.Vector.filled(4, 0x7FFFFFFF, fx.Int32) @@ -240,28 +238,21 @@ def _finish_q_fragments( for q_src in q_chunks: q_f32 = fx.Vector(q_src).to(fx.Float32) q_f32_chunks.append(q_f32) - q_i32 = q_f32.bitcast(fx.Int32) - q_abs_i32 = q_i32 & abs_mask - q_abs = q_abs_i32.bitcast(fx.Float32) - chunk_max = q_abs.reduce("max") - local_max = fx.maxnumf(local_max, chunk_max) + q_abs = (q_f32.bitcast(fx.Int32) & abs_mask).bitcast(fx.Float32) + local_max = fx.maxnumf(local_max, q_abs.reduce("max")) for sh in [8, 4, 2, 1]: local_max = fx.maxnumf(local_max, dpp_utils.dpp_xor_f32(local_max, sh)) - query_scale_lane = fx.Float32( - arith.select( - local_max > fx.Float32(0.0), - local_max * fx.Float32(1.0 / FP8_MAX).ir_value(), - fx.Float32(1.0), - ) + query_scale_lane = (local_max > fx.Float32(0.0)).select( + local_max * fx.Float32(1.0 / FP8_MAX), + fx.Float32(1.0), ) - inv_query_scale = rcp_f32(query_scale_lane) + inv_query_scale = fx.Float32(rcp_f32(query_scale_lane)) q_words = [] for q_f32 in q_f32_chunks: p = q_f32 * inv_query_scale lo = rocdl.cvt_pk_fp8_f32(T.i32, p[0], p[1], fx.Int32(0), False) q_words.append(rocdl.cvt_pk_fp8_f32(T.i32, p[2], p[3], lo, True)) - q_w0, q_w1 = q_words if lane16id == fx.Int32(0): fx.ptr_store( @@ -269,25 +260,25 @@ def _finish_q_fragments( softmax_base + local_qhead_idx, ) - v01 = fx.Vector.from_elements([q_w0, q_w1], dtype=fx.Int32) + q_vec = fx.Vector.from_elements(q_words, dtype=fx.Int32) if const_expr(q_lanes_per_head < MFMA_N): if lane16id < fx.Int32(q_lanes_per_head): - fx.ptr_store(v01, logits_base + lds_q_base) + fx.ptr_store(q_vec, logits_base + lds_q_base) else: - fx.ptr_store(v01, logits_base + lds_q_base) + fx.ptr_store(q_vec, logits_base + lds_q_base) q_frags = [] gpu.barrier() - query_scale_lane = fx.ptr_load(softmax_base + (lane16id), result_type=fx.Vector.make_type(1, fx.Float32))[ - 0 - ].ir_value() + query_scale_lane = fx.ptr_load(softmax_base + lane16id, result_type=fx.Vector.make_type(1, fx.Float32))[0] for qkhe in range_constexpr(qkhe_loop): for qkr in range_constexpr(2): lds_rd = lane16id * fx.Int32(head_dim // 8) + fx.Int32(qkhe * 8) + rowid * fx.Int32(2) + fx.Int32(qkr) - q_v1 = fx.ptr_load( - fx.recast_iter(fx.Int64, logits_base) + (lds_rd), result_type=fx.Vector.make_type(1, fx.Int64) + q_frags.append( + fx.ptr_load( + fx.recast_iter(fx.Int64, logits_base) + lds_rd, + result_type=fx.Vector.make_type(1, fx.Int64), + )[0] ) - q_frags.append(q_v1[0]) return q_frags, query_scale_lane @@ -327,33 +318,6 @@ def _prefetch_mtp_group_query( return qi_val, qhi_pos, q_chunks -def _finish_mtp_group_q_fragments( - logits_base, - softmax_base, - mtp_prefetch, - lane16id, - rowid, - local_qhead_idx, - *, - head_dim: int, - qkhe_loop: int, - q_lanes_per_head: int, -): - qi_val, qhi_pos, q_chunks = mtp_prefetch - q_frags, query_scale_lane = _finish_q_fragments( - logits_base, - softmax_base, - q_chunks, - lane16id, - rowid, - local_qhead_idx, - head_dim=head_dim, - qkhe_loop=qkhe_loop, - q_lanes_per_head=q_lanes_per_head, - ) - return qi_val, qhi_pos, q_frags, query_scale_lane - - def _normalize_pa_output(running_sum, outs, zero_f): one_f = fx.Float32(1.0).ir_value() safe_sum = arith.select(running_sum > zero_f, running_sum, one_f) @@ -379,7 +343,7 @@ def _make_pa_phase_helpers( scale_base, stride_ks_block, stride_ks_head, - softmax_scale_base, + softmax_scale, k_scale_val, v_scale_val, warp_id, @@ -531,7 +495,7 @@ def _qk_and_intra_softmax( if const_expr(per_token_kv): k_scale_vecs, v_scale_vecs = preloaded_scales - query_scale = fx.Float32(query_scale_lane * softmax_scale_base) + query_scale = fx.Float32(query_scale_lane * softmax_scale) d_out = [] for td in range_constexpr(TLOOP): acc = fx.Vector.filled(4, 0.0, fx.Float32) @@ -1031,7 +995,7 @@ def compile_pa_decode_metadata( _OUTPUT_DTYPE = _PA_OUTPUT_DTYPES[output_dtype_str] if softmax_scale is None: softmax_scale = 1.0 / (head_dim**0.5) - _softmax_scale = float(softmax_scale) + softmax_scale = float(softmax_scale) parts_per_block = KV_BLOCK_SIZE // KV_COMPUTE_BLOCK # Per-token mode adds one cross-warp vmax slot per lane. @@ -1142,12 +1106,6 @@ def _divide_addr(addr, dtype, width): if const_expr(per_token_kv): scale_base = fx.recast_iter(fx.Float32, fx.add_offset(lds.buf.ptr, fx.Int32(scale_off // 4))) - _softmax_scale_const = fx.Float32(_softmax_scale) - NEG_INF = fx.Float32(float("-inf")) - ZERO_F = fx.Float32(0.0) - c_cps = fx.Int32(KV_COMPUTE_BLOCK) # 256-token partition - c_one = fx.Int32(1) - local_qhead_idx = warp_id * fx.Int32(4) + rowid _k_tok_thread_base, _c_tok_stride_dw, _k_he_off_dw = _build_pa_k_thread_invariants( warp_id, @@ -1206,7 +1164,7 @@ def _divide_addr(addr, dtype, width): scale_base=scale_base, stride_ks_block=stride_ks_block, stride_ks_head=stride_ks_head, - softmax_scale_base=_softmax_scale_const, + softmax_scale=fx.Float32(softmax_scale), k_scale_val=k_scale_val, v_scale_val=v_scale_val, warp_id=warp_id, @@ -1219,15 +1177,12 @@ def _divide_addr(addr, dtype, width): def _unwrap(v): return v.ir_value() if hasattr(v, "ir_value") else v - c_ql = fx.Int32(query_length) - c_zero_i32 = fx.Int32(0) - # Negative partial_idx writes final output; split rows reserve the first QL slots. - _is_direct = partial_idx < c_zero_i32 - _po_row_base = partial_idx + c_ql + _is_direct = partial_idx < fx.Int32(0) + _po_row_base = partial_idx + fx.Int32(query_length) num_parts_in_work = kv_end - kv_start - last_part_idx_val = num_parts_in_work - c_one + last_part_idx_val = num_parts_in_work - fx.Int32(1) _mtp_groups = math.ceil(query_length * query_group_size / 16) @@ -1238,7 +1193,7 @@ def _unwrap(v): for _mtp_g in range_constexpr(_mtp_groups): if const_expr(_mtp_g > 0): gpu.barrier() - mtp_prefetch = _prefetch_mtp_group_query( + _qi, _qhi, q_chunks = _prefetch_mtp_group_query( q_tiles, copy_q, q_reg, @@ -1253,16 +1208,14 @@ def _unwrap(v): query_group_size=query_group_size, q_lanes_per_head=_Q_LANES_PER_HEAD, ) - _qi, _qhi, _qfrags, _qscale = _finish_mtp_group_q_fragments( + _qfrags, _qscale = _finish_q_fragments( logits_base, softmax_base, - mtp_prefetch, + q_chunks, lane16id, rowid, local_qhead_idx, head_dim=head_dim, - qkhe_loop=_QKHELOOP, - q_lanes_per_head=_Q_LANES_PER_HEAD, ) qi_per_mtp.append(_qi) qhi_per_mtp.append(_qhi) @@ -1281,7 +1234,7 @@ def _unwrap(v): copy_i32, i32_reg, )[0] - first_tile_tok = urem_const(local_last_part, parts_per_block) * c_cps + first_tile_tok = urem_const(local_last_part, parts_per_block) * fx.Int32(KV_COMPUTE_BLOCK) first_k_base = _compute_block_base_dw_i64(first_phys_block, stride_k_block, _k_head_off) scale_scalars0 = _load_kv_scale_scalars(first_tile_tok, first_phys_block) k_flat0 = _load_k_flat( @@ -1311,7 +1264,10 @@ def _unpack_states_kv(flat): return states, list(flat[state_end:k_end]), tuple(flat[k_end:]) init_states = [ - tuple([NEG_INF, ZERO_F] + [fx.Vector.filled(4, 0.0, fx.Float32) for _ in range_constexpr(_VHELOOP)]) + tuple( + [fx.Float32(float("-inf")), fx.Float32(0.0)] + + [fx.Vector.filled(4, 0.0, fx.Float32) for _ in range_constexpr(_VHELOOP)] + ) for _ in range(_mtp_groups) ] @@ -1325,12 +1281,12 @@ def _unpack_states_kv(flat): # Process partition zero last for online-softmax stability. rel_part = last_part_idx_val - ib lp = local_part_start + rel_part - next_rel = rel_part - c_one - next_rel_clamped = (next_rel >= c_zero_i32).select(next_rel, c_zero_i32) + next_rel = rel_part - fx.Int32(1) + next_rel_clamped = (next_rel >= fx.Int32(0)).select(next_rel, fx.Int32(0)) next_lp = local_part_start + next_rel_clamped k_ops = unflatten_k(k_flat, qkhe_loop=_QKHELOOP) - partition_start = lp * c_cps # within-sequence token offset of this 256-tile + partition_start = lp * fx.Int32(KV_COMPUTE_BLOCK) phys_block = _copy_load( kv_page_indices, @@ -1338,7 +1294,7 @@ def _unpack_states_kv(flat): copy_i32, i32_reg, )[0] - tile_token_offset = urem_const(lp, parts_per_block) * c_cps + tile_token_offset = urem_const(lp, parts_per_block) * fx.Int32(KV_COMPUTE_BLOCK) v_base = _compute_block_base_dw_i64(phys_block, stride_v_block, _v_head_off) if const_expr(per_token_kv): v_ops, k_scale_vecs, v_scale_vecs = _load_v_and_scales( @@ -1397,7 +1353,7 @@ def _unpack_states_kv(flat): copy_i32, i32_reg, )[0] - next_tile_tok = urem_const(next_lp, parts_per_block) * c_cps + next_tile_tok = urem_const(next_lp, parts_per_block) * fx.Int32(KV_COMPUTE_BLOCK) next_k_base = _compute_block_base_dw_i64(next_phys_block, stride_k_block, _k_head_off) next_scale_scalars = _load_kv_scale_scalars(next_tile_tok, next_phys_block) k_next_flat = _load_k_flat( @@ -1423,7 +1379,7 @@ def _unpack_states_kv(flat): running_max = fx.Float32(rmax_raw) running_sum = fx.Float32(rsum_raw) outs = [fx.Vector(out_raw) for out_raw in outs_raw] - outelems_norm = _normalize_pa_output(running_sum, outs, ZERO_F) + outelems_norm = _normalize_pa_output(running_sum, outs, fx.Float32(0.0)) qi_val_mg = qi_per_mtp[_mtp_g] qhi_pos_mg = qhi_per_mtp[_mtp_g] qhead = kv_h * fx.Int32(query_group_size) + qhi_pos_mg @@ -1447,7 +1403,7 @@ def _unpack_states_kv(flat): fx.slice(partial_out_tiles, (None, po_off // fx.Int32(4))), ) - safe_sum_lse = (running_sum > ZERO_F).select(running_sum, fx.Float32(1.0)) + safe_sum_lse = (running_sum > fx.Float32(0.0)).select(running_sum, fx.Float32(1.0)) log_sum = fmath.log(safe_sum_lse, fastmath=arith.FastMathFlags.fast) lse_val = running_max + log_sum pl_off = _po_row * stride_pl_ql + qhead From 9d507c1203136ffa6e0a84f9d613ac03700aa41d Mon Sep 17 00:00:00 2001 From: fsx950223 Date: Mon, 27 Jul 2026 04:50:14 +0000 Subject: [PATCH 3/7] perf(pa): Optimize metadata decode and persist grid tuning Signed-off-by: fsx950223 Co-authored-by: Cursor --- kernels/attention/pa_decode_fp8.py | 50 +++- kernels/attention/pa_metadata.py | 273 +++++++++--------- kernels/attention/pa_metadata_grid_tuning.csv | 19 ++ kernels/attention/pa_metadata_tuning.py | 139 +++++++++ scripts/tune_pa_metadata_grid.py | 227 +++++++++++++++ tests/kernels/test_pa.py | 1 + tests/unit/test_pa_metadata_tuning.py | 108 +++++++ 7 files changed, 662 insertions(+), 155 deletions(-) create mode 100644 kernels/attention/pa_metadata_grid_tuning.csv create mode 100644 kernels/attention/pa_metadata_tuning.py create mode 100644 scripts/tune_pa_metadata_grid.py create mode 100644 tests/unit/test_pa_metadata_tuning.py diff --git a/kernels/attention/pa_decode_fp8.py b/kernels/attention/pa_decode_fp8.py index 5bbb46c98..996d9feef 100644 --- a/kernels/attention/pa_decode_fp8.py +++ b/kernels/attention/pa_decode_fp8.py @@ -16,18 +16,16 @@ import torch +from flydsl.runtime.device import get_rocm_arch from kernels.attention.pa_decode_swa import compile_pa_decode_sw, compile_pa_decode_sw_reduce from kernels.attention.pa_decode_tile import pa_decode_tile from kernels.attention.pa_metadata import compile_pa_decode_metadata +from kernels.attention.pa_metadata_tuning import lookup_pa_metadata_grid_multiplier from kernels.common.tensor_shim import _run_compiled from kernels.common.utils import cdiv # ── Kernel geometry constants ──────────────────────────────────────── KV_COMPUTE_BLOCK = 256 # tile size (matches SP3 kTileKV) -# Persistent-grid oversubscription for the metadata decode path: launch -# CU_count * this many workgroups so the HW keeps multiple workgroups resident -# per CU (memory-latency hiding). 1 = original (1 wg/CU). -_PA_METADATA_GRID_OVERSUB = 3 MFMA_N = 16 @@ -44,6 +42,9 @@ def get_pa_metadata( num_query_heads: int, num_kv_heads: int, partition_size: int = KV_COMPUTE_BLOCK, + *, + per_token_kv: bool | None = None, + grid_multiplier: int | None = None, ): """Compute PA metadata (worklist, reduce maps) via get_pa_metadata_v1. @@ -59,6 +60,11 @@ def get_pa_metadata( NOTE: the consuming decode kernel must interpret kv_start/kv_end as partition indices accordingly. + Exact shape/device matches are loaded from ``pa_metadata_grid_tuning.csv``. + Missing entries default to ``grid_multiplier=1``. ``per_token_kv`` selects + scale-mode-specific tuning, and ``grid_multiplier`` is the tuner's explicit + candidate override. + Returns a dict with: work_indptr, work_info_flat, reduce_indptr, reduce_final_map, reduce_partial_map, num_sm, partial_output, partial_lse, stride_po_partial, stride_pl_partial. @@ -71,13 +77,25 @@ def get_pa_metadata( head_size = query.shape[-1] props = torch.cuda.get_device_properties(dev) - # Oversubscribe the persistent grid: the decode kernel is memory-latency-bound - # and only ~3 workgroups/CU fit by VGPR, but the worklist defaults to 1 wg/CU - # (grid = CU count). Distributing work across num_cu = CU_count * OVERSUB bins - # (and launching that many workgroups) lets the HW keep multiple workgroups - # resident per CU → more waves in flight → better latency hiding. - base_cu = props.multi_processor_count - num_sm = base_cu * _PA_METADATA_GRID_OVERSUB + num_blocks = key_cache.shape[0] + if grid_multiplier is None and per_token_kv is not None: + grid_multiplier = lookup_pa_metadata_grid_multiplier( + arch=get_rocm_arch(), + num_cu=props.multi_processor_count, + batch_size=batch_size, + num_blocks=num_blocks, + query_length=query_length, + per_token_kv=per_token_kv, + num_query_heads=num_query_heads, + num_kv_heads=num_kv_heads, + head_dim=head_size, + block_size=key_cache.shape[-2], + ) + if grid_multiplier is None: + grid_multiplier = 1 + if grid_multiplier < 1: + raise ValueError("grid_multiplier must be positive") + num_sm = props.multi_processor_count * grid_multiplier num_sm = (num_sm // num_kv_heads) * num_kv_heads # keep divisible by num_kv_heads seqlens_qo_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=dev) * query_length @@ -496,7 +514,15 @@ def pa_decode_ps_launch( "CUDA graph capture requires precomputed `metadata`; " "call `get_pa_metadata()` before capture and pass it via `metadata=`." ) - metadata = get_pa_metadata(query, key_cache, context_lengths, kv_indptr, num_query_heads, num_kv_heads) + metadata = get_pa_metadata( + query, + key_cache, + context_lengths, + kv_indptr, + num_query_heads, + num_kv_heads, + per_token_kv=per_token_kv, + ) work_indptr = metadata["work_indptr"] work_info_flat = metadata["work_info_flat"] diff --git a/kernels/attention/pa_metadata.py b/kernels/attention/pa_metadata.py index 36c9ed852..746d3c916 100644 --- a/kernels/attention/pa_metadata.py +++ b/kernels/attention/pa_metadata.py @@ -21,7 +21,7 @@ import flydsl.expr as fx from flydsl.expr import arith, const_expr, gpu, range_constexpr, rocdl from flydsl.expr import math as fmath -from flydsl.expr.typing import Int32, T +from flydsl.expr.typing import T from flydsl.runtime.device import get_rocm_arch from kernels.attention.pa_common import _compute_block_base_dw_i64 from kernels.common import dpp_utils @@ -126,9 +126,7 @@ def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = LDS_SOFTMAX_BYTES = 2 * NUM_WARPS * MFMA_N * 4 # 512 LDS_SCALE_V_PADDING = 4 # break K/V same-bank paired writes - LDS_SCALE_V_OFFSET = KV_COMPUTE_BLOCK + LDS_SCALE_V_PADDING - LDS_SCALE_BYTES = (LDS_SCALE_V_OFFSET + KV_COMPUTE_BLOCK) * 4 # K/V per-token scale staging FP8_MAX = 240.0 @@ -413,19 +411,6 @@ def _load_v_and_scales( tile_token_offset_i32, scale_scalars, ): - if const_expr(per_token_kv): - scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id - k_scale_scalar, v_scale_scalar = scale_scalars - fx.ptr_store( - fx.Vector.from_elements([k_scale_scalar], dtype=fx.Float32), - scale_base + scale_stage_token, - ) - fx.ptr_store( - fx.Vector.from_elements([v_scale_scalar], dtype=fx.Float32), - scale_base + (fx.Int32(LDS_SCALE_V_OFFSET) + scale_stage_token), - ) - rocdl.sched_barrier(0) - v_results = [] for vt in range_constexpr(VTLOOP): vhe_data = [] @@ -442,13 +427,24 @@ def _load_v_and_scales( v_results.append(vhe_data) if const_expr(per_token_kv): + scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id + k_scale_scalar, v_scale_scalar = scale_scalars + fx.ptr_store( + fx.Vector.from_elements([k_scale_scalar], dtype=fx.Float32), + scale_base + scale_stage_token, + ) + fx.ptr_store( + fx.Vector.from_elements([v_scale_scalar], dtype=fx.Float32), + scale_base + (fx.Int32(LDS_SCALE_V_OFFSET) + scale_stage_token), + ) + rocdl.sched_barrier(0) gpu.barrier() k_scale_vecs = [] v_scale_vecs = [] for td in range_constexpr(TLOOP): scale_row_base = kv_tok_thread_base + fx.Int32(td * MFMA_N) k_scale_vecs.append( - fx.ptr_load(scale_base + (scale_row_base), result_type=fx.Vector.make_type(4, fx.Float32)) + fx.ptr_load(scale_base + scale_row_base, result_type=fx.Vector.make_type(4, fx.Float32)) ) v_scale_vecs.append( fx.ptr_load( @@ -500,7 +496,10 @@ def _qk_and_intra_softmax( for td in range_constexpr(TLOOP): acc = fx.Vector.filled(4, 0.0, fx.Float32) for k_step in range_constexpr(qkhe_loop * 2): - acc = rocdl.mfma_f32_16x16x32_fp8_fp8(T.f32x4, [k_ops[td][k_step], q_frags[k_step], acc, 0, 0, 0]) + acc = rocdl.mfma_f32_16x16x32_fp8_fp8( + T.f32x4, + [k_ops[td][k_step], q_frags[k_step], acc, 0, 0, 0], + ) if const_expr(per_token_kv): d_out.append(fx.Vector(acc) * (k_scale_vecs[td] * query_scale)) else: @@ -523,7 +522,7 @@ def _qk_and_intra_softmax( return d_out, v_scale_vecs return d_out - def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs): + def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs, *, v_normalization=None): partition_max = neg_inf partition_sum = zero_f max_vec = fx.ptr_load(softmax_base + (sm_rd_max_offs[0]), result_type=fx.Vector.make_type(4, fx.Float32)) @@ -557,26 +556,27 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs): arith.unwrap(partition_sum), arith.unwrap(sum_vec[w]), fastmath=arith.FastMathFlags.contract ) - accum_sum = arith.mulf(arith.unwrap(accum_scale), arith.unwrap(rsum), fastmath=arith.FastMathFlags.contract) - rsum = arith.addf(accum_sum, arith.unwrap(partition_sum), fastmath=arith.FastMathFlags.contract) - rmax = new_rmax - for vhe in range_constexpr(vhe_loop): - outs[vhe] = outs[vhe] * fx.Float32(accum_scale) - if const_expr(per_token_kv): - v_max_global = zero_f - vmax_vec = fx.ptr_load(softmax_base + (sm_vmax_rd_offs[0]), result_type=fx.Vector.make_type(4, fx.Float32)) - for w in range_constexpr(NUM_WARPS): - w_vmax = vmax_vec[w] - v_max_global = fx.maxnumf(v_max_global, w_vmax) - v_max_scaled = v_max_global * fx.Float32(1.0 / FP8_MAX).ir_value() - v_max_safe_scaled = v_max_scaled + fx.Float32(1e-8 / FP8_MAX).ir_value() - norm_factor = rcp_f32(v_max_safe_scaled) - v_correction = v_max_scaled + if const_expr(v_normalization is None): + v_max_global = zero_f + vmax_vec = fx.ptr_load( + softmax_base + (sm_vmax_rd_offs[0]), result_type=fx.Vector.make_type(4, fx.Float32) + ) + for w in range_constexpr(NUM_WARPS): + w_vmax = vmax_vec[w] + v_max_global = fx.maxnumf(v_max_global, w_vmax) + v_correction = v_max_global * fx.Float32(1.0 / FP8_MAX).ir_value() + norm_factor = rcp_f32(v_correction + fx.Float32(1e-8 / FP8_MAX).ir_value()) + normalized_v_scales = [ + fx.Vector(v_scale_vecs[td]) * fx.Float32(norm_factor) for td in range_constexpr(TLOOP) + ] + else: + v_correction, normalized_v_scales = v_normalization for td in range_constexpr(TLOOP): - d_out[td] = d_out[td] * (v_scale_vecs[td] * fx.Float32(norm_factor)) + d_out[td] = d_out[td] * normalized_v_scales[td] else: v_correction = v_scale_val + normalized_v_scales = [] for td in range_constexpr(TLOOP): pv = fx.Vector(d_out[td]) @@ -585,7 +585,13 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs): elem_base = prob_wr_thread_base + fx.Int32(td * MFMA_N * (PROB_ROW_STRIDE_BYTES // 4)) pk_vec = fx.Vector.from_elements([pk], dtype=fx.Int32) fx.ptr_store(pk_vec, logits_base + elem_base) - return rmax, rsum, outs, v_correction + + accum_sum = arith.mulf(arith.unwrap(accum_scale), arith.unwrap(rsum), fastmath=arith.FastMathFlags.contract) + rsum = arith.addf(accum_sum, arith.unwrap(partition_sum), fastmath=arith.FastMathFlags.contract) + rmax = new_rmax + for vhe in range_constexpr(vhe_loop): + outs[vhe] = outs[vhe] * fx.Float32(accum_scale) + return rmax, rsum, outs, v_correction, normalized_v_scales def _pv_mfma(v_ops, outs, v_correction): fm_contract = arith.FastMathFlags.contract @@ -665,7 +671,7 @@ def pa_metadata_v1_kernel( reduce_indptr_ptr: fx.Tensor, # [num_batches + 1] i32 (output) reduce_final_map_ptr: fx.Tensor, # [num_batches * 2] i32 (output, flattened) reduce_partial_map_ptr: fx.Tensor, # [max_split_tiles] i32 (output) - num_batches: Int32, + num_batches: fx.Int32, ): scalar_layout = fx.make_layout(1, 1) work_layout = fx.make_layout(4, 1) @@ -693,7 +699,7 @@ def _divide_buffer(tensor, tile_layout): c0 = fx.Int32(0) c1 = fx.Int32(1) c_qlen = fx.Int32(query_length) - c_nb = num_batches # Int32 runtime + c_nb = num_batches # fx.Int32 runtime c_nspk = fx.Int32(num_splits_per_khead) c_numcu = fx.Int32(num_cu) c_ws = fx.Int32(warp_size) @@ -885,7 +891,7 @@ def launch_pa_metadata_v1( reduce_indptr: fx.Tensor, reduce_final_map: fx.Tensor, reduce_partial_map: fx.Tensor, - num_batches: Int32, + num_batches: fx.Int32, stream: fx.Stream = fx.Stream(None), ): pa_metadata_v1_kernel( @@ -988,7 +994,6 @@ def compile_pa_decode_metadata( _QKHELOOP = head_dim // QKHE_PER_FETCH _VHELOOP = head_dim // MFMA_N // NUM_WARPS _Q_LANES_PER_HEAD = head_dim // Q_ELEMS_PER_LANE - _N_K_h = TLOOP * _QKHELOOP * 2 if query_input_dtype not in ("bf16", "f16"): raise ValueError(f"`compile_pa_decode_metadata` only supports bf16/f16 queries, got {query_input_dtype!r}") _QUERY_DTYPE = fx.BFloat16 if query_input_dtype == "bf16" else fx.Float16 @@ -1011,7 +1016,7 @@ class SharedStorage: buf: fx.Array[fx.Int32, _LDS_TOTAL_BYTES // 4, 16] @flyc.kernel(known_block_size=(BLOCK_THREADS, 1, 1)) - def pa_decode_metadata_kenrel( + def pa_decode_metadata_kernel( out_ptr: fx.Int64, # output [batch, num_q_heads, head_dim] partial_out_ptr: fx.Int64, # partial output [num_partials, 1, nhead, head_dim] fp32 partial_lse_ptr: fx.Int64, # partial LSE [num_partials, 1, nhead, 1] fp32 @@ -1026,20 +1031,20 @@ def pa_decode_metadata_kenrel( kv_page_indices_ptr: fx.Int64, # [total_pages] int32 kv_indptr_ptr: fx.Int64, # [num_seqs + 1] int32 — prefix sum of pages per seq partition_indptr_ptr: fx.Int64, # [num_seqs + 1] int32 — prefix sum of partitions per seq - stride_q_seq: Int32, - stride_q_head: Int32, - stride_k_block: Int32, - stride_k_head: Int32, - stride_v_block: Int32, - stride_v_head: Int32, - stride_out_seq: Int32, - stride_out_head: Int32, - stride_po_partial: Int32, # stride for partial_output partial dim (nhead * head_dim) - stride_pl_partial: Int32, # stride for partial_lse partial dim (nhead) - stride_ks_block: Int32, # key_scale stride for block dim (num_kv_heads * KV_BLOCK_SIZE); 0 for per-tensor - stride_ks_head: Int32, # key_scale stride for head dim (KV_BLOCK_SIZE); 0 for per-tensor - stride_po_ql: Int32, # stride for partial_output query-length dim (num_query_heads * head_dim) - stride_pl_ql: Int32, # stride for partial_lse query-length dim (num_query_heads) + stride_q_seq: fx.Int32, + stride_q_head: fx.Int32, + stride_k_block: fx.Int32, + stride_k_head: fx.Int32, + stride_v_block: fx.Int32, + stride_v_head: fx.Int32, + stride_out_seq: fx.Int32, + stride_out_head: fx.Int32, + stride_po_partial: fx.Int32, # stride for partial_output partial dim (nhead * head_dim) + stride_pl_partial: fx.Int32, # stride for partial_lse partial dim (nhead) + stride_ks_block: fx.Int32, # key_scale stride for block dim (num_kv_heads * KV_BLOCK_SIZE); 0 for per-tensor + stride_ks_head: fx.Int32, # key_scale stride for head dim (KV_BLOCK_SIZE); 0 for per-tensor + stride_po_ql: fx.Int32, # stride for partial_output query-length dim (num_query_heads * head_dim) + stride_pl_ql: fx.Int32, # stride for partial_lse query-length dim (num_query_heads) ): tid = fx.Int32(gpu.thread_id("x")) cu_id = fx.Int32(gpu.block_id("x")) # CU index (0..num_sm-1) @@ -1227,41 +1232,13 @@ def _unwrap(v): context_len + fx.Int32(1 - query_length) + qi_per_mtp[_mtp_g] for _mtp_g in range(_mtp_groups) ] - local_last_part = local_part_start + last_part_idx_val - first_phys_block = _copy_load( - kv_page_indices, - kv_page_base + udiv_const(local_last_part, parts_per_block), - copy_i32, - i32_reg, - )[0] - first_tile_tok = urem_const(local_last_part, parts_per_block) * fx.Int32(KV_COMPUTE_BLOCK) - first_k_base = _compute_block_base_dw_i64(first_phys_block, stride_k_block, _k_head_off) - scale_scalars0 = _load_kv_scale_scalars(first_tile_tok, first_phys_block) - k_flat0 = _load_k_flat( - k_global_ptr, - global_copy_16b, - global_reg_16b, - first_k_base, - first_tile_tok, - _k_tok_thread_base, - _c_tok_stride_dw, - _k_he_off_dw, - qkhe_loop=_QKHELOOP, - ) - state_width = 2 + _VHELOOP - def _pack_states_kv(states, k_flat, scale_scalars): - flat = [_unwrap(value) for state in states for value in state] - flat.extend(_unwrap(v) for v in k_flat) - flat.extend(_unwrap(v) for v in scale_scalars) - return flat + def _pack_states(states): + return [_unwrap(value) for state in states for value in state] - def _unpack_states_kv(flat): - state_end = state_width * _mtp_groups - k_end = state_end + _N_K_h - states = [tuple(flat[state_width * i : state_width * (i + 1)]) for i in range(_mtp_groups)] - return states, list(flat[state_end:k_end]), tuple(flat[k_end:]) + def _unpack_states(flat): + return [tuple(flat[state_width * i : state_width * (i + 1)]) for i in range(_mtp_groups)] init_states = [ tuple( @@ -1275,17 +1252,12 @@ def _unpack_states_kv(flat): fx.Int32(0), num_parts_in_work, fx.Int32(1), - init=_pack_states_kv(init_states, k_flat0, scale_scalars0), + init=_pack_states(init_states), ): - cur_states, k_flat, scale_scalars = _unpack_states_kv(state) + cur_states = _unpack_states(state) # Process partition zero last for online-softmax stability. rel_part = last_part_idx_val - ib lp = local_part_start + rel_part - next_rel = rel_part - fx.Int32(1) - next_rel_clamped = (next_rel >= fx.Int32(0)).select(next_rel, fx.Int32(0)) - next_lp = local_part_start + next_rel_clamped - - k_ops = unflatten_k(k_flat, qkhe_loop=_QKHELOOP) partition_start = lp * fx.Int32(KV_COMPUTE_BLOCK) phys_block = _copy_load( @@ -1295,6 +1267,22 @@ def _unpack_states_kv(flat): i32_reg, )[0] tile_token_offset = urem_const(lp, parts_per_block) * fx.Int32(KV_COMPUTE_BLOCK) + k_base = _compute_block_base_dw_i64(phys_block, stride_k_block, _k_head_off) + scale_scalars = _load_kv_scale_scalars(tile_token_offset, phys_block) + if const_expr(per_token_kv): + rocdl.sched_barrier(0) + k_flat = _load_k_flat( + k_global_ptr, + global_copy_16b, + global_reg_16b, + k_base, + tile_token_offset, + _k_tok_thread_base, + _c_tok_stride_dw, + _k_he_off_dw, + qkhe_loop=_QKHELOOP, + ) + k_ops = unflatten_k(k_flat, qkhe_loop=_QKHELOOP) v_base = _compute_block_base_dw_i64(phys_block, stride_v_block, _v_head_off) if const_expr(per_token_kv): v_ops, k_scale_vecs, v_scale_vecs = _load_v_and_scales( @@ -1308,10 +1296,10 @@ def _unpack_states_kv(flat): tile_token_offset, scale_scalars, ) + + v_normalization = None new_states = [] for _mtp_g in range_constexpr(_mtp_groups): - if const_expr(_mtp_g > 0): - gpu.barrier() state = cur_states[_mtp_g] rmax, rsum = state[0], state[1] outs = [state[2 + vhe] for vhe in range_constexpr(_VHELOOP)] @@ -1335,42 +1323,27 @@ def _unpack_states_kv(flat): ) v_scales = None - # Stage vmax before the cross-wave normalization reads it. - if const_expr(per_token_kv): + if const_expr(per_token_kv and _mtp_g == 0): _store_vmax_warp(partition_start, seq_end=context_len, v_scale_vecs=v_scales) gpu.barrier() - rmax, rsum, outs, v_correction = _cross_warp_softmax_and_prob_pack( - d_out, rmax, rsum, outs, v_scales + rmax, rsum, outs, v_correction, normalized_v_scales = _cross_warp_softmax_and_prob_pack( + d_out, + rmax, + rsum, + outs, + v_scales, + v_normalization=v_normalization, ) + if const_expr(per_token_kv and _mtp_g == 0): + v_normalization = (v_correction, normalized_v_scales) gpu.barrier() outs = _pv_mfma(v_ops, outs, v_correction) new_states.append(tuple([rmax, rsum] + outs)) - next_phys_block = _copy_load( - kv_page_indices, - kv_page_base + udiv_const(next_lp, parts_per_block), - copy_i32, - i32_reg, - )[0] - next_tile_tok = urem_const(next_lp, parts_per_block) * fx.Int32(KV_COMPUTE_BLOCK) - next_k_base = _compute_block_base_dw_i64(next_phys_block, stride_k_block, _k_head_off) - next_scale_scalars = _load_kv_scale_scalars(next_tile_tok, next_phys_block) - k_next_flat = _load_k_flat( - k_global_ptr, - global_copy_16b, - global_reg_16b, - next_k_base, - next_tile_tok, - _k_tok_thread_base, - _c_tok_stride_dw, - _k_he_off_dw, - qkhe_loop=_QKHELOOP, - ) - - results = yield _pack_states_kv(new_states, k_next_flat, next_scale_scalars) + results = yield _pack_states(new_states) - final_states, _, _ = _unpack_states_kv(results) + final_states = _unpack_states(results) for _mtp_g in range_constexpr(_mtp_groups): final_state = final_states[_mtp_g] @@ -1443,7 +1416,7 @@ def launch_pa_decode_metadata( num_sm, stream: fx.Stream = fx.Stream(None), ): - pa_decode_metadata_kenrel( + pa_decode_metadata_kernel( out, po, pl, @@ -1476,7 +1449,7 @@ def launch_pa_decode_metadata( return { "launch": launch_pa_decode_metadata, - "kernel": pa_decode_metadata_kenrel, + "kernel": pa_decode_metadata_kernel, } @@ -1489,7 +1462,8 @@ def compile_pa_metadata_reduce( head_dim: int, output_dtype_str: str, ): - block_threads = head_dim + vec_width = 2 if head_dim % 2 == 0 else 1 + block_threads = head_dim // vec_width assert 0 < block_threads <= 1024, "head_dim must fit in one workgroup" output_dtype = _PA_OUTPUT_DTYPES[output_dtype_str] @@ -1501,10 +1475,10 @@ def pa_metadata_reduce_kernel( reduce_indptr_ptr: fx.Tensor, reduce_final_map_ptr: fx.Tensor, reduce_partial_map_ptr: fx.Tensor, - stride_out_seq: Int32, - stride_out_head: Int32, - stride_po_row: Int32, # num_query_heads * head_dim - stride_pl_row: Int32, # num_query_heads + stride_out_seq: fx.Int32, + stride_out_head: fx.Int32, + stride_po_row: fx.Int32, # num_query_heads * head_dim + stride_pl_row: fx.Int32, # num_query_heads ): tid = fx.Int32(gpu.thread_id("x")) g = fx.Int32(gpu.block_id("x")) @@ -1528,13 +1502,22 @@ def _flat_divide(tensor): copy_i32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Int32) copy_f32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Float32) + copy_f32_vec = fx.make_copy_atom( + fx.rocdl.BufferCopy64b() if vec_width == 2 else fx.rocdl.BufferCopy32b(), + fx.Float32, + ) copy_output = fx.make_copy_atom( - fx.rocdl.BufferCopy32b() if output_dtype.width == 32 else fx.rocdl.BufferCopy16b(), + ( + fx.rocdl.BufferCopy64b() + if output_dtype.width * vec_width == 64 + else fx.rocdl.BufferCopy32b() if output_dtype.width * vec_width == 32 else fx.rocdl.BufferCopy16b() + ), output_dtype, ) reg_i32 = fx.make_rmem_tensor(1, fx.Int32) reg_f32 = fx.make_rmem_tensor(1, fx.Float32) - reg_output = fx.make_rmem_tensor(1, output_dtype) + reg_f32_vec = fx.make_rmem_tensor(vec_width, fx.Float32) + reg_output = fx.make_rmem_tensor(vec_width, output_dtype) lo = _copy_load(reduce_indptr, g, copy_i32, reg_i32)[0] hi = _copy_load(reduce_indptr, g + fx.Int32(1), copy_i32, reg_i32)[0] @@ -1547,34 +1530,38 @@ def _flat_divide(tensor): lo, hi, fx.Int32(1), - init=[fx.Float32(float("-inf")), fx.Float32(0.0), fx.Float32(0.0)], + init=[ + fx.Float32(float("-inf")), + fx.Float32(0.0), + fx.Vector.filled(vec_width, 0.0, fx.Float32), + ], ): m = fx.Float32(st[0]) denom = fx.Float32(st[1]) - acc = fx.Float32(st[2]) + acc = fx.Vector(st[2]) prow = _copy_load(reduce_partial_map, slot, copy_i32, reg_i32)[0] + qi lse = _copy_load(partial_lse, prow * stride_pl_row + qhead, copy_f32, reg_f32)[0] v = _copy_load( partial_output, - prow * stride_po_row + qhead * fx.Int32(head_dim) + tid, - copy_f32, - reg_f32, - )[0] + prow * stride_po_row + qhead * fx.Int32(head_dim) + tid * fx.Int32(vec_width), + copy_f32_vec, + reg_f32_vec, + ) m_new = m.maximumf(lse) scale_old = exp2_f32_fast((m - m_new) * fx.Float32(LOG2E)) w = exp2_f32_fast((lse - m_new) * fx.Float32(LOG2E)) denom_new = denom * scale_old + w - acc_new = acc * scale_old + v * w + acc_new = acc * fx.Float32(scale_old) + fx.Vector(v) * fx.Float32(w) results = yield [m_new, denom_new, acc_new] denom_f = fx.Float32(results[1]) - acc_f = fx.Float32(results[2]) + acc_f = fx.Vector(results[2]) safe_denom = (denom_f > fx.Float32(0.0)).select(denom_f, fx.Float32(1.0)) out_acc = acc_f * rcp_f32(safe_denom) - out_off = out_row * stride_out_seq + qhead * stride_out_head + tid - output_val = fx.Float32(out_acc).to(output_dtype) - fx.memref_store_vec(fx.Vector.from_elements([output_val], dtype=output_dtype), reg_output) + out_off = out_row * stride_out_seq + qhead * stride_out_head + tid * fx.Int32(vec_width) + output_val = fx.Vector(out_acc).to(output_dtype) + fx.memref_store_vec(output_val, reg_output) fx.copy(copy_output, reg_output, fx.slice(final_output, (None, out_off))) @flyc.jit diff --git a/kernels/attention/pa_metadata_grid_tuning.csv b/kernels/attention/pa_metadata_grid_tuning.csv new file mode 100644 index 000000000..9389628ec --- /dev/null +++ b/kernels/attention/pa_metadata_grid_tuning.csv @@ -0,0 +1,19 @@ +arch,num_cu,batch_size,num_blocks,context_length,query_length,per_token_kv,num_query_heads,num_kv_heads,head_dim,block_size,grid_multiplier,latency_us +gfx942,80,1,1,1024,1,0,16,1,128,1024,1,11.440 +gfx942,80,3,6,1027,4,1,16,1,128,1024,1,21.160 +gfx942,80,8,256,32768,4,1,16,1,128,1024,1,122.200 +gfx942,80,16,64,4096,1,0,16,1,128,1024,1,20.680 +gfx942,80,16,128,8192,4,0,16,1,128,1024,2,58.801 +gfx942,80,16,128,8192,4,1,16,1,128,1024,1,75.000 +gfx942,80,32,256,8192,1,0,16,1,128,1024,3,35.520 +gfx942,80,64,256,4096,1,1,16,1,128,1024,3,38.360 +gfx942,80,64,256,4096,2,0,16,1,128,1024,2,53.600 +gfx942,80,64,256,4096,2,1,16,1,128,1024,2,59.920 +gfx942,80,64,256,4096,3,0,16,1,128,1024,2,71.401 +gfx942,80,64,256,4096,3,1,16,1,128,1024,2,81.800 +gfx942,80,64,512,8192,1,0,16,1,128,1024,3,49.520 +gfx942,80,81,648,8192,1,0,16,1,128,1024,3,65.840 +gfx942,80,81,648,8192,4,0,16,1,128,1024,2,184.720 +gfx942,80,81,648,8192,4,1,16,1,128,1024,1,263.400 +gfx942,80,128,1024,8192,4,1,16,1,128,1024,1,415.640 +gfx942,80,256,512,2048,4,1,16,1,128,1024,1,232.000 diff --git a/kernels/attention/pa_metadata_tuning.py b/kernels/attention/pa_metadata_tuning.py new file mode 100644 index 000000000..5a117315d --- /dev/null +++ b/kernels/attention/pa_metadata_tuning.py @@ -0,0 +1,139 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +from __future__ import annotations + +import csv +import functools +from pathlib import Path +from typing import Iterable, Mapping + +PA_METADATA_TUNING_CSV = Path(__file__).with_name("pa_metadata_grid_tuning.csv") + +PA_METADATA_TUNING_FIELDS = ( + "arch", + "num_cu", + "batch_size", + "num_blocks", + "context_length", + "query_length", + "per_token_kv", + "num_query_heads", + "num_kv_heads", + "head_dim", + "block_size", + "grid_multiplier", + "latency_us", +) + +_KEY_FIELDS = ( + "arch", + "num_cu", + "batch_size", + "num_blocks", + "query_length", + "per_token_kv", + "num_query_heads", + "num_kv_heads", + "head_dim", + "block_size", +) + +_INT_FIELDS = set(_KEY_FIELDS) - {"arch", "per_token_kv"} | { + "context_length", + "grid_multiplier", +} + + +def _normalize_row(row: Mapping) -> dict: + normalized = {field: row[field] for field in PA_METADATA_TUNING_FIELDS} + normalized["arch"] = str(normalized["arch"]) + for field in _INT_FIELDS: + normalized[field] = int(normalized[field]) + per_token_kv = normalized["per_token_kv"] + if isinstance(per_token_kv, str): + if per_token_kv not in {"0", "1"}: + raise ValueError(f"per_token_kv must be 0 or 1, got {per_token_kv!r}") + per_token_kv = per_token_kv == "1" + normalized["per_token_kv"] = bool(per_token_kv) + normalized["latency_us"] = float(normalized["latency_us"]) + if normalized["grid_multiplier"] < 1: + raise ValueError("grid_multiplier must be positive") + return normalized + + +def _row_key(row: Mapping) -> tuple: + return tuple(row[field] for field in _KEY_FIELDS) + + +def read_pa_metadata_tuning_rows(path: Path = PA_METADATA_TUNING_CSV) -> list[dict]: + path = Path(path) + if not path.is_file(): + return [] + with path.open(newline="") as handle: + reader = csv.DictReader(handle) + missing = set(PA_METADATA_TUNING_FIELDS) - set(reader.fieldnames or ()) + if missing: + raise ValueError(f"{path} is missing columns: {sorted(missing)}") + rows = [_normalize_row(row) for row in reader] + keys = [_row_key(row) for row in rows] + if len(keys) != len(set(keys)): + raise ValueError(f"{path} contains duplicate tuning keys") + return rows + + +@functools.lru_cache(maxsize=None) +def _load_pa_metadata_tuning(path: Path) -> dict[tuple, int]: + return {_row_key(row): row["grid_multiplier"] for row in read_pa_metadata_tuning_rows(path)} + + +def lookup_pa_metadata_grid_multiplier( + *, + arch: str, + num_cu: int, + batch_size: int, + num_blocks: int, + query_length: int, + per_token_kv: bool, + num_query_heads: int, + num_kv_heads: int, + head_dim: int, + block_size: int, + path: Path = PA_METADATA_TUNING_CSV, +) -> int | None: + key = ( + str(arch), + int(num_cu), + int(batch_size), + int(num_blocks), + int(query_length), + bool(per_token_kv), + int(num_query_heads), + int(num_kv_heads), + int(head_dim), + int(block_size), + ) + return _load_pa_metadata_tuning(Path(path)).get(key) + + +def write_pa_metadata_tuning_rows( + rows: Iterable[Mapping], + path: Path = PA_METADATA_TUNING_CSV, +) -> None: + path = Path(path) + merged = {_row_key(row): row for row in read_pa_metadata_tuning_rows(path)} + for row in rows: + normalized = _normalize_row(row) + merged[_row_key(normalized)] = normalized + + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_suffix(f"{path.suffix}.tmp") + with temporary.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=PA_METADATA_TUNING_FIELDS) + writer.writeheader() + for key in sorted(merged): + row = dict(merged[key]) + row["per_token_kv"] = int(row["per_token_kv"]) + writer.writerow(row) + temporary.replace(path) + _load_pa_metadata_tuning.cache_clear() diff --git a/scripts/tune_pa_metadata_grid.py b/scripts/tune_pa_metadata_grid.py new file mode 100644 index 000000000..3a4f890ee --- /dev/null +++ b/scripts/tune_pa_metadata_grid.py @@ -0,0 +1,227 @@ +#!/usr/bin/env python3 +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Tune PA metadata grid multipliers and persist winners to CSV. + +Example: + python3 scripts/tune_pa_metadata_grid.py \ + --shape 16,4096,1,per_tensor \ + --shape 81,8192,4,per_token +""" + +from __future__ import annotations + +import argparse +import statistics +import sys +from pathlib import Path + +import torch + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from flydsl.runtime.device import get_rocm_arch +from kernels.attention.pa_decode_fp8 import get_pa_metadata, pa_decode_ps_launch +from kernels.attention.pa_metadata_tuning import ( + PA_METADATA_TUNING_CSV, + write_pa_metadata_tuning_rows, +) + + +def _parse_int_list(value: str) -> list[int]: + values = [int(item) for item in value.split(",")] + if not values or any(item < 1 for item in values): + raise argparse.ArgumentTypeError("expected comma-separated positive integers") + return values + + +def _parse_shape(value: str) -> tuple[int, int, int, bool]: + try: + batch_size, context_length, query_length, quant_mode = value.split(",") + per_token_kv = {"per_tensor": False, "per_token": True}[quant_mode] + parsed = (int(batch_size), int(context_length), int(query_length), per_token_kv) + except (KeyError, TypeError, ValueError) as error: + raise argparse.ArgumentTypeError("shape must be BATCH,CONTEXT,QUERY_LENGTH,per_tensor|per_token") from error + if any(item < 1 for item in parsed[:3]): + raise argparse.ArgumentTypeError("shape dimensions must be positive") + return parsed + + +def _measure_candidate( + *, + batch_size: int, + context_length: int, + query_length: int, + per_token_kv: bool, + grid_multiplier: int, + num_query_heads: int, + num_kv_heads: int, + head_dim: int, + block_size: int, + warmup: int, + iterations: int, + rounds: int, + device: torch.device, +) -> tuple[float, int]: + pages_per_sequence = (context_length + block_size - 1) // block_size + num_blocks = batch_size * pages_per_sequence + context_lengths = torch.full((batch_size,), context_length, dtype=torch.int32, device=device) + query = torch.zeros( + (batch_size * query_length, num_query_heads, head_dim), + dtype=torch.bfloat16, + device=device, + ) + kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * pages_per_sequence + kv_page_indices = torch.arange(num_blocks, dtype=torch.int32, device=device) + key = torch.zeros( + (num_blocks, num_kv_heads, head_dim // 16, block_size, 16), + dtype=torch.float8_e4m3fnuz, + device=device, + ) + value = torch.ones( + (num_blocks, num_kv_heads, block_size // 16, head_dim, 16), + dtype=torch.float8_e4m3fnuz, + device=device, + ) + if per_token_kv: + key_scale = torch.ones((num_blocks, num_kv_heads, block_size), dtype=torch.float32, device=device) + value_scale = torch.ones_like(key_scale) + else: + key_scale = torch.ones((1,), dtype=torch.float32, device=device) + value_scale = torch.ones((1,), dtype=torch.float32, device=device) + + metadata = get_pa_metadata( + query, + key, + context_lengths, + kv_indptr, + num_query_heads, + num_kv_heads, + per_token_kv=per_token_kv, + grid_multiplier=grid_multiplier, + ) + output = torch.empty_like(query) + + def launch(): + pa_decode_ps_launch( + output, + query, + key, + value, + context_lengths, + kv_page_indices, + kv_indptr, + 1.0 / head_dim**0.5, + key_scale=key_scale, + value_scale=value_scale, + metadata=metadata, + ) + + for _ in range(warmup): + launch() + torch.cuda.synchronize(device) + + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for _ in range(iterations): + launch() + torch.cuda.synchronize(device) + + timings = [] + for _ in range(rounds): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + graph.replay() + end.record() + end.synchronize() + timings.append(start.elapsed_time(end) * 1000.0 / iterations) + + if not torch.allclose( + output.float(), + torch.ones_like(output, dtype=torch.float32), + atol=2e-2, + rtol=2e-2, + ): + raise RuntimeError(f"grid_multiplier={grid_multiplier} failed correctness") + return statistics.median(timings), num_blocks + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--shape", action="append", type=_parse_shape, required=True) + parser.add_argument("--candidates", type=_parse_int_list, default=[1, 2, 3]) + parser.add_argument("--warmup", type=int, default=5) + parser.add_argument("--iterations", type=int, default=100) + parser.add_argument("--rounds", type=int, default=3) + parser.add_argument("--device", type=int, default=0) + parser.add_argument("--num-query-heads", type=int, default=16) + parser.add_argument("--num-kv-heads", type=int, default=1) + parser.add_argument("--head-dim", type=int, default=128) + parser.add_argument("--block-size", type=int, default=1024) + parser.add_argument("--output", type=Path, default=PA_METADATA_TUNING_CSV) + args = parser.parse_args() + + if args.warmup < 0 or args.iterations < 1 or args.rounds < 1: + parser.error("warmup must be non-negative; iterations and rounds must be positive") + + torch.cuda.set_device(args.device) + device = torch.device("cuda", args.device) + arch = get_rocm_arch() + num_cu = torch.cuda.get_device_properties(device).multi_processor_count + tuned_rows = [] + + for batch_size, context_length, query_length, per_token_kv in args.shape: + timings = {} + num_blocks = 0 + for grid_multiplier in args.candidates: + latency_us, num_blocks = _measure_candidate( + batch_size=batch_size, + context_length=context_length, + query_length=query_length, + per_token_kv=per_token_kv, + grid_multiplier=grid_multiplier, + num_query_heads=args.num_query_heads, + num_kv_heads=args.num_kv_heads, + head_dim=args.head_dim, + block_size=args.block_size, + warmup=args.warmup, + iterations=args.iterations, + rounds=args.rounds, + device=device, + ) + timings[grid_multiplier] = latency_us + print( + f"B={batch_size} C={context_length} Q={query_length} " + f"per_token={per_token_kv} grid={grid_multiplier}: {latency_us:.3f} us" + ) + + best_grid = min(timings, key=timings.get) + best_latency = timings[best_grid] + print(f" winner: grid={best_grid}, {best_latency:.3f} us") + tuned_rows.append( + { + "arch": arch, + "num_cu": num_cu, + "batch_size": batch_size, + "num_blocks": num_blocks, + "context_length": context_length, + "query_length": query_length, + "per_token_kv": per_token_kv, + "num_query_heads": args.num_query_heads, + "num_kv_heads": args.num_kv_heads, + "head_dim": args.head_dim, + "block_size": args.block_size, + "grid_multiplier": best_grid, + "latency_us": best_latency, + } + ) + torch.cuda.empty_cache() + + write_pa_metadata_tuning_rows(tuned_rows, args.output) + print(f"Wrote {len(tuned_rows)} tuned configuration(s) to {args.output}") + + +if __name__ == "__main__": + main() diff --git a/tests/kernels/test_pa.py b/tests/kernels/test_pa.py index 04179b328..4957e302a 100644 --- a/tests/kernels/test_pa.py +++ b/tests/kernels/test_pa.py @@ -751,6 +751,7 @@ def gluon_call() -> None: kv_indptr, num_query_heads, num_kv_heads, + per_token_kv=quant_mode == "per_token", ) ps_key_scale: torch.Tensor = key_scale_original ps_value_scale: torch.Tensor = value_scale_original diff --git a/tests/unit/test_pa_metadata_tuning.py b/tests/unit/test_pa_metadata_tuning.py new file mode 100644 index 000000000..d9b175cf5 --- /dev/null +++ b/tests/unit/test_pa_metadata_tuning.py @@ -0,0 +1,108 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import csv + +import pytest + +from kernels.attention.pa_metadata_tuning import ( + PA_METADATA_TUNING_CSV, + lookup_pa_metadata_grid_multiplier, + read_pa_metadata_tuning_rows, + write_pa_metadata_tuning_rows, +) + + +def _row(**overrides): + row = { + "arch": "gfx942", + "num_cu": 80, + "batch_size": 81, + "num_blocks": 648, + "context_length": 8192, + "query_length": 4, + "per_token_kv": True, + "num_query_heads": 16, + "num_kv_heads": 1, + "head_dim": 128, + "block_size": 1024, + "grid_multiplier": 1, + "latency_us": 263.4, + } + row.update(overrides) + return row + + +def test_default_pa_metadata_tuning_table(): + rows = read_pa_metadata_tuning_rows() + assert rows + assert ( + lookup_pa_metadata_grid_multiplier( + arch="gfx942", + num_cu=80, + batch_size=81, + num_blocks=648, + query_length=4, + per_token_kv=True, + num_query_heads=16, + num_kv_heads=1, + head_dim=128, + block_size=1024, + ) + == 1 + ) + assert PA_METADATA_TUNING_CSV.is_file() + + +def test_pa_metadata_tuning_csv_roundtrip_and_upsert(tmp_path): + path = tmp_path / "pa_tuning.csv" + write_pa_metadata_tuning_rows([_row()], path) + write_pa_metadata_tuning_rows( + [_row(grid_multiplier=2, latency_us=250.0), _row(batch_size=32, num_blocks=256)], path + ) + + rows = read_pa_metadata_tuning_rows(path) + assert len(rows) == 2 + assert ( + lookup_pa_metadata_grid_multiplier( + arch="gfx942", + num_cu=80, + batch_size=81, + num_blocks=648, + query_length=4, + per_token_kv=True, + num_query_heads=16, + num_kv_heads=1, + head_dim=128, + block_size=1024, + path=path, + ) + == 2 + ) + assert ( + lookup_pa_metadata_grid_multiplier( + arch="gfx950", + num_cu=80, + batch_size=81, + num_blocks=648, + query_length=4, + per_token_kv=True, + num_query_heads=16, + num_kv_heads=1, + head_dim=128, + block_size=1024, + path=path, + ) + is None + ) + + +def test_pa_metadata_tuning_rejects_invalid_boolean(tmp_path): + path = tmp_path / "invalid.csv" + with path.open("w", newline="") as handle: + writer = csv.DictWriter(handle, fieldnames=_row()) + writer.writeheader() + writer.writerow(_row(per_token_kv="true")) + + with pytest.raises(ValueError, match="per_token_kv"): + read_pa_metadata_tuning_rows(path) From 67cfd4c1584a050a12e6930e0bccf1521beb123b Mon Sep 17 00:00:00 2001 From: fsx950223 Date: Mon, 27 Jul 2026 06:00:52 +0000 Subject: [PATCH 4/7] refactor(pa): Share common memory copy helpers Signed-off-by: fsx950223 Co-authored-by: Cursor --- kernels/attention/pa_metadata.py | 138 ++++++++++++++++++------------- kernels/common/utils.py | 25 ++++++ 2 files changed, 105 insertions(+), 58 deletions(-) diff --git a/kernels/attention/pa_metadata.py b/kernels/attention/pa_metadata.py index 746d3c916..d1dc70470 100644 --- a/kernels/attention/pa_metadata.py +++ b/kernels/attention/pa_metadata.py @@ -28,8 +28,12 @@ from kernels.common.kernels_common import get_warp_size from kernels.common.tensor_shim import _run_compiled from kernels.common.utils import ( + copy_load, + copy_store, exp2_f32_fast, + global_pointer_from_addr, is_pow2, + load_global_16b, rcp_f32, udiv_const, unflatten_k, @@ -45,26 +49,6 @@ } -def _global_pointer_from_addr(addr, dtype, *, alignment: int): - ptr_type = fx.PointerType.get( - elem_ty=dtype.ir_type, - address_space=fx.AddressSpace.Global, - alignment=alignment, - ) - return fx.inttoptr(ptr_type, addr) - - -def _copy_load(source, offset, copy_atom, register): - fx.copy(copy_atom, fx.slice(source, (None, fx.Int32(offset))), register) - return fx.memref_load_vec(register) - - -def _load_global_16b(global_ptr, byte_offset, copy_atom, reg): - src = fx.make_view(global_ptr + byte_offset, fx.make_layout(16, 1)) - fx.copy(copy_atom, src, reg) - return fx.memref_load_vec(reg).bitcast(fx.Int64) - - def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = None): """Return buffer sizes and dtypes for the PA worklist tensors. @@ -154,7 +138,7 @@ def _load_k_flat( kbo_dw = kbo * c_tok_stride_dw for qkhe in range_constexpr(qkhe_loop): ka_dw = k_block_base_dw_i64 + fx.Int64(kbo_dw + k_he_off_dw[qkhe]) - k2 = _load_global_16b(k_global_ptr, ka_dw * fx.Int64(4), k_copy_atom, k_reg) + k2 = load_global_16b(k_global_ptr, ka_dw * fx.Int64(4), k_copy_atom, k_reg) k2_words = fx.Vector(k2) k_flat.append(k2_words[0]) k_flat.append(k2_words[1]) @@ -312,7 +296,7 @@ def _prefetch_mtp_group_query( q_tile = q_elem // fx.Int32(4) q_chunks = [] for qwi in range_constexpr(Q_CHUNKS_PER_LANE): - q_chunks.append(_copy_load(q_tiles, q_tile + fx.Int32(qwi), q_copy_atom, q_reg)) + q_chunks.append(copy_load(q_tiles, q_tile + fx.Int32(qwi), q_copy_atom, q_reg)) return qi_val, qhi_pos, q_chunks @@ -401,8 +385,8 @@ def _load_kv_scale_scalars(tile_token_offset_i32, phys_block): scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id scale_global_token = tile_token_offset_i32 + scale_stage_token scale_offset = fx.Int32(scale_block_base + scale_global_token) - k_scale_scalar = _copy_load(ks_tiles, scale_offset, scale_copy_atom, scale_reg)[0] - v_scale_scalar = _copy_load(vs_tiles, scale_offset, scale_copy_atom, scale_reg)[0] + k_scale_scalar = copy_load(ks_tiles, scale_offset, scale_copy_atom, scale_reg)[0] + v_scale_scalar = copy_load(vs_tiles, scale_offset, scale_copy_atom, scale_reg)[0] return k_scale_scalar, v_scale_scalar return () @@ -422,7 +406,7 @@ def _load_v_and_scales( else: va_dw_delta = vhead_elem_dw[vhe] + (v_token_in_block >> fx.Int32(2)) va_byte = (v_block_base_dw + fx.Int64(va_dw_delta)) * fx.Int64(4) - v_i64x2 = _load_global_16b(v_global_ptr, va_byte, v_copy_atom, v_reg) + v_i64x2 = load_global_16b(v_global_ptr, va_byte, v_copy_atom, v_reg) vhe_data.append(v_i64x2) v_results.append(vhe_data) @@ -707,15 +691,23 @@ def _divide_buffer(tensor, tile_layout): lane = fx.Int32(gpu.thread_id("x")) def _num_part(batch_idx): - ctxv = _copy_load(context_lens, batch_idx, copy_i32, load_reg)[0] + ctxv = copy_load(context_lens, batch_idx, copy_i32, load_reg)[0] return fx.Int32(arith.ceildivui(ctxv.ir_value(), c_kvg.ir_value())) - def _store(divided, off, val): - fx.memref_store_vec(fx.Vector.from_elements([fx.Int32(val)], dtype=fx.Int32), store_reg) - fx.copy(copy_i32, store_reg, fx.slice(divided, (None, off))) - - _store(work_indptr, 0, 0) - _store(reduce_indptr, 0, 0) + copy_store( + work_indptr, + 0, + copy_i32, + store_reg, + fx.Vector.from_elements([fx.Int32(0)], dtype=fx.Int32), + ) + copy_store( + reduce_indptr, + 0, + copy_i32, + store_reg, + fx.Vector.from_elements([fx.Int32(0)], dtype=fx.Int32), + ) # Lane-strided partition count followed by a wave reduction. b = lane @@ -762,8 +754,8 @@ def _remain_for_cid(cid_val): remain_kv = pages - kvblk_ do_finish = remain_ >= remain_kv # fx bool - qo_start = _copy_load(sq, batch_, copy_i32, load_reg)[0] - qo_end = _copy_load(sq, batch_ + c1, copy_i32, load_reg)[0] + qo_start = copy_load(sq, batch_, copy_i32, load_reg)[0] + qo_end = copy_load(sq, batch_ + c1, copy_i32, load_reg)[0] kv_start = kvbeg_ + kvblk_ # same for both branches f_kv_end = kvend_ # min(kv_start + remain_kv, kvend_) == kvend_ @@ -820,7 +812,13 @@ def _remain_for_cid(cid_val): do_reduce = do_finish & nsplit_pos num_splits = nsplit_ + c1 # Non-reduce writes are overwritten before their slots become visible. - _store(reduce_indptr, grt_ + c1, lri_ + num_splits) + copy_store( + reduce_indptr, + grt_ + c1, + copy_i32, + store_reg, + fx.Vector.from_elements([fx.Int32(lri_ + num_splits)], dtype=fx.Int32), + ) fx.memref_store_vec( fx.Vector.from_elements([qo_start, qo_end], dtype=fx.Int32), store_reg2, @@ -830,12 +828,24 @@ def _remain_for_cid(cid_val): sidx = c0 while sidx < rcount: val = pidx_ - (nsplit_ - sidx) * c_qlen - _store(reduce_partial_map, lri_ + sidx, val) + copy_store( + reduce_partial_map, + lri_ + sidx, + copy_i32, + store_reg, + fx.Vector.from_elements([fx.Int32(val)], dtype=fx.Int32), + ) sidx = sidx + c1 next_num_works = do_finish.select(f_nworks2, s_nworks2) # The last same-value-race write before advancing cid is authoritative. - _store(work_indptr, cid_ + c1, next_num_works) + copy_store( + work_indptr, + cid_ + c1, + copy_i32, + store_reg, + fx.Vector.from_elements([fx.Int32(next_num_works)], dtype=fx.Int32), + ) ( cid_, @@ -873,13 +883,25 @@ def _remain_for_cid(cid_val): it_t = cid while it_t <= c_numcu: - _store(work_indptr, it_t, num_works) + copy_store( + work_indptr, + it_t, + copy_i32, + store_reg, + fx.Vector.from_elements([fx.Int32(num_works)], dtype=fx.Int32), + ) it_t = it_t + c1 c_rip_size = c_nb + c1 # reduce_indptr length = num_batches + 1 it_r = global_reduce_tile_idx while it_r < c_rip_size: - _store(reduce_indptr, it_r, last_reduce_indptr) + copy_store( + reduce_indptr, + it_r, + copy_i32, + store_reg, + fx.Vector.from_elements([fx.Int32(last_reduce_indptr)], dtype=fx.Int32), + ) it_r = it_r + c1 @flyc.jit @@ -1054,7 +1076,7 @@ def pa_decode_metadata_kernel( warp_id = tid >> fx.Int32(6) def _divide_addr(addr, dtype, width): - ptr = _global_pointer_from_addr(addr, dtype, alignment=width * dtype.width // 8) + ptr = global_pointer_from_addr(addr, dtype, alignment=width * dtype.width // 8) flat = fx.make_view(ptr, fx.make_layout(_FLAT_BUFFER_ELEMENTS, 1)) return fx.logical_divide( fx.rocdl.make_buffer_tensor(flat), @@ -1092,8 +1114,8 @@ def _divide_addr(addr, dtype, width): partial_out_reg = fx.make_rmem_tensor(4, fx.Float32) partial_lse_reg = fx.make_rmem_tensor(1, fx.Float32) - k_global_ptr = _global_pointer_from_addr(key_cache_ptr, fx.Uint8, alignment=16) - v_global_ptr = _global_pointer_from_addr(value_cache_ptr, fx.Uint8, alignment=16) + k_global_ptr = global_pointer_from_addr(key_cache_ptr, fx.Uint8, alignment=16) + v_global_ptr = global_pointer_from_addr(value_cache_ptr, fx.Uint8, alignment=16) global_copy_16b = fx.make_copy_atom(fx.UniversalCopy128b(), fx.Uint8) global_reg_16b = fx.make_rmem_tensor(16, fx.Uint8) @@ -1101,8 +1123,8 @@ def _divide_addr(addr, dtype, width): k_scale_val = fx.Float32(1.0) v_scale_val = fx.Float32(1.0) else: - k_scale_val = _copy_load(key_scales, 0, copy_f32, scale_reg)[0] - v_scale_val = _copy_load(value_scales, 0, copy_f32, scale_reg)[0] + k_scale_val = copy_load(key_scales, 0, copy_f32, scale_reg)[0] + v_scale_val = copy_load(value_scales, 0, copy_f32, scale_reg)[0] lds = fx.SharedAllocator().allocate(SharedStorage).peek() logits_base = lds.buf.ptr @@ -1119,30 +1141,30 @@ def _divide_addr(addr, dtype, width): qkhe_loop=_QKHELOOP, ) - work_start = _copy_load(work_indptr, cu_id, copy_i32, i32_reg)[0] - work_end = _copy_load(work_indptr, cu_id + fx.Int32(1), copy_i32, i32_reg)[0] + work_start = copy_load(work_indptr, cu_id, copy_i32, i32_reg)[0] + work_end = copy_load(work_indptr, cu_id + fx.Int32(1), copy_i32, i32_reg)[0] for work_idx in range(work_start, work_end, fx.Int32(1)): # Two naturally aligned dwordx4 loads cover the eight work fields. info_tile = work_idx * fx.Int32(2) - wi_lo_v = _copy_load(work_info, info_tile, copy_i32x4, i32x4_reg) + wi_lo_v = copy_load(work_info, info_tile, copy_i32x4, i32x4_reg) batch_idx = wi_lo_v[0] partial_idx = wi_lo_v[1] qo_start = wi_lo_v[2] - wi_hi_v = _copy_load(work_info, info_tile + fx.Int32(1), copy_i32x4, i32x4_reg) + wi_hi_v = copy_load(work_info, info_tile + fx.Int32(1), copy_i32x4, i32x4_reg) kv_start = wi_hi_v[0] kv_end = wi_hi_v[1] q_head_range = wi_hi_v[3] # Convert cumulative partition indices to sequence-local page indices. - kv_part_base = _copy_load(partition_indptr, batch_idx, copy_i32, i32_reg)[0] - kv_page_base = _copy_load(kv_indptr, batch_idx, copy_i32, i32_reg)[0] + kv_part_base = copy_load(partition_indptr, batch_idx, copy_i32, i32_reg)[0] + kv_page_base = copy_load(kv_indptr, batch_idx, copy_i32, i32_reg)[0] local_part_start = kv_start - kv_part_base q_head_start = q_head_range & fx.Int32(0xFFFF) kv_h = udiv_const(q_head_start, query_group_size) - context_len = _copy_load(context_lengths, batch_idx, copy_i32, i32_reg)[0] + context_len = copy_load(context_lengths, batch_idx, copy_i32, i32_reg)[0] _k_head_off = kv_h * stride_k_head _v_head_off = kv_h * stride_v_head @@ -1260,7 +1282,7 @@ def _unpack_states(flat): lp = local_part_start + rel_part partition_start = lp * fx.Int32(KV_COMPUTE_BLOCK) - phys_block = _copy_load( + phys_block = copy_load( kv_page_indices, kv_page_base + udiv_const(lp, parts_per_block), copy_i32, @@ -1519,11 +1541,11 @@ def _flat_divide(tensor): reg_f32_vec = fx.make_rmem_tensor(vec_width, fx.Float32) reg_output = fx.make_rmem_tensor(vec_width, output_dtype) - lo = _copy_load(reduce_indptr, g, copy_i32, reg_i32)[0] - hi = _copy_load(reduce_indptr, g + fx.Int32(1), copy_i32, reg_i32)[0] + lo = copy_load(reduce_indptr, g, copy_i32, reg_i32)[0] + hi = copy_load(reduce_indptr, g + fx.Int32(1), copy_i32, reg_i32)[0] if hi > lo: - qo_start = _copy_load(reduce_final_map, g * fx.Int32(2), copy_i32, reg_i32)[0] + qo_start = copy_load(reduce_final_map, g * fx.Int32(2), copy_i32, reg_i32)[0] out_row = qo_start + qi for slot, st in range( @@ -1539,9 +1561,9 @@ def _flat_divide(tensor): m = fx.Float32(st[0]) denom = fx.Float32(st[1]) acc = fx.Vector(st[2]) - prow = _copy_load(reduce_partial_map, slot, copy_i32, reg_i32)[0] + qi - lse = _copy_load(partial_lse, prow * stride_pl_row + qhead, copy_f32, reg_f32)[0] - v = _copy_load( + prow = copy_load(reduce_partial_map, slot, copy_i32, reg_i32)[0] + qi + lse = copy_load(partial_lse, prow * stride_pl_row + qhead, copy_f32, reg_f32)[0] + v = copy_load( partial_output, prow * stride_po_row + qhead * fx.Int32(head_dim) + tid * fx.Int32(vec_width), copy_f32_vec, diff --git a/kernels/common/utils.py b/kernels/common/utils.py index e16138100..6628b9a66 100644 --- a/kernels/common/utils.py +++ b/kernels/common/utils.py @@ -15,6 +15,31 @@ from kernels.common.mem_ops import global_ptr_from_addr as global_ptr_from_addr +def global_pointer_from_addr(addr, dtype, *, alignment: int): + ptr_type = fx.PointerType.get( + elem_ty=dtype.ir_type, + address_space=fx.AddressSpace.Global, + alignment=alignment, + ) + return fx.inttoptr(ptr_type, addr) + + +def copy_load(source, offset, copy_atom, register): + fx.copy(copy_atom, fx.slice(source, (None, fx.Int32(offset))), register) + return fx.memref_load_vec(register) + + +def copy_store(destination, offset, copy_atom, register, value): + fx.memref_store_vec(value, register) + fx.copy(copy_atom, register, fx.slice(destination, (None, fx.Int32(offset)))) + + +def load_global_16b(global_ptr, byte_offset, copy_atom, register): + source = fx.make_view(global_ptr + byte_offset, fx.make_layout(16, 1)) + fx.copy(copy_atom, source, register) + return fx.memref_load_vec(register).bitcast(fx.Int64) + + def rcp_f32(value): return rocdl.rcp(T.f32, value) From 0a051cec39d301c84565fc8eecaec37f7d2943d2 Mon Sep 17 00:00:00 2001 From: fsx950223 Date: Mon, 27 Jul 2026 07:26:23 +0000 Subject: [PATCH 5/7] refactor(pa): Use FlyDSL autotune persistence Signed-off-by: fsx950223 Co-authored-by: Cursor --- kernels/attention/pa_decode_fp8.py | 10 +- kernels/attention/pa_metadata_grid_tuning.csv | 19 - kernels/attention/pa_metadata_tuning.py | 399 +++++++++++++----- python/flydsl/autotune.py | 8 + scripts/tune_pa_metadata_grid.py | 227 ---------- tests/unit/test_pa_metadata_tuning.py | 115 ++--- 6 files changed, 335 insertions(+), 443 deletions(-) delete mode 100644 kernels/attention/pa_metadata_grid_tuning.csv delete mode 100644 scripts/tune_pa_metadata_grid.py diff --git a/kernels/attention/pa_decode_fp8.py b/kernels/attention/pa_decode_fp8.py index 996d9feef..77b94d3a9 100644 --- a/kernels/attention/pa_decode_fp8.py +++ b/kernels/attention/pa_decode_fp8.py @@ -16,7 +16,6 @@ import torch -from flydsl.runtime.device import get_rocm_arch from kernels.attention.pa_decode_swa import compile_pa_decode_sw, compile_pa_decode_sw_reduce from kernels.attention.pa_decode_tile import pa_decode_tile from kernels.attention.pa_metadata import compile_pa_decode_metadata @@ -60,10 +59,10 @@ def get_pa_metadata( NOTE: the consuming decode kernel must interpret kv_start/kv_end as partition indices accordingly. - Exact shape/device matches are loaded from ``pa_metadata_grid_tuning.csv``. - Missing entries default to ``grid_multiplier=1``. ``per_token_kv`` selects - scale-mode-specific tuning, and ``grid_multiplier`` is the tuner's explicit - candidate override. + Exact shape/device matches are loaded from FlyDSL's persistent Autotuner + cache. Missing entries default to ``grid_multiplier=1``. ``per_token_kv`` + selects scale-mode-specific tuning, and ``grid_multiplier`` is the tuner's + explicit candidate override. Returns a dict with: work_indptr, work_info_flat, reduce_indptr, reduce_final_map, reduce_partial_map, num_sm, partial_output, @@ -80,7 +79,6 @@ def get_pa_metadata( num_blocks = key_cache.shape[0] if grid_multiplier is None and per_token_kv is not None: grid_multiplier = lookup_pa_metadata_grid_multiplier( - arch=get_rocm_arch(), num_cu=props.multi_processor_count, batch_size=batch_size, num_blocks=num_blocks, diff --git a/kernels/attention/pa_metadata_grid_tuning.csv b/kernels/attention/pa_metadata_grid_tuning.csv deleted file mode 100644 index 9389628ec..000000000 --- a/kernels/attention/pa_metadata_grid_tuning.csv +++ /dev/null @@ -1,19 +0,0 @@ -arch,num_cu,batch_size,num_blocks,context_length,query_length,per_token_kv,num_query_heads,num_kv_heads,head_dim,block_size,grid_multiplier,latency_us -gfx942,80,1,1,1024,1,0,16,1,128,1024,1,11.440 -gfx942,80,3,6,1027,4,1,16,1,128,1024,1,21.160 -gfx942,80,8,256,32768,4,1,16,1,128,1024,1,122.200 -gfx942,80,16,64,4096,1,0,16,1,128,1024,1,20.680 -gfx942,80,16,128,8192,4,0,16,1,128,1024,2,58.801 -gfx942,80,16,128,8192,4,1,16,1,128,1024,1,75.000 -gfx942,80,32,256,8192,1,0,16,1,128,1024,3,35.520 -gfx942,80,64,256,4096,1,1,16,1,128,1024,3,38.360 -gfx942,80,64,256,4096,2,0,16,1,128,1024,2,53.600 -gfx942,80,64,256,4096,2,1,16,1,128,1024,2,59.920 -gfx942,80,64,256,4096,3,0,16,1,128,1024,2,71.401 -gfx942,80,64,256,4096,3,1,16,1,128,1024,2,81.800 -gfx942,80,64,512,8192,1,0,16,1,128,1024,3,49.520 -gfx942,80,81,648,8192,1,0,16,1,128,1024,3,65.840 -gfx942,80,81,648,8192,4,0,16,1,128,1024,2,184.720 -gfx942,80,81,648,8192,4,1,16,1,128,1024,1,263.400 -gfx942,80,128,1024,8192,4,1,16,1,128,1024,1,415.640 -gfx942,80,256,512,2048,4,1,16,1,128,1024,1,232.000 diff --git a/kernels/attention/pa_metadata_tuning.py b/kernels/attention/pa_metadata_tuning.py index 5a117315d..3465cccf2 100644 --- a/kernels/attention/pa_metadata_tuning.py +++ b/kernels/attention/pa_metadata_tuning.py @@ -1,95 +1,74 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2025 FlyDSL Project Contributors +"""Runtime lookup and offline tuning for PA metadata grid multipliers. + +Run a fresh search with: + FLYDSL_AUTOTUNE=1 python3 -m kernels.attention.pa_metadata_tuning \ + --shape 81,8192,4,per_token +""" + from __future__ import annotations -import csv -import functools -from pathlib import Path -from typing import Iterable, Mapping +import argparse +import statistics +from collections.abc import Callable, Iterable -PA_METADATA_TUNING_CSV = Path(__file__).with_name("pa_metadata_grid_tuning.csv") +from flydsl.autotune import Autotuner, Config -PA_METADATA_TUNING_FIELDS = ( - "arch", - "num_cu", +PA_METADATA_GRID_KEY = [ "batch_size", "num_blocks", - "context_length", "query_length", "per_token_kv", - "num_query_heads", - "num_kv_heads", - "head_dim", - "block_size", - "grid_multiplier", - "latency_us", -) - -_KEY_FIELDS = ( - "arch", "num_cu", - "batch_size", - "num_blocks", - "query_length", - "per_token_kv", "num_query_heads", "num_kv_heads", "head_dim", "block_size", -) - -_INT_FIELDS = set(_KEY_FIELDS) - {"arch", "per_token_kv"} | { - "context_length", - "grid_multiplier", -} - - -def _normalize_row(row: Mapping) -> dict: - normalized = {field: row[field] for field in PA_METADATA_TUNING_FIELDS} - normalized["arch"] = str(normalized["arch"]) - for field in _INT_FIELDS: - normalized[field] = int(normalized[field]) - per_token_kv = normalized["per_token_kv"] - if isinstance(per_token_kv, str): - if per_token_kv not in {"0", "1"}: - raise ValueError(f"per_token_kv must be 0 or 1, got {per_token_kv!r}") - per_token_kv = per_token_kv == "1" - normalized["per_token_kv"] = bool(per_token_kv) - normalized["latency_us"] = float(normalized["latency_us"]) - if normalized["grid_multiplier"] < 1: - raise ValueError("grid_multiplier must be positive") - return normalized - - -def _row_key(row: Mapping) -> tuple: - return tuple(row[field] for field in _KEY_FIELDS) - - -def read_pa_metadata_tuning_rows(path: Path = PA_METADATA_TUNING_CSV) -> list[dict]: - path = Path(path) - if not path.is_file(): - return [] - with path.open(newline="") as handle: - reader = csv.DictReader(handle) - missing = set(PA_METADATA_TUNING_FIELDS) - set(reader.fieldnames or ()) - if missing: - raise ValueError(f"{path} is missing columns: {sorted(missing)}") - rows = [_normalize_row(row) for row in reader] - keys = [_row_key(row) for row in rows] - if len(keys) != len(set(keys)): - raise ValueError(f"{path} contains duplicate tuning keys") - return rows - - -@functools.lru_cache(maxsize=None) -def _load_pa_metadata_tuning(path: Path) -> dict[tuple, int]: - return {_row_key(row): row["grid_multiplier"] for row in read_pa_metadata_tuning_rows(path)} +] + + +def run_pa_metadata_grid_config( + batch_size, + num_blocks, + query_length, + per_token_kv, + num_cu, + num_query_heads, + num_kv_heads, + head_dim, + block_size, + runner, + grid_multiplier, +): + if runner is None: + raise RuntimeError("runner is required when benchmarking a grid config") + return runner(grid_multiplier) + + +def make_pa_metadata_grid_autotuner( + *, + candidates: Iterable[int] = (1, 2, 3), + warmup: int = 5, + rep: int = 3, + do_bench: Callable | None = None, +) -> Autotuner: + return Autotuner( + fn=run_pa_metadata_grid_config, + configs=[Config(grid_multiplier=int(candidate)) for candidate in candidates], + key=PA_METADATA_GRID_KEY, + warmup=warmup, + rep=rep, + do_bench_fn=do_bench, + ) + + +_RUNTIME_TUNER = make_pa_metadata_grid_autotuner(warmup=0, rep=1) def lookup_pa_metadata_grid_multiplier( *, - arch: str, num_cu: int, batch_size: int, num_blocks: int, @@ -99,41 +78,245 @@ def lookup_pa_metadata_grid_multiplier( num_kv_heads: int, head_dim: int, block_size: int, - path: Path = PA_METADATA_TUNING_CSV, ) -> int | None: - key = ( - str(arch), - int(num_cu), - int(batch_size), - int(num_blocks), - int(query_length), - bool(per_token_kv), - int(num_query_heads), - int(num_kv_heads), - int(head_dim), - int(block_size), + args = ( + batch_size, + num_blocks, + query_length, + per_token_kv, + num_cu, + num_query_heads, + num_kv_heads, + head_dim, + block_size, + None, + ) + get_cached_config = getattr(_RUNTIME_TUNER, "get_cached_config", None) + if get_cached_config is not None: + config = get_cached_config(*args) + else: + config = _RUNTIME_TUNER.cache.get(_RUNTIME_TUNER._make_key(args, {})) + if config is None: + return None + return int(config.kwargs["grid_multiplier"]) + + +def _parse_int_list(value: str) -> list[int]: + values = [int(item) for item in value.split(",")] + if not values or any(item < 1 for item in values): + raise argparse.ArgumentTypeError("expected comma-separated positive integers") + return values + + +def _parse_shape(value: str) -> tuple[int, int, int, bool]: + try: + batch_size, context_length, query_length, quant_mode = value.split(",") + per_token_kv = {"per_tensor": False, "per_token": True}[quant_mode] + parsed = (int(batch_size), int(context_length), int(query_length), per_token_kv) + except (KeyError, TypeError, ValueError) as error: + raise argparse.ArgumentTypeError("shape must be BATCH,CONTEXT,QUERY_LENGTH,per_tensor|per_token") from error + if any(item < 1 for item in parsed[:3]): + raise argparse.ArgumentTypeError("shape dimensions must be positive") + return parsed + + +def _tune_shape( + *, + batch_size: int, + context_length: int, + query_length: int, + per_token_kv: bool, + candidates: list[int], + num_query_heads: int, + num_kv_heads: int, + head_dim: int, + block_size: int, + warmup: int, + iterations: int, + rounds: int, + device, +): + import torch + + from kernels.attention.pa_decode_fp8 import get_pa_metadata, pa_decode_ps_launch + + pages_per_sequence = (context_length + block_size - 1) // block_size + num_blocks = batch_size * pages_per_sequence + num_cu = torch.cuda.get_device_properties(device).multi_processor_count + context_lengths = torch.full((batch_size,), context_length, dtype=torch.int32, device=device) + query = torch.zeros( + (batch_size * query_length, num_query_heads, head_dim), + dtype=torch.bfloat16, + device=device, ) - return _load_pa_metadata_tuning(Path(path)).get(key) - - -def write_pa_metadata_tuning_rows( - rows: Iterable[Mapping], - path: Path = PA_METADATA_TUNING_CSV, -) -> None: - path = Path(path) - merged = {_row_key(row): row for row in read_pa_metadata_tuning_rows(path)} - for row in rows: - normalized = _normalize_row(row) - merged[_row_key(normalized)] = normalized - - path.parent.mkdir(parents=True, exist_ok=True) - temporary = path.with_suffix(f"{path.suffix}.tmp") - with temporary.open("w", newline="") as handle: - writer = csv.DictWriter(handle, fieldnames=PA_METADATA_TUNING_FIELDS) - writer.writeheader() - for key in sorted(merged): - row = dict(merged[key]) - row["per_token_kv"] = int(row["per_token_kv"]) - writer.writerow(row) - temporary.replace(path) - _load_pa_metadata_tuning.cache_clear() + kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * pages_per_sequence + kv_page_indices = torch.arange(num_blocks, dtype=torch.int32, device=device) + key = torch.zeros( + (num_blocks, num_kv_heads, head_dim // 16, block_size, 16), + dtype=torch.float8_e4m3fnuz, + device=device, + ) + value = torch.ones( + (num_blocks, num_kv_heads, block_size // 16, head_dim, 16), + dtype=torch.float8_e4m3fnuz, + device=device, + ) + if per_token_kv: + key_scale = torch.ones((num_blocks, num_kv_heads, block_size), dtype=torch.float32, device=device) + value_scale = torch.ones_like(key_scale) + else: + key_scale = torch.ones((1,), dtype=torch.float32, device=device) + value_scale = torch.ones((1,), dtype=torch.float32, device=device) + + output = torch.empty_like(query) + graphs = {} + metadata_resources = [] + for grid_multiplier in candidates: + metadata = get_pa_metadata( + query, + key, + context_lengths, + kv_indptr, + num_query_heads, + num_kv_heads, + per_token_kv=per_token_kv, + grid_multiplier=grid_multiplier, + ) + metadata_resources.append(metadata) + + def launch(metadata=metadata): + pa_decode_ps_launch( + output, + query, + key, + value, + context_lengths, + kv_page_indices, + kv_indptr, + 1.0 / head_dim**0.5, + key_scale=key_scale, + value_scale=value_scale, + metadata=metadata, + ) + + launch() + torch.cuda.synchronize(device) + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph): + for _ in range(iterations): + launch() + graphs[grid_multiplier] = graph + torch.cuda.synchronize(device) + + active_grid = {"value": None} + timings_us = {} + + def replay_grid(grid_multiplier): + active_grid["value"] = grid_multiplier + graphs[grid_multiplier].replay() + return grid_multiplier + + def graph_bench(call, warmup, rep): + for _ in range(warmup): + call() + torch.cuda.synchronize(device) + samples_ms = [] + for _ in range(rep): + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + call() + end.record() + end.synchronize() + samples_ms.append(start.elapsed_time(end) / iterations) + latency_ms = statistics.median(samples_ms) + timings_us[active_grid["value"]] = latency_ms * 1000.0 + return latency_ms + + tuner = make_pa_metadata_grid_autotuner( + candidates=candidates, + warmup=warmup, + rep=rounds, + do_bench=graph_bench, + ) + tuner_args = ( + batch_size, + num_blocks, + query_length, + per_token_kv, + num_cu, + num_query_heads, + num_kv_heads, + head_dim, + block_size, + replay_grid, + ) + tuner(*tuner_args) + get_cached_config = getattr(tuner, "get_cached_config", None) + if get_cached_config is not None: + best_config = get_cached_config(*tuner_args) + else: + best_config = tuner.cache.get(tuner._make_key(tuner_args, {})) + best_grid = int(best_config.kwargs["grid_multiplier"]) + torch.cuda.synchronize(device) + + if not torch.allclose( + output.float(), + torch.ones_like(output, dtype=torch.float32), + atol=2e-2, + rtol=2e-2, + ): + raise RuntimeError(f"grid_multiplier={best_grid} failed correctness") + return best_grid, timings_us.get(best_grid), tuner.cache_file + + +def main() -> None: + import torch + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--shape", action="append", type=_parse_shape, required=True) + parser.add_argument("--candidates", type=_parse_int_list, default=[1, 2, 3]) + parser.add_argument("--warmup", type=int, default=5) + parser.add_argument("--iterations", type=int, default=100) + parser.add_argument("--rounds", type=int, default=3) + parser.add_argument("--device", type=int, default=0) + parser.add_argument("--num-query-heads", type=int, default=16) + parser.add_argument("--num-kv-heads", type=int, default=1) + parser.add_argument("--head-dim", type=int, default=128) + parser.add_argument("--block-size", type=int, default=1024) + args = parser.parse_args() + + if args.warmup < 0 or args.iterations < 1 or args.rounds < 1: + parser.error("warmup must be non-negative; iterations and rounds must be positive") + + torch.cuda.set_device(args.device) + device = torch.device("cuda", args.device) + cache_file = None + + for batch_size, context_length, query_length, per_token_kv in args.shape: + best_grid, best_latency, cache_file = _tune_shape( + batch_size=batch_size, + context_length=context_length, + query_length=query_length, + per_token_kv=per_token_kv, + candidates=args.candidates, + num_query_heads=args.num_query_heads, + num_kv_heads=args.num_kv_heads, + head_dim=args.head_dim, + block_size=args.block_size, + warmup=args.warmup, + iterations=args.iterations, + rounds=args.rounds, + device=device, + ) + if best_latency is None: + print(f" cached winner: grid={best_grid}") + else: + print(f" winner: grid={best_grid}, {best_latency:.3f} us") + torch.cuda.empty_cache() + + print(f"Autotune cache: {cache_file}") + + +if __name__ == "__main__": + main() diff --git a/python/flydsl/autotune.py b/python/flydsl/autotune.py index 3b4cad77b..8e36c45df 100644 --- a/python/flydsl/autotune.py +++ b/python/flydsl/autotune.py @@ -364,6 +364,14 @@ def _run_config(self, config, args, kwargs): self._reset_tensors(args, merged) return self._run_with_hints(config.compiler_opts(), args, merged) + def get_cached_config(self, *args, **kwargs): + """Return the persisted/in-memory config for these arguments, if any.""" + return self.cache.get(self._make_key(args, kwargs)) + + @property + def cache_file(self): + return self._cache_file + def __call__(self, *args, **kwargs): key = self._make_key(args, kwargs) force = _tuning_enabled() diff --git a/scripts/tune_pa_metadata_grid.py b/scripts/tune_pa_metadata_grid.py deleted file mode 100644 index 3a4f890ee..000000000 --- a/scripts/tune_pa_metadata_grid.py +++ /dev/null @@ -1,227 +0,0 @@ -#!/usr/bin/env python3 -# SPDX-License-Identifier: Apache-2.0 -# Copyright (c) 2025 FlyDSL Project Contributors - -"""Tune PA metadata grid multipliers and persist winners to CSV. - -Example: - python3 scripts/tune_pa_metadata_grid.py \ - --shape 16,4096,1,per_tensor \ - --shape 81,8192,4,per_token -""" - -from __future__ import annotations - -import argparse -import statistics -import sys -from pathlib import Path - -import torch - -sys.path.insert(0, str(Path(__file__).resolve().parents[1])) - -from flydsl.runtime.device import get_rocm_arch -from kernels.attention.pa_decode_fp8 import get_pa_metadata, pa_decode_ps_launch -from kernels.attention.pa_metadata_tuning import ( - PA_METADATA_TUNING_CSV, - write_pa_metadata_tuning_rows, -) - - -def _parse_int_list(value: str) -> list[int]: - values = [int(item) for item in value.split(",")] - if not values or any(item < 1 for item in values): - raise argparse.ArgumentTypeError("expected comma-separated positive integers") - return values - - -def _parse_shape(value: str) -> tuple[int, int, int, bool]: - try: - batch_size, context_length, query_length, quant_mode = value.split(",") - per_token_kv = {"per_tensor": False, "per_token": True}[quant_mode] - parsed = (int(batch_size), int(context_length), int(query_length), per_token_kv) - except (KeyError, TypeError, ValueError) as error: - raise argparse.ArgumentTypeError("shape must be BATCH,CONTEXT,QUERY_LENGTH,per_tensor|per_token") from error - if any(item < 1 for item in parsed[:3]): - raise argparse.ArgumentTypeError("shape dimensions must be positive") - return parsed - - -def _measure_candidate( - *, - batch_size: int, - context_length: int, - query_length: int, - per_token_kv: bool, - grid_multiplier: int, - num_query_heads: int, - num_kv_heads: int, - head_dim: int, - block_size: int, - warmup: int, - iterations: int, - rounds: int, - device: torch.device, -) -> tuple[float, int]: - pages_per_sequence = (context_length + block_size - 1) // block_size - num_blocks = batch_size * pages_per_sequence - context_lengths = torch.full((batch_size,), context_length, dtype=torch.int32, device=device) - query = torch.zeros( - (batch_size * query_length, num_query_heads, head_dim), - dtype=torch.bfloat16, - device=device, - ) - kv_indptr = torch.arange(batch_size + 1, dtype=torch.int32, device=device) * pages_per_sequence - kv_page_indices = torch.arange(num_blocks, dtype=torch.int32, device=device) - key = torch.zeros( - (num_blocks, num_kv_heads, head_dim // 16, block_size, 16), - dtype=torch.float8_e4m3fnuz, - device=device, - ) - value = torch.ones( - (num_blocks, num_kv_heads, block_size // 16, head_dim, 16), - dtype=torch.float8_e4m3fnuz, - device=device, - ) - if per_token_kv: - key_scale = torch.ones((num_blocks, num_kv_heads, block_size), dtype=torch.float32, device=device) - value_scale = torch.ones_like(key_scale) - else: - key_scale = torch.ones((1,), dtype=torch.float32, device=device) - value_scale = torch.ones((1,), dtype=torch.float32, device=device) - - metadata = get_pa_metadata( - query, - key, - context_lengths, - kv_indptr, - num_query_heads, - num_kv_heads, - per_token_kv=per_token_kv, - grid_multiplier=grid_multiplier, - ) - output = torch.empty_like(query) - - def launch(): - pa_decode_ps_launch( - output, - query, - key, - value, - context_lengths, - kv_page_indices, - kv_indptr, - 1.0 / head_dim**0.5, - key_scale=key_scale, - value_scale=value_scale, - metadata=metadata, - ) - - for _ in range(warmup): - launch() - torch.cuda.synchronize(device) - - graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): - for _ in range(iterations): - launch() - torch.cuda.synchronize(device) - - timings = [] - for _ in range(rounds): - start = torch.cuda.Event(enable_timing=True) - end = torch.cuda.Event(enable_timing=True) - start.record() - graph.replay() - end.record() - end.synchronize() - timings.append(start.elapsed_time(end) * 1000.0 / iterations) - - if not torch.allclose( - output.float(), - torch.ones_like(output, dtype=torch.float32), - atol=2e-2, - rtol=2e-2, - ): - raise RuntimeError(f"grid_multiplier={grid_multiplier} failed correctness") - return statistics.median(timings), num_blocks - - -def main() -> None: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--shape", action="append", type=_parse_shape, required=True) - parser.add_argument("--candidates", type=_parse_int_list, default=[1, 2, 3]) - parser.add_argument("--warmup", type=int, default=5) - parser.add_argument("--iterations", type=int, default=100) - parser.add_argument("--rounds", type=int, default=3) - parser.add_argument("--device", type=int, default=0) - parser.add_argument("--num-query-heads", type=int, default=16) - parser.add_argument("--num-kv-heads", type=int, default=1) - parser.add_argument("--head-dim", type=int, default=128) - parser.add_argument("--block-size", type=int, default=1024) - parser.add_argument("--output", type=Path, default=PA_METADATA_TUNING_CSV) - args = parser.parse_args() - - if args.warmup < 0 or args.iterations < 1 or args.rounds < 1: - parser.error("warmup must be non-negative; iterations and rounds must be positive") - - torch.cuda.set_device(args.device) - device = torch.device("cuda", args.device) - arch = get_rocm_arch() - num_cu = torch.cuda.get_device_properties(device).multi_processor_count - tuned_rows = [] - - for batch_size, context_length, query_length, per_token_kv in args.shape: - timings = {} - num_blocks = 0 - for grid_multiplier in args.candidates: - latency_us, num_blocks = _measure_candidate( - batch_size=batch_size, - context_length=context_length, - query_length=query_length, - per_token_kv=per_token_kv, - grid_multiplier=grid_multiplier, - num_query_heads=args.num_query_heads, - num_kv_heads=args.num_kv_heads, - head_dim=args.head_dim, - block_size=args.block_size, - warmup=args.warmup, - iterations=args.iterations, - rounds=args.rounds, - device=device, - ) - timings[grid_multiplier] = latency_us - print( - f"B={batch_size} C={context_length} Q={query_length} " - f"per_token={per_token_kv} grid={grid_multiplier}: {latency_us:.3f} us" - ) - - best_grid = min(timings, key=timings.get) - best_latency = timings[best_grid] - print(f" winner: grid={best_grid}, {best_latency:.3f} us") - tuned_rows.append( - { - "arch": arch, - "num_cu": num_cu, - "batch_size": batch_size, - "num_blocks": num_blocks, - "context_length": context_length, - "query_length": query_length, - "per_token_kv": per_token_kv, - "num_query_heads": args.num_query_heads, - "num_kv_heads": args.num_kv_heads, - "head_dim": args.head_dim, - "block_size": args.block_size, - "grid_multiplier": best_grid, - "latency_us": best_latency, - } - ) - torch.cuda.empty_cache() - - write_pa_metadata_tuning_rows(tuned_rows, args.output) - print(f"Wrote {len(tuned_rows)} tuned configuration(s) to {args.output}") - - -if __name__ == "__main__": - main() diff --git a/tests/unit/test_pa_metadata_tuning.py b/tests/unit/test_pa_metadata_tuning.py index d9b175cf5..b8ff31e5f 100644 --- a/tests/unit/test_pa_metadata_tuning.py +++ b/tests/unit/test_pa_metadata_tuning.py @@ -1,108 +1,57 @@ # SPDX-License-Identifier: Apache-2.0 # Copyright (c) 2025 FlyDSL Project Contributors -import csv +import importlib -import pytest +import kernels.attention.pa_metadata_tuning as tuning -from kernels.attention.pa_metadata_tuning import ( - PA_METADATA_TUNING_CSV, - lookup_pa_metadata_grid_multiplier, - read_pa_metadata_tuning_rows, - write_pa_metadata_tuning_rows, -) - -def _row(**overrides): - row = { - "arch": "gfx942", +def _lookup(**overrides): + kwargs = { "num_cu": 80, "batch_size": 81, "num_blocks": 648, - "context_length": 8192, "query_length": 4, "per_token_kv": True, "num_query_heads": 16, "num_kv_heads": 1, "head_dim": 128, "block_size": 1024, - "grid_multiplier": 1, - "latency_us": 263.4, } - row.update(overrides) - return row + kwargs.update(overrides) + return tuning.lookup_pa_metadata_grid_multiplier(**kwargs) -def test_default_pa_metadata_tuning_table(): - rows = read_pa_metadata_tuning_rows() - assert rows - assert ( - lookup_pa_metadata_grid_multiplier( - arch="gfx942", - num_cu=80, - batch_size=81, - num_blocks=648, - query_length=4, - per_token_kv=True, - num_query_heads=16, - num_kv_heads=1, - head_dim=128, - block_size=1024, - ) - == 1 - ) - assert PA_METADATA_TUNING_CSV.is_file() +def test_pa_metadata_tuner_persists_and_runtime_reads(tmp_path, monkeypatch): + monkeypatch.setenv("FLYDSL_AUTOTUNE_CACHE_DIR", str(tmp_path)) + importlib.reload(tuning) + active_grid = {"value": None} + def runner(grid_multiplier): + active_grid["value"] = grid_multiplier + return grid_multiplier -def test_pa_metadata_tuning_csv_roundtrip_and_upsert(tmp_path): - path = tmp_path / "pa_tuning.csv" - write_pa_metadata_tuning_rows([_row()], path) - write_pa_metadata_tuning_rows( - [_row(grid_multiplier=2, latency_us=250.0), _row(batch_size=32, num_blocks=256)], path - ) + def bench(call, warmup, rep): + call() + return {1: 2.0, 2: 1.0}[active_grid["value"]] - rows = read_pa_metadata_tuning_rows(path) - assert len(rows) == 2 - assert ( - lookup_pa_metadata_grid_multiplier( - arch="gfx942", - num_cu=80, - batch_size=81, - num_blocks=648, - query_length=4, - per_token_kv=True, - num_query_heads=16, - num_kv_heads=1, - head_dim=128, - block_size=1024, - path=path, - ) - == 2 - ) - assert ( - lookup_pa_metadata_grid_multiplier( - arch="gfx950", - num_cu=80, - batch_size=81, - num_blocks=648, - query_length=4, - per_token_kv=True, - num_query_heads=16, - num_kv_heads=1, - head_dim=128, - block_size=1024, - path=path, - ) - is None + tuner = tuning.make_pa_metadata_grid_autotuner( + candidates=[1, 2], + warmup=0, + rep=1, + do_bench=bench, ) + tuner_args = (81, 648, 4, True, 80, 16, 1, 128, 1024, runner) + tuner(*tuner_args) + + assert tuner.get_cached_config(*tuner_args).kwargs["grid_multiplier"] == 2 + assert tuner.cache_file == tmp_path / "run_pa_metadata_grid_config.json" + assert tuner.cache_file.is_file() + importlib.reload(tuning) + assert _lookup() == 2 + assert _lookup(batch_size=32) is None -def test_pa_metadata_tuning_rejects_invalid_boolean(tmp_path): - path = tmp_path / "invalid.csv" - with path.open("w", newline="") as handle: - writer = csv.DictWriter(handle, fieldnames=_row()) - writer.writeheader() - writer.writerow(_row(per_token_kv="true")) - with pytest.raises(ValueError, match="per_token_kv"): - read_pa_metadata_tuning_rows(path) +def test_pa_metadata_tuner_key_excludes_context_length(): + assert "context_length" not in tuning.PA_METADATA_GRID_KEY From af1967e0bda2969c5a9194286ace90c56e81a00b Mon Sep 17 00:00:00 2001 From: fsx950223 Date: Tue, 28 Jul 2026 06:24:17 +0000 Subject: [PATCH 6/7] refactor(pa): Simplify metadata query and scale flow Signed-off-by: fsx950223 Co-authored-by: Cursor --- kernels/attention/pa_common.py | 22 +++++ kernels/attention/pa_metadata.py | 140 ++++++++++++------------------- 2 files changed, 77 insertions(+), 85 deletions(-) diff --git a/kernels/attention/pa_common.py b/kernels/attention/pa_common.py index fef339dc4..28f455cb1 100644 --- a/kernels/attention/pa_common.py +++ b/kernels/attention/pa_common.py @@ -11,6 +11,7 @@ import flydsl.expr as fx from flydsl.expr import arith, const_expr, range_constexpr from kernels.common import buffer_ops +from kernels.common.utils import copy_load # PA Q-tiling constants (identical across all PA decode kernels). MFMA_N = 16 @@ -53,3 +54,24 @@ def _prefetch_q_chunks( ) ) return q_chunks + + +@flyc.jit +def _prefetch_q_chunks_tile( + q_tiles, + q_copy_atom, + q_reg, + q_base, + lane16id, + *, + q_lanes_per_head, +): + q_load_lane = lane16id + if const_expr(q_lanes_per_head < MFMA_N): + q_load_lane = (lane16id < fx.Int32(q_lanes_per_head)).select(lane16id, fx.Int32(0)) + q_elem = q_base + q_load_lane * fx.Int32(Q_ELEMS_PER_LANE) + q_tile = q_elem // fx.Int32(4) + q_chunks = [] + for qwi in range_constexpr(Q_CHUNKS_PER_LANE): + q_chunks.append(copy_load(q_tiles, q_tile + fx.Int32(qwi), q_copy_atom, q_reg)) + return q_chunks diff --git a/kernels/attention/pa_metadata.py b/kernels/attention/pa_metadata.py index d1dc70470..5df7da4ec 100644 --- a/kernels/attention/pa_metadata.py +++ b/kernels/attention/pa_metadata.py @@ -19,11 +19,10 @@ import flydsl.compiler as flyc import flydsl.expr as fx -from flydsl.expr import arith, const_expr, gpu, range_constexpr, rocdl +from flydsl.expr import arith, as_ir_value, const_expr, gpu, range_constexpr, rocdl from flydsl.expr import math as fmath from flydsl.expr.typing import T -from flydsl.runtime.device import get_rocm_arch -from kernels.attention.pa_common import _compute_block_base_dw_i64 +from kernels.attention.pa_common import _compute_block_base_dw_i64, _prefetch_q_chunks_tile from kernels.common import dpp_utils from kernels.common.kernels_common import get_warp_size from kernels.common.tensor_shim import _run_compiled @@ -101,8 +100,6 @@ def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = Q_ELEMS_PER_LANE = 8 -Q_CHUNKS_PER_LANE = Q_ELEMS_PER_LANE // 4 - PROB_ROW_STRIDE_BYTES = 40 # 32 data + 8 padding -> 0 bank conflict LDS_LOGITS_BYTES = NUM_WARPS * 4 * MFMA_N * PROB_ROW_STRIDE_BYTES # 10240 @@ -289,14 +286,14 @@ def _prefetch_mtp_group_query( ) q_row = batch_idx * fx.Int32(query_length) + qi_for_q q_base = q_row * stride_q_seq + (kv_h * fx.Int32(query_group_size) + local_qhead_idx_for_q) * stride_q_head - q_load_lane = lane16id - if const_expr(q_lanes_per_head < MFMA_N): - q_load_lane = (lane16id < fx.Int32(q_lanes_per_head)).select(lane16id, fx.Int32(0)) - q_elem = q_base + q_load_lane * fx.Int32(Q_ELEMS_PER_LANE) - q_tile = q_elem // fx.Int32(4) - q_chunks = [] - for qwi in range_constexpr(Q_CHUNKS_PER_LANE): - q_chunks.append(copy_load(q_tiles, q_tile + fx.Int32(qwi), q_copy_atom, q_reg)) + q_chunks = _prefetch_q_chunks_tile( + q_tiles, + q_copy_atom, + q_reg, + q_base, + lane16id, + q_lanes_per_head=q_lanes_per_head, + ) return qi_val, qhi_pos, q_chunks @@ -410,6 +407,8 @@ def _load_v_and_scales( vhe_data.append(v_i64x2) v_results.append(vhe_data) + k_scale_vecs = [] + v_scale_vecs = [] if const_expr(per_token_kv): scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id k_scale_scalar, v_scale_scalar = scale_scalars @@ -423,8 +422,6 @@ def _load_v_and_scales( ) rocdl.sched_barrier(0) gpu.barrier() - k_scale_vecs = [] - v_scale_vecs = [] for td in range_constexpr(TLOOP): scale_row_base = kv_tok_thread_base + fx.Int32(td * MFMA_N) k_scale_vecs.append( @@ -436,9 +433,8 @@ def _load_v_and_scales( result_type=fx.Vector.make_type(4, fx.Float32), ) ) - return v_results, k_scale_vecs, v_scale_vecs - return v_results + return v_results, k_scale_vecs, v_scale_vecs def _store_vmax_warp(partition_start, *, seq_end=None, v_scale_vecs=None): if const_expr(per_token_kv): @@ -470,11 +466,8 @@ def _qk_and_intra_softmax( causal_bound, query_scale_lane, *, - preloaded_scales=None, + k_scale_vecs=None, ): - if const_expr(per_token_kv): - k_scale_vecs, v_scale_vecs = preloaded_scales - query_scale = fx.Float32(query_scale_lane * softmax_scale) d_out = [] for td in range_constexpr(TLOOP): @@ -502,8 +495,6 @@ def _qk_and_intra_softmax( softmax_base + sm_max_off, ) - if const_expr(per_token_kv): - return d_out, v_scale_vecs return d_out def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs, *, v_normalization=None): @@ -955,7 +946,7 @@ def get_pa_metadata_v1( if num_cu is None: num_cu = torch.cuda.get_device_properties(dev).multi_processor_count num_batches = context_lens.shape[0] - warp_size = get_warp_size(get_rocm_arch()) + warp_size = get_warp_size() compiled = compile_pa_metadata_v1( num_cu=num_cu, @@ -1013,13 +1004,13 @@ def compile_pa_decode_metadata( raise ValueError(f"compile_pa_decode_metadata only supports block_size={KV_BLOCK_SIZE}, got {block_size}") if head_dim % QKHE_PER_FETCH != 0 or head_dim % (MFMA_N * NUM_WARPS) != 0 or head_dim % Q_ELEMS_PER_LANE != 0: raise ValueError(f"Unsupported head_dim={head_dim}; must be a multiple of {MFMA_N * NUM_WARPS}.") - _QKHELOOP = head_dim // QKHE_PER_FETCH - _VHELOOP = head_dim // MFMA_N // NUM_WARPS - _Q_LANES_PER_HEAD = head_dim // Q_ELEMS_PER_LANE + QKHELOOP = head_dim // QKHE_PER_FETCH + VHELOOP = head_dim // MFMA_N // NUM_WARPS + Q_LANES_PER_HEAD = head_dim // Q_ELEMS_PER_LANE if query_input_dtype not in ("bf16", "f16"): raise ValueError(f"`compile_pa_decode_metadata` only supports bf16/f16 queries, got {query_input_dtype!r}") - _QUERY_DTYPE = fx.BFloat16 if query_input_dtype == "bf16" else fx.Float16 - _OUTPUT_DTYPE = _PA_OUTPUT_DTYPES[output_dtype_str] + QUERY_DTYPE = fx.BFloat16 if query_input_dtype == "bf16" else fx.Float16 + OUTPUT_DTYPE = _PA_OUTPUT_DTYPES[output_dtype_str] if softmax_scale is None: softmax_scale = 1.0 / (head_dim**0.5) softmax_scale = float(softmax_scale) @@ -1083,8 +1074,8 @@ def _divide_addr(addr, dtype, width): fx.make_layout(width, 1), ) - q_tiles = _divide_addr(query_ptr, _QUERY_DTYPE, 4) - out_tiles = _divide_addr(out_ptr, _OUTPUT_DTYPE, 4) + q_tiles = _divide_addr(query_ptr, QUERY_DTYPE, 4) + out_tiles = _divide_addr(out_ptr, OUTPUT_DTYPE, 4) partial_out_tiles = _divide_addr(partial_out_ptr, fx.Float32, 4) partial_lse = _divide_addr(partial_lse_ptr, fx.Float32, 1) context_lengths = _divide_addr(context_lengths_ptr, fx.Int32, 1) @@ -1099,18 +1090,18 @@ def _divide_addr(addr, dtype, width): copy_i32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Int32) copy_i32x4 = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Int32) copy_f32 = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Float32) - copy_q = fx.make_copy_atom(fx.rocdl.BufferCopy64b(), _QUERY_DTYPE) + copy_q = fx.make_copy_atom(fx.rocdl.BufferCopy64b(), QUERY_DTYPE) copy_out = fx.make_copy_atom( - fx.rocdl.BufferCopy128b() if _OUTPUT_DTYPE.width == 32 else fx.rocdl.BufferCopy64b(), - _OUTPUT_DTYPE, + fx.rocdl.BufferCopy128b() if OUTPUT_DTYPE.width == 32 else fx.rocdl.BufferCopy64b(), + OUTPUT_DTYPE, ) copy_f32x4 = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Float32) i32_reg = fx.make_rmem_tensor(1, fx.Int32) i32x4_reg = fx.make_rmem_tensor(4, fx.Int32) scale_reg = fx.make_rmem_tensor(1, fx.Float32) - q_reg = fx.make_rmem_tensor(4, _QUERY_DTYPE) - out_reg = fx.make_rmem_tensor(4, _OUTPUT_DTYPE) + q_reg = fx.make_rmem_tensor(4, QUERY_DTYPE) + out_reg = fx.make_rmem_tensor(4, OUTPUT_DTYPE) partial_out_reg = fx.make_rmem_tensor(4, fx.Float32) partial_lse_reg = fx.make_rmem_tensor(1, fx.Float32) @@ -1138,7 +1129,7 @@ def _divide_addr(addr, dtype, width): warp_id, lane16id, rowid, - qkhe_loop=_QKHELOOP, + qkhe_loop=QKHELOOP, ) work_start = copy_load(work_indptr, cu_id, copy_i32, i32_reg)[0] @@ -1200,10 +1191,6 @@ def _divide_addr(addr, dtype, width): head_dim=head_dim, ) - # Q stays in registers while K/V are reused across all MTP groups. - def _unwrap(v): - return v.ir_value() if hasattr(v, "ir_value") else v - # Negative partial_idx writes final output; split rows reserve the first QL slots. _is_direct = partial_idx < fx.Int32(0) _po_row_base = partial_idx + fx.Int32(query_length) @@ -1233,7 +1220,7 @@ def _unwrap(v): mtp_group_idx=_mtp_g, query_length=query_length, query_group_size=query_group_size, - q_lanes_per_head=_Q_LANES_PER_HEAD, + q_lanes_per_head=Q_LANES_PER_HEAD, ) _qfrags, _qscale = _finish_q_fragments( logits_base, @@ -1254,10 +1241,10 @@ def _unwrap(v): context_len + fx.Int32(1 - query_length) + qi_per_mtp[_mtp_g] for _mtp_g in range(_mtp_groups) ] - state_width = 2 + _VHELOOP + state_width = 2 + VHELOOP def _pack_states(states): - return [_unwrap(value) for state in states for value in state] + return [as_ir_value(value) for state in states for value in state] def _unpack_states(flat): return [tuple(flat[state_width * i : state_width * (i + 1)]) for i in range(_mtp_groups)] @@ -1265,7 +1252,7 @@ def _unpack_states(flat): init_states = [ tuple( [fx.Float32(float("-inf")), fx.Float32(0.0)] - + [fx.Vector.filled(4, 0.0, fx.Float32) for _ in range_constexpr(_VHELOOP)] + + [fx.Vector.filled(4, 0.0, fx.Float32) for _ in range_constexpr(VHELOOP)] ) for _ in range(_mtp_groups) ] @@ -1302,51 +1289,34 @@ def _unpack_states(flat): _k_tok_thread_base, _c_tok_stride_dw, _k_he_off_dw, - qkhe_loop=_QKHELOOP, + qkhe_loop=QKHELOOP, ) - k_ops = unflatten_k(k_flat, qkhe_loop=_QKHELOOP) + k_ops = unflatten_k(k_flat, qkhe_loop=QKHELOOP) v_base = _compute_block_base_dw_i64(phys_block, stride_v_block, _v_head_off) - if const_expr(per_token_kv): - v_ops, k_scale_vecs, v_scale_vecs = _load_v_and_scales( - v_base, - tile_token_offset, - scale_scalars, - ) - else: - v_ops = _load_v_and_scales( - v_base, - tile_token_offset, - scale_scalars, - ) + v_ops, k_scale_vecs, v_scale_vecs = _load_v_and_scales( + v_base, + tile_token_offset, + scale_scalars, + ) v_normalization = None new_states = [] for _mtp_g in range_constexpr(_mtp_groups): state = cur_states[_mtp_g] rmax, rsum = state[0], state[1] - outs = [state[2 + vhe] for vhe in range_constexpr(_VHELOOP)] - - if const_expr(per_token_kv): - d_out, v_scales = _qk_and_intra_softmax( - k_ops, - partition_start, - q_frags_per_mtp[_mtp_g], - causal_bound_per_mtp[_mtp_g], - query_scale_lane=qscale_per_mtp[_mtp_g], - preloaded_scales=(k_scale_vecs, v_scale_vecs), - ) - else: - d_out = _qk_and_intra_softmax( - k_ops, - partition_start, - q_frags_per_mtp[_mtp_g], - causal_bound_per_mtp[_mtp_g], - query_scale_lane=qscale_per_mtp[_mtp_g], - ) - v_scales = None + outs = [state[2 + vhe] for vhe in range_constexpr(VHELOOP)] + + d_out = _qk_and_intra_softmax( + k_ops, + partition_start, + q_frags_per_mtp[_mtp_g], + causal_bound_per_mtp[_mtp_g], + query_scale_lane=qscale_per_mtp[_mtp_g], + k_scale_vecs=k_scale_vecs, + ) if const_expr(per_token_kv and _mtp_g == 0): - _store_vmax_warp(partition_start, seq_end=context_len, v_scale_vecs=v_scales) + _store_vmax_warp(partition_start, seq_end=context_len, v_scale_vecs=v_scale_vecs) gpu.barrier() rmax, rsum, outs, v_correction, normalized_v_scales = _cross_warp_softmax_and_prob_pack( @@ -1354,7 +1324,7 @@ def _unpack_states(flat): rmax, rsum, outs, - v_scales, + v_scale_vecs, v_normalization=v_normalization, ) if const_expr(per_token_kv and _mtp_g == 0): @@ -1370,7 +1340,7 @@ def _unpack_states(flat): for _mtp_g in range_constexpr(_mtp_groups): final_state = final_states[_mtp_g] rmax_raw, rsum_raw = final_state[0], final_state[1] - outs_raw = [final_state[2 + vhe] for vhe in range_constexpr(_VHELOOP)] + outs_raw = [final_state[2 + vhe] for vhe in range_constexpr(VHELOOP)] running_max = fx.Float32(rmax_raw) running_sum = fx.Float32(rsum_raw) outs = [fx.Vector(out_raw) for out_raw in outs_raw] @@ -1381,14 +1351,14 @@ def _unpack_states(flat): if _is_direct: out_row = qo_start + qi_val_mg - for vhe in range_constexpr(_VHELOOP): + for vhe in range_constexpr(VHELOOP): hs_base = fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + rowid * fx.Int32(4) out_off = out_row * stride_out_seq + qhead * stride_out_head + hs_base - fx.memref_store_vec(outelems_norm[vhe].to(_OUTPUT_DTYPE), out_reg) + fx.memref_store_vec(outelems_norm[vhe].to(OUTPUT_DTYPE), out_reg) fx.copy(copy_out, out_reg, fx.slice(out_tiles, (None, out_off // fx.Int32(4)))) else: _po_row = _po_row_base + qi_val_mg - for vhe in range_constexpr(_VHELOOP): + for vhe in range_constexpr(VHELOOP): hs_base = fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + rowid * fx.Int32(4) po_off = _po_row * stride_po_ql + qhead * fx.Int32(head_dim) + hs_base fx.memref_store_vec(outelems_norm[vhe], partial_out_reg) From eed3b819171cdf216733e0b6692b89e368a713fd Mon Sep 17 00:00:00 2001 From: fsx950223 Date: Tue, 28 Jul 2026 07:18:01 +0000 Subject: [PATCH 7/7] refactor(pa): Simplify metadata constants and loop state Signed-off-by: fsx950223 Co-authored-by: Cursor --- kernels/attention/pa_metadata.py | 415 ++++++++++++++----------------- 1 file changed, 185 insertions(+), 230 deletions(-) diff --git a/kernels/attention/pa_metadata.py b/kernels/attention/pa_metadata.py index 5df7da4ec..7a5e474f3 100644 --- a/kernels/attention/pa_metadata.py +++ b/kernels/attention/pa_metadata.py @@ -24,8 +24,8 @@ from flydsl.expr.typing import T from kernels.attention.pa_common import _compute_block_base_dw_i64, _prefetch_q_chunks_tile from kernels.common import dpp_utils -from kernels.common.kernels_common import get_warp_size -from kernels.common.tensor_shim import _run_compiled +from kernels.common.kernels_common import dtype_to_elem_type, get_warp_size +from kernels.common.tensor_shim import _run_compiled, get_dtype_str from kernels.common.utils import ( copy_load, copy_store, @@ -41,11 +41,6 @@ _WORK_INFO_FIELDS = 8 _FLAT_BUFFER_ELEMENTS = 1 << 30 -_PA_OUTPUT_DTYPES = { - "f32": fx.Float32, - "f16": fx.Float16, - "bf16": fx.BFloat16, -} def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = None): @@ -80,7 +75,7 @@ def get_pa_metadata_info_v1(batch_size: int, num_head_k: int = 1, num_cu: int = NUM_WARPS = 4 -WARP_SIZE = 64 +WARP_SIZE = get_warp_size() BLOCK_THREADS = NUM_WARPS * WARP_SIZE # 256 @@ -131,11 +126,11 @@ def _load_k_flat( tile_tok_base = tile_token_offset_i32 + k_tok_thread_base for td in range_constexpr(TLOOP): - kbo = tile_tok_base + fx.Int32(td * MFMA_N) + kbo = tile_tok_base + td * MFMA_N kbo_dw = kbo * c_tok_stride_dw for qkhe in range_constexpr(qkhe_loop): ka_dw = k_block_base_dw_i64 + fx.Int64(kbo_dw + k_he_off_dw[qkhe]) - k2 = load_global_16b(k_global_ptr, ka_dw * fx.Int64(4), k_copy_atom, k_reg) + k2 = load_global_16b(k_global_ptr, ka_dw * 4, k_copy_atom, k_reg) k2_words = fx.Vector(k2) k_flat.append(k2_words[0]) k_flat.append(k2_words[1]) @@ -150,10 +145,10 @@ def _build_pa_k_thread_invariants( *, qkhe_loop: int = 2, ): - k_tok_thread_base = warp_id * fx.Int32(TOKENS_PER_WARP) + lane16id - c_tok_stride_dw = fx.Int32(FP8_ELEMS_16B // 4) - c_he_stride_dw = fx.Int32(KV_BLOCK_SIZE * FP8_ELEMS_16B // 4) - k_he_off_dw = [rowid * c_he_stride_dw + fx.Int32(qkhe * 4) * c_he_stride_dw for qkhe in range(qkhe_loop)] + k_tok_thread_base = warp_id * TOKENS_PER_WARP + lane16id + c_tok_stride_dw = FP8_ELEMS_16B // 4 + c_he_stride_dw = KV_BLOCK_SIZE * FP8_ELEMS_16B // 4 + k_he_off_dw = [rowid * c_he_stride_dw + qkhe * 4 * c_he_stride_dw for qkhe in range(qkhe_loop)] return k_tok_thread_base, c_tok_stride_dw, k_he_off_dw @@ -166,32 +161,32 @@ def _compute_mtp_group_state( query_group_size, ): g_off = mtp_group_idx * 16 - lane_pair_raw = lane16id + fx.Int32(g_off) - c_total_pairs = fx.Int32(query_length * query_group_size) - c_pair_max = fx.Int32(query_length * query_group_size - 1) - c_ql_m1 = fx.Int32(query_length - 1) + lane_pair_raw = lane16id + g_off + total_pairs = query_length * query_group_size + pair_max = total_pairs - 1 + qlen_minus1 = query_length - 1 if const_expr((query_length * query_group_size) % MFMA_N == 0): lane_pair = lane_pair_raw else: - lane_pair = arith.select(lane_pair_raw < c_total_pairs, lane_pair_raw, c_pair_max) + lane_pair = (lane_pair_raw < total_pairs).select(lane_pair_raw, pair_max) qi_raw = udiv_const(lane_pair, query_group_size) if const_expr((query_length * query_group_size) % MFMA_N == 0): qi_val = qi_raw else: - qi_val = arith.select(qi_raw < c_ql_m1, qi_raw, c_ql_m1) + qi_val = (qi_raw < qlen_minus1).select(qi_raw, qlen_minus1) qhi_pos = urem_const(lane_pair, query_group_size) - lqh_pair_raw = local_qhead_idx + fx.Int32(g_off) + lqh_pair_raw = local_qhead_idx + g_off if const_expr((query_length * query_group_size) % MFMA_N == 0): lqh_pair = lqh_pair_raw else: - lqh_pair = arith.select(lqh_pair_raw < c_total_pairs, lqh_pair_raw, c_pair_max) + lqh_pair = (lqh_pair_raw < total_pairs).select(lqh_pair_raw, pair_max) lqi_raw = udiv_const(lqh_pair, query_group_size) if const_expr((query_length * query_group_size) % MFMA_N == 0): qi_for_q = lqi_raw else: - qi_for_q = arith.select(lqi_raw < c_ql_m1, lqi_raw, c_ql_m1) + qi_for_q = (lqi_raw < qlen_minus1).select(lqi_raw, qlen_minus1) local_qhead_idx_for_q = urem_const(lqh_pair, query_group_size) return qi_val, qhi_pos, qi_for_q, local_qhead_idx_for_q @@ -209,7 +204,7 @@ def _finish_q_fragments( ): qkhe_loop = head_dim // QKHE_PER_FETCH q_lanes_per_head = head_dim // Q_ELEMS_PER_LANE - lds_q_base = local_qhead_idx * fx.Int32(head_dim // 4) + lane16id * 2 + lds_q_base = local_qhead_idx * (head_dim // 4) + lane16id * 2 abs_mask = fx.Vector.filled(4, 0x7FFFFFFF, fx.Int32) q_f32_chunks = [] @@ -222,10 +217,7 @@ def _finish_q_fragments( for sh in [8, 4, 2, 1]: local_max = fx.maxnumf(local_max, dpp_utils.dpp_xor_f32(local_max, sh)) - query_scale_lane = (local_max > fx.Float32(0.0)).select( - local_max * fx.Float32(1.0 / FP8_MAX), - fx.Float32(1.0), - ) + query_scale_lane = (local_max > 0.0).select(local_max * (1.0 / FP8_MAX), 1.0) inv_query_scale = fx.Float32(rcp_f32(query_scale_lane)) q_words = [] for q_f32 in q_f32_chunks: @@ -233,7 +225,7 @@ def _finish_q_fragments( lo = rocdl.cvt_pk_fp8_f32(T.i32, p[0], p[1], fx.Int32(0), False) q_words.append(rocdl.cvt_pk_fp8_f32(T.i32, p[2], p[3], lo, True)) - if lane16id == fx.Int32(0): + if lane16id == 0: fx.ptr_store( fx.Vector.from_elements([query_scale_lane], dtype=fx.Float32), softmax_base + local_qhead_idx, @@ -241,7 +233,7 @@ def _finish_q_fragments( q_vec = fx.Vector.from_elements(q_words, dtype=fx.Int32) if const_expr(q_lanes_per_head < MFMA_N): - if lane16id < fx.Int32(q_lanes_per_head): + if lane16id < q_lanes_per_head: fx.ptr_store(q_vec, logits_base + lds_q_base) else: fx.ptr_store(q_vec, logits_base + lds_q_base) @@ -251,7 +243,7 @@ def _finish_q_fragments( query_scale_lane = fx.ptr_load(softmax_base + lane16id, result_type=fx.Vector.make_type(1, fx.Float32))[0] for qkhe in range_constexpr(qkhe_loop): for qkr in range_constexpr(2): - lds_rd = lane16id * fx.Int32(head_dim // 8) + fx.Int32(qkhe * 8) + rowid * fx.Int32(2) + fx.Int32(qkr) + lds_rd = lane16id * (head_dim // 8) + qkhe * 8 + rowid * 2 + qkr q_frags.append( fx.ptr_load( fx.recast_iter(fx.Int64, logits_base) + lds_rd, @@ -284,8 +276,8 @@ def _prefetch_mtp_group_query( query_length=query_length, query_group_size=query_group_size, ) - q_row = batch_idx * fx.Int32(query_length) + qi_for_q - q_base = q_row * stride_q_seq + (kv_h * fx.Int32(query_group_size) + local_qhead_idx_for_q) * stride_q_head + q_row = batch_idx * query_length + qi_for_q + q_base = q_row * stride_q_seq + (kv_h * query_group_size + local_qhead_idx_for_q) * stride_q_head q_chunks = _prefetch_q_chunks_tile( q_tiles, q_copy_atom, @@ -297,9 +289,8 @@ def _prefetch_mtp_group_query( return qi_val, qhi_pos, q_chunks -def _normalize_pa_output(running_sum, outs, zero_f): - one_f = fx.Float32(1.0).ir_value() - safe_sum = arith.select(running_sum > zero_f, running_sum, one_f) +def _normalize_pa_output(running_sum, outs): + safe_sum = (running_sum > 0.0).select(running_sum, 1.0) inv_sum = rcp_f32(safe_sum) return [out * fx.Float32(inv_sum) for out in outs] @@ -333,39 +324,34 @@ def _make_pa_phase_helpers( qkhe_loop = head_dim // QKHE_PER_FETCH vhe_loop = head_dim // MFMA_N // NUM_WARPS - vhead_elems = [ - fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + lane16id for vhe in range(vhe_loop) - ] - v_tok_thread_off = [fx.Int32(vt * TOKENS_PER_WARP) + rowid * fx.Int32(MFMA_N) for vt in range(VTLOOP)] + vhead_elems = [warp_id * MFMA_N + lane16id + vhe * NUM_WARPS * MFMA_N for vhe in range(vhe_loop)] + v_tok_thread_off = [rowid * MFMA_N + vt * TOKENS_PER_WARP for vt in range(VTLOOP)] if const_expr(trans_v): - vhead_elem_dw = [vhead_elems[vhe] * fx.Int32(FP8_ELEMS_16B // 4) for vhe in range(vhe_loop)] + vhead_elem_dw = [value * (FP8_ELEMS_16B // 4) for value in vhead_elems] else: - vhead_elem_dw = [vhead_elems[vhe] * fx.Int32(KV_BLOCK_SIZE // 4) for vhe in range(vhe_loop)] + vhead_elem_dw = [value * (KV_BLOCK_SIZE // 4) for value in vhead_elems] - kv_tok_thread_base = warp_id * fx.Int32(TOKENS_PER_WARP) + rowid * fx.Int32(4) - rowid_8x8 = rowid >> fx.Int32(1) - offset_in_slot = rowid & fx.Int32(1) + kv_tok_thread_base = warp_id * TOKENS_PER_WARP + rowid * 4 + rowid_8x8 = rowid >> 1 + offset_in_slot = rowid & 1 prob_row_i32 = PROB_ROW_STRIDE_BYTES // 4 prob_row_i64 = PROB_ROW_STRIDE_BYTES // 8 prob_wr_thread_base = ( - warp_id * fx.Int32(4 * MFMA_N * prob_row_i32) - + lane16id * fx.Int32(prob_row_i32) - + rowid_8x8 * fx.Int32(2) - + offset_in_slot + warp_id * (4 * MFMA_N * prob_row_i32) + lane16id * prob_row_i32 + rowid_8x8 * 2 + offset_in_slot ) - pv_prob_read_base = rowid * fx.Int32(MFMA_N * prob_row_i64) + lane16id * fx.Int32(prob_row_i64) + pv_prob_read_base = rowid * (MFMA_N * prob_row_i64) + lane16id * prob_row_i64 - sm_lane_wave_base = lane16id * fx.Int32(NUM_WARPS) + sm_lane_wave_base = lane16id * NUM_WARPS sm_max_off = sm_lane_wave_base + warp_id - sm_sum_off = fx.Int32(NUM_WARPS * MFMA_N) + sm_lane_wave_base + warp_id - sm_rd_max_offs = [sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] - sm_rd_sum_offs = [fx.Int32(NUM_WARPS * MFMA_N) + sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] + sm_sum_off = NUM_WARPS * MFMA_N + sm_lane_wave_base + warp_id + sm_rd_max_offs = [sm_lane_wave_base + w for w in range(NUM_WARPS)] + sm_rd_sum_offs = [NUM_WARPS * MFMA_N + sm_lane_wave_base + w for w in range(NUM_WARPS)] sm_vmax_wr_off = None sm_vmax_rd_offs = None if const_expr(per_token_kv): - sm_vmax_wr_off = fx.Int32(2 * NUM_WARPS * MFMA_N) + sm_lane_wave_base + warp_id - sm_vmax_rd_offs = [fx.Int32(2 * NUM_WARPS * MFMA_N) + sm_lane_wave_base + fx.Int32(w) for w in range(NUM_WARPS)] + sm_vmax_wr_off = 2 * NUM_WARPS * MFMA_N + sm_lane_wave_base + warp_id + sm_vmax_rd_offs = [2 * NUM_WARPS * MFMA_N + sm_lane_wave_base + w for w in range(NUM_WARPS)] neg_inf = fx.Float32(float("-inf")) zero_f = fx.Float32(0.0) @@ -373,15 +359,15 @@ def _make_pa_phase_helpers( pv_prob_i64_elems = [] for vt in range_constexpr(VTLOOP): for j in range_constexpr(2): - p_elem = fx.Int32(vt * 4 * MFMA_N * (PROB_ROW_STRIDE_BYTES // 8)) + pv_prob_read_base + fx.Int32(j) + p_elem = vt * 4 * MFMA_N * (PROB_ROW_STRIDE_BYTES // 8) + pv_prob_read_base + j pv_prob_i64_elems.append(p_elem) def _load_kv_scale_scalars(tile_token_offset_i32, phys_block): if const_expr(per_token_kv): scale_block_base = phys_block * stride_ks_block + kv_h * stride_ks_head - scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id + scale_stage_token = warp_id * WARP_SIZE + rowid * MFMA_N + lane16id scale_global_token = tile_token_offset_i32 + scale_stage_token - scale_offset = fx.Int32(scale_block_base + scale_global_token) + scale_offset = scale_block_base + scale_global_token k_scale_scalar = copy_load(ks_tiles, scale_offset, scale_copy_atom, scale_reg)[0] v_scale_scalar = copy_load(vs_tiles, scale_offset, scale_copy_atom, scale_reg)[0] return k_scale_scalar, v_scale_scalar @@ -398,11 +384,11 @@ def _load_v_and_scales( for vhe in range_constexpr(vhe_loop): v_token_in_block = tile_token_offset_i32 + v_tok_thread_off[vt] if const_expr(trans_v): - vt_group = v_token_in_block >> fx.Int32(4) - va_dw_delta = vt_group * fx.Int32(head_dim * FP8_ELEMS_16B // 4) + vhead_elem_dw[vhe] + vt_group = v_token_in_block >> 4 + va_dw_delta = vt_group * (head_dim * FP8_ELEMS_16B // 4) + vhead_elem_dw[vhe] else: - va_dw_delta = vhead_elem_dw[vhe] + (v_token_in_block >> fx.Int32(2)) - va_byte = (v_block_base_dw + fx.Int64(va_dw_delta)) * fx.Int64(4) + va_dw_delta = vhead_elem_dw[vhe] + (v_token_in_block >> 2) + va_byte = (v_block_base_dw + fx.Int64(va_dw_delta)) * 4 v_i64x2 = load_global_16b(v_global_ptr, va_byte, v_copy_atom, v_reg) vhe_data.append(v_i64x2) v_results.append(vhe_data) @@ -410,7 +396,7 @@ def _load_v_and_scales( k_scale_vecs = [] v_scale_vecs = [] if const_expr(per_token_kv): - scale_stage_token = warp_id * fx.Int32(WARP_SIZE) + rowid * fx.Int32(MFMA_N) + lane16id + scale_stage_token = warp_id * WARP_SIZE + rowid * MFMA_N + lane16id k_scale_scalar, v_scale_scalar = scale_scalars fx.ptr_store( fx.Vector.from_elements([k_scale_scalar], dtype=fx.Float32), @@ -418,18 +404,18 @@ def _load_v_and_scales( ) fx.ptr_store( fx.Vector.from_elements([v_scale_scalar], dtype=fx.Float32), - scale_base + (fx.Int32(LDS_SCALE_V_OFFSET) + scale_stage_token), + scale_base + (LDS_SCALE_V_OFFSET + scale_stage_token), ) rocdl.sched_barrier(0) gpu.barrier() for td in range_constexpr(TLOOP): - scale_row_base = kv_tok_thread_base + fx.Int32(td * MFMA_N) + scale_row_base = kv_tok_thread_base + td * MFMA_N k_scale_vecs.append( fx.ptr_load(scale_base + scale_row_base, result_type=fx.Vector.make_type(4, fx.Float32)) ) v_scale_vecs.append( fx.ptr_load( - scale_base + (fx.Int32(LDS_SCALE_V_OFFSET) + scale_row_base), + scale_base + (LDS_SCALE_V_OFFSET + scale_row_base), result_type=fx.Vector.make_type(4, fx.Float32), ) ) @@ -446,16 +432,16 @@ def _store_vmax_warp(partition_start, *, seq_end=None, v_scale_vecs=None): vs = (_token_vec_i32(kv_tok_base, td) < seq_end).select(vs, zero_f) v_max_warp = fx.maxnumf(v_max_warp, vs.reduce("max")) for sh in [32, 16]: - v_max_warp = fx.maxnumf(v_max_warp, v_max_warp.shuffle_xor(fx.Int32(sh), fx.Int32(WARP_SIZE))) + v_max_warp = fx.maxnumf(v_max_warp, v_max_warp.shuffle_xor(sh, WARP_SIZE)) fx.ptr_store( fx.Vector.from_elements([v_max_warp], dtype=fx.Float32), softmax_base + sm_vmax_wr_off, ) def _token_vec_i32(kv_tok_base, td: int): - kv_tok_td_base = kv_tok_base + fx.Int32(td * MFMA_N) + kv_tok_td_base = kv_tok_base + td * MFMA_N return fx.Vector.from_elements( - [kv_tok_td_base + fx.Int32(i) for i in range_constexpr(4)], + [kv_tok_td_base + i for i in range_constexpr(4)], dtype=fx.Int32, ) @@ -468,7 +454,7 @@ def _qk_and_intra_softmax( *, k_scale_vecs=None, ): - query_scale = fx.Float32(query_scale_lane * softmax_scale) + query_scale = query_scale_lane * softmax_scale d_out = [] for td in range_constexpr(TLOOP): acc = fx.Vector.filled(4, 0.0, fx.Float32) @@ -489,7 +475,7 @@ def _qk_and_intra_softmax( d_out[td] = logits_vec qk_max = fx.maxnumf(qk_max, fx.Vector(logits_vec).reduce("max")) for sh in [32, 16]: - qk_max = fx.maxnumf(qk_max, qk_max.shuffle_xor(fx.Int32(sh), fx.Int32(WARP_SIZE))) + qk_max = fx.maxnumf(qk_max, qk_max.shuffle_xor(sh, WARP_SIZE)) fx.ptr_store( fx.Vector.from_elements([qk_max], dtype=fx.Float32), softmax_base + sm_max_off, @@ -505,24 +491,20 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs, *, partition_max = fx.maxnumf(partition_max, max_vec[w]) new_rmax = fx.maxnumf(rmax, partition_max) - safe_eff_max = arith.select(partition_max > neg_inf, new_rmax, zero_f) + safe_eff_max = (partition_max > neg_inf).select(new_rmax, zero_f) local_exp_sum = zero_f for td in range_constexpr(TLOOP): - diff_vec = fx.Vector(d_out[td]) - fx.Float32(safe_eff_max) - p_vec = exp2_f32_fast(diff_vec * fx.Float32(LOG2E)) + diff_vec = fx.Vector(d_out[td]) - safe_eff_max + p_vec = exp2_f32_fast(diff_vec * LOG2E) local_exp_sum = local_exp_sum + fx.Vector(p_vec).reduce("add") d_out[td] = p_vec for sh in [32, 16]: - local_exp_sum = local_exp_sum + local_exp_sum.shuffle_xor(fx.Int32(sh), fx.Int32(WARP_SIZE)) + local_exp_sum = local_exp_sum + local_exp_sum.shuffle_xor(sh, WARP_SIZE) fx.ptr_store( fx.Vector.from_elements([local_exp_sum], dtype=fx.Float32), softmax_base + sm_sum_off, ) - accum_scale = arith.select( - rmax > neg_inf, - exp2_f32_fast((rmax - new_rmax) * fx.Float32(LOG2E).ir_value()), - zero_f, - ) + accum_scale = (rmax > neg_inf).select(exp2_f32_fast((rmax - new_rmax) * LOG2E), zero_f) gpu.barrier() sum_vec = fx.ptr_load(softmax_base + (sm_rd_sum_offs[0]), result_type=fx.Vector.make_type(4, fx.Float32)) @@ -540,8 +522,8 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs, *, for w in range_constexpr(NUM_WARPS): w_vmax = vmax_vec[w] v_max_global = fx.maxnumf(v_max_global, w_vmax) - v_correction = v_max_global * fx.Float32(1.0 / FP8_MAX).ir_value() - norm_factor = rcp_f32(v_correction + fx.Float32(1e-8 / FP8_MAX).ir_value()) + v_correction = v_max_global * (1.0 / FP8_MAX) + norm_factor = rcp_f32(v_correction + 1e-8 / FP8_MAX) normalized_v_scales = [ fx.Vector(v_scale_vecs[td]) * fx.Float32(norm_factor) for td in range_constexpr(TLOOP) ] @@ -557,7 +539,7 @@ def _cross_warp_softmax_and_prob_pack(d_out, rmax, rsum, outs, v_scale_vecs, *, pv = fx.Vector(d_out[td]) lo = rocdl.cvt_pk_fp8_f32(T.i32, pv[0], pv[1], fx.Int32(0), False) pk = rocdl.cvt_pk_fp8_f32(T.i32, pv[2], pv[3], lo, True) - elem_base = prob_wr_thread_base + fx.Int32(td * MFMA_N * (PROB_ROW_STRIDE_BYTES // 4)) + elem_base = prob_wr_thread_base + td * MFMA_N * (PROB_ROW_STRIDE_BYTES // 4) pk_vec = fx.Vector.from_elements([pk], dtype=fx.Int32) fx.ptr_store(pk_vec, logits_base + elem_base) @@ -625,19 +607,18 @@ def compile_pa_metadata_v1( gqa: int, kv_granularity: int, query_length: int, - warp_size: int, ): """Compile the single-wave worklist scheduler for a fixed shape config.""" assert is_pow2(kv_granularity), "kv_granularity must be power of 2" assert num_cu % num_heads_k == 0, "num_cu must be divisible by num_heads_k" num_splits_per_khead = num_cu // num_heads_k _shuffle_offsets = [] - _o = warp_size // 2 + _o = WARP_SIZE // 2 while _o >= 1: _shuffle_offsets.append(_o) _o //= 2 - @flyc.kernel(known_block_size=(warp_size, 1, 1)) + @flyc.kernel(known_block_size=(WARP_SIZE, 1, 1)) def pa_metadata_v1_kernel( seqlens_qo_indptr_ptr: fx.Tensor, # [num_batches + 1] i32 (cumulative qo seqlens) context_lens_ptr: fx.Tensor, # [num_batches] i32 @@ -671,106 +652,93 @@ def _divide_buffer(tensor, tile_layout): store_reg2 = fx.make_rmem_tensor(2, fx.Int32) store_reg4 = fx.make_rmem_tensor(4, fx.Int32) - c0 = fx.Int32(0) - c1 = fx.Int32(1) - c_qlen = fx.Int32(query_length) - c_nb = num_batches # fx.Int32 runtime - c_nspk = fx.Int32(num_splits_per_khead) - c_numcu = fx.Int32(num_cu) - c_ws = fx.Int32(warp_size) - c_kvg = fx.Int32(kv_granularity) lane = fx.Int32(gpu.thread_id("x")) def _num_part(batch_idx): ctxv = copy_load(context_lens, batch_idx, copy_i32, load_reg)[0] - return fx.Int32(arith.ceildivui(ctxv.ir_value(), c_kvg.ir_value())) + return (ctxv + kv_granularity - 1) // kv_granularity copy_store( work_indptr, 0, copy_i32, store_reg, - fx.Vector.from_elements([fx.Int32(0)], dtype=fx.Int32), + fx.Vector.from_elements([0], dtype=fx.Int32), ) copy_store( reduce_indptr, 0, copy_i32, store_reg, - fx.Vector.from_elements([fx.Int32(0)], dtype=fx.Int32), + fx.Vector.from_elements([0], dtype=fx.Int32), ) # Lane-strided partition count followed by a wave reduction. b = lane - sum_blocks = c0 - while b < c_nb: + sum_blocks = fx.Int32(0) + while b < num_batches: nblk = _num_part(b) - b = b + c_ws + b = b + WARP_SIZE sum_blocks = sum_blocks + nblk for sh in _shuffle_offsets: - sum_blocks = sum_blocks + sum_blocks.shuffle_xor(fx.Int32(sh), c_ws) + sum_blocks = sum_blocks + sum_blocks.shuffle_xor(sh, WARP_SIZE) - average = fx.Int32(arith.divui(sum_blocks.ir_value(), c_nspk.ir_value())) - reminder = fx.Int32(arith.remui(sum_blocks.ir_value(), c_nspk.ir_value())) + average = sum_blocks // num_splits_per_khead + reminder = sum_blocks % num_splits_per_khead def _remain_for_cid(cid_val): - mod = fx.Int32(arith.remui(cid_val.ir_value(), c_nspk.ir_value())) - return average + (mod < reminder).select(c1, c0) + mod = cid_val % num_splits_per_khead + return average + (mod < reminder).select(1, 0) - cid = c0 - num_works = c0 + cid = fx.Int32(0) + num_works = fx.Int32(0) for khead in range_constexpr(num_heads_k): qh_start = khead * gqa qh_end = (khead + 1) * gqa qhr_const = (qh_end << 16) | (qh_start & 0xFFFF) # python int constant - kvend0 = _num_part(c0) # partitions in batch 0 (cumulative kv end) - remain0 = _remain_for_cid(cid) - - cid_ = cid - batch_ = c0 - kvblk_ = c0 - nsplit_ = c0 - nworks_ = num_works - pidx_ = c0 - kvbeg_ = c0 - kvend_ = kvend0 - remain_ = remain0 - lri_ = c0 - grt_ = c0 - - while (cid_ < c_numcu) & (batch_ < c_nb): - pages = kvend_ - kvbeg_ + batch_ = fx.Int32(0) + kvblk_ = fx.Int32(0) + nsplit_ = fx.Int32(0) + pidx_ = fx.Int32(0) + kvbeg_ = fx.Int32(0) + kvend = _num_part(0) # partitions in batch 0 (cumulative kv end) + remain = _remain_for_cid(cid) + lri_ = fx.Int32(0) + grt_ = fx.Int32(0) + + while (cid < num_cu) & (batch_ < num_batches): + pages = kvend - kvbeg_ remain_kv = pages - kvblk_ - do_finish = remain_ >= remain_kv # fx bool + do_finish = remain >= remain_kv # fx bool qo_start = copy_load(sq, batch_, copy_i32, load_reg)[0] - qo_end = copy_load(sq, batch_ + c1, copy_i32, load_reg)[0] + qo_end = copy_load(sq, batch_ + 1, copy_i32, load_reg)[0] kv_start = kvbeg_ + kvblk_ # same for both branches - f_kv_end = kvend_ # min(kv_start + remain_kv, kvend_) == kvend_ - nsplit_pos = nsplit_ > c0 - f_ploc = nsplit_pos.select(pidx_, fx.Int32(-1)) - f_pidx2 = nsplit_pos.select(pidx_ + c_qlen, pidx_) - f_nworks2 = nworks_ + c1 - f_remain2 = remain_ - remain_kv - f_batch2 = batch_ + c1 + f_kv_end = kvend # min(kv_start + remain_kv, kvend) == kvend + nsplit_pos = nsplit_ > 0 + f_ploc = nsplit_pos.select(pidx_, -1) + f_pidx2 = nsplit_pos.select(pidx_ + query_length, pidx_) + f_nworks2 = num_works + 1 + f_remain2 = remain - remain_kv + f_batch2 = batch_ + 1 # Select evaluates both values, so clamp the speculative load. - nb_in_range = f_batch2 < c_nb - safe_idx = nb_in_range.select(f_batch2, c0) - f_new_pages = nb_in_range.select(_num_part(safe_idx), c0) - f_kvbeg2 = kvend_ - f_kvend2 = kvend_ + f_new_pages - - s_emit = remain_ > c0 - s_kv_end_raw = kv_start + remain_ - s_kv_end = (s_kv_end_raw < kvend_).select(s_kv_end_raw, kvend_) - s_nworks2 = s_emit.select(nworks_ + c1, nworks_) - s_pidx2 = s_emit.select(pidx_ + c_qlen, pidx_) - s_kvblk2 = s_emit.select(kvblk_ + remain_, kvblk_) - s_nsplit2 = s_emit.select(nsplit_ + c1, nsplit_) - s_cid2 = cid_ + c1 + nb_in_range = f_batch2 < num_batches + safe_idx = nb_in_range.select(f_batch2, 0) + f_new_pages = nb_in_range.select(_num_part(safe_idx), 0) + f_kvbeg2 = kvend + f_kvend2 = kvend + f_new_pages + + s_emit = remain > 0 + s_kv_end_raw = kv_start + remain + s_kv_end = (s_kv_end_raw < kvend).select(s_kv_end_raw, kvend) + s_nworks2 = s_emit.select(num_works + 1, num_works) + s_pidx2 = s_emit.select(pidx_ + query_length, pidx_) + s_kvblk2 = s_emit.select(kvblk_ + remain, kvblk_) + s_nsplit2 = s_emit.select(nsplit_ + 1, nsplit_) + s_cid2 = cid + 1 s_remain2 = _remain_for_cid(s_cid2) w_ploc = do_finish.select(f_ploc, pidx_) @@ -782,8 +750,8 @@ def _remain_for_cid(cid_val): qo_end, kv_start, w_kv_end, - c0, - fx.Int32(qhr_const), + 0, + qhr_const, ] for half in range_constexpr(2): start = half * 4 @@ -797,93 +765,91 @@ def _remain_for_cid(cid_val): fx.copy( copy_i32x4, store_reg4, - fx.slice(work_info, (None, nworks_ * fx.Int32(2) + fx.Int32(half))), + fx.slice(work_info, (None, num_works * 2 + half)), ) do_reduce = do_finish & nsplit_pos - num_splits = nsplit_ + c1 + num_splits = nsplit_ + 1 # Non-reduce writes are overwritten before their slots become visible. copy_store( reduce_indptr, - grt_ + c1, + grt_ + 1, copy_i32, store_reg, - fx.Vector.from_elements([fx.Int32(lri_ + num_splits)], dtype=fx.Int32), + fx.Vector.from_elements([lri_ + num_splits], dtype=fx.Int32), ) fx.memref_store_vec( fx.Vector.from_elements([qo_start, qo_end], dtype=fx.Int32), store_reg2, ) fx.copy(copy_i32x2, store_reg2, fx.slice(reduce_final_map, (None, grt_))) - rcount = do_reduce.select(num_splits, c0) - sidx = c0 + rcount = do_reduce.select(num_splits, 0) + sidx = fx.Int32(0) while sidx < rcount: - val = pidx_ - (nsplit_ - sidx) * c_qlen + val = pidx_ - (nsplit_ - sidx) * query_length copy_store( reduce_partial_map, lri_ + sidx, copy_i32, store_reg, - fx.Vector.from_elements([fx.Int32(val)], dtype=fx.Int32), + fx.Vector.from_elements([val], dtype=fx.Int32), ) - sidx = sidx + c1 + sidx = sidx + 1 next_num_works = do_finish.select(f_nworks2, s_nworks2) # The last same-value-race write before advancing cid is authoritative. copy_store( work_indptr, - cid_ + c1, + cid + 1, copy_i32, store_reg, - fx.Vector.from_elements([fx.Int32(next_num_works)], dtype=fx.Int32), + fx.Vector.from_elements([next_num_works], dtype=fx.Int32), ) ( - cid_, + cid, batch_, kvblk_, nsplit_, - nworks_, + num_works, pidx_, kvbeg_, - kvend_, - remain_, + kvend, + remain, lri_, grt_, ) = ( - do_finish.select(cid_, s_cid2), + do_finish.select(cid, s_cid2), do_finish.select(f_batch2, batch_), - do_finish.select(c0, s_kvblk2), - do_finish.select(c0, s_nsplit2), + do_finish.select(0, s_kvblk2), + do_finish.select(0, s_nsplit2), next_num_works, do_finish.select(f_pidx2, s_pidx2), do_finish.select(f_kvbeg2, kvbeg_), - do_finish.select(f_kvend2, kvend_), + do_finish.select(f_kvend2, kvend), do_finish.select(f_remain2, s_remain2), - lri_ + do_reduce.select(num_splits, c0), - grt_ + do_reduce.select(c1, c0), + lri_ + do_reduce.select(num_splits, 0), + grt_ + do_reduce.select(1, 0), ) - cid = cid_ - num_works = nworks_ last_reduce_indptr = lri_ global_reduce_tile_idx = grt_ - in_range = cid < c_numcu - cid = in_range.select(cid + c1, cid) + in_range = cid < num_cu + cid = in_range.select(cid + 1, cid) it_t = cid - while it_t <= c_numcu: + while it_t <= num_cu: copy_store( work_indptr, it_t, copy_i32, store_reg, - fx.Vector.from_elements([fx.Int32(num_works)], dtype=fx.Int32), + fx.Vector.from_elements([num_works], dtype=fx.Int32), ) - it_t = it_t + c1 + it_t = it_t + 1 - c_rip_size = c_nb + c1 # reduce_indptr length = num_batches + 1 + c_rip_size = num_batches + 1 # reduce_indptr length = num_batches + 1 it_r = global_reduce_tile_idx while it_r < c_rip_size: copy_store( @@ -891,9 +857,9 @@ def _remain_for_cid(cid_val): it_r, copy_i32, store_reg, - fx.Vector.from_elements([fx.Int32(last_reduce_indptr)], dtype=fx.Int32), + fx.Vector.from_elements([last_reduce_indptr], dtype=fx.Int32), ) - it_r = it_r + c1 + it_r = it_r + 1 @flyc.jit def launch_pa_metadata_v1( @@ -916,7 +882,7 @@ def launch_pa_metadata_v1( reduce_final_map, reduce_partial_map, num_batches, - ).launch(grid=(1, 1, 1), block=(warp_size, 1, 1), stream=stream) + ).launch(grid=(1, 1, 1), block=(WARP_SIZE, 1, 1), stream=stream) return {"kernel": pa_metadata_v1_kernel, "launch": launch_pa_metadata_v1} @@ -946,7 +912,6 @@ def get_pa_metadata_v1( if num_cu is None: num_cu = torch.cuda.get_device_properties(dev).multi_processor_count num_batches = context_lens.shape[0] - warp_size = get_warp_size() compiled = compile_pa_metadata_v1( num_cu=num_cu, @@ -954,7 +919,6 @@ def get_pa_metadata_v1( gqa=query_group_size, kv_granularity=kv_granularity, query_length=query_length, - warp_size=warp_size, ) work_info_flat = work_info.view(-1) @@ -983,7 +947,7 @@ def compile_pa_decode_metadata( query_length: int = 1, query_input_dtype: str = "bf16", head_dim: int = 128, - block_size: int = None, + block_size: int = KV_BLOCK_SIZE, output_dtype_str: str = "bf16", ): """Compile a PS-mode PA decode kernel. @@ -998,8 +962,6 @@ def compile_pa_decode_metadata( (``work_info[1]``) ``< 0`` writes the final output directly to ``out``, while ``>= 0`` writes a partial slot that ``pa_reduce_v1`` later combines. """ - if block_size is None: - block_size = KV_BLOCK_SIZE if block_size != KV_BLOCK_SIZE: raise ValueError(f"compile_pa_decode_metadata only supports block_size={KV_BLOCK_SIZE}, got {block_size}") if head_dim % QKHE_PER_FETCH != 0 or head_dim % (MFMA_N * NUM_WARPS) != 0 or head_dim % Q_ELEMS_PER_LANE != 0: @@ -1009,8 +971,8 @@ def compile_pa_decode_metadata( Q_LANES_PER_HEAD = head_dim // Q_ELEMS_PER_LANE if query_input_dtype not in ("bf16", "f16"): raise ValueError(f"`compile_pa_decode_metadata` only supports bf16/f16 queries, got {query_input_dtype!r}") - QUERY_DTYPE = fx.BFloat16 if query_input_dtype == "bf16" else fx.Float16 - OUTPUT_DTYPE = _PA_OUTPUT_DTYPES[output_dtype_str] + QUERY_DTYPE = dtype_to_elem_type(query_input_dtype) + OUTPUT_DTYPE = dtype_to_elem_type(output_dtype_str) if softmax_scale is None: softmax_scale = 1.0 / (head_dim**0.5) softmax_scale = float(softmax_scale) @@ -1062,9 +1024,9 @@ def pa_decode_metadata_kernel( tid = fx.Int32(gpu.thread_id("x")) cu_id = fx.Int32(gpu.block_id("x")) # CU index (0..num_sm-1) - lane16id = tid & fx.Int32(15) - rowid = (tid >> fx.Int32(4)) & fx.Int32(3) - warp_id = tid >> fx.Int32(6) + lane16id = tid & 15 + rowid = (tid >> 4) & 3 + warp_id = tid >> 6 def _divide_addr(addr, dtype, width): ptr = global_pointer_from_addr(addr, dtype, alignment=width * dtype.width // 8) @@ -1111,8 +1073,8 @@ def _divide_addr(addr, dtype, width): global_reg_16b = fx.make_rmem_tensor(16, fx.Uint8) if const_expr(per_token_kv): - k_scale_val = fx.Float32(1.0) - v_scale_val = fx.Float32(1.0) + k_scale_val = 1.0 + v_scale_val = 1.0 else: k_scale_val = copy_load(key_scales, 0, copy_f32, scale_reg)[0] v_scale_val = copy_load(value_scales, 0, copy_f32, scale_reg)[0] @@ -1124,7 +1086,7 @@ def _divide_addr(addr, dtype, width): if const_expr(per_token_kv): scale_base = fx.recast_iter(fx.Float32, fx.add_offset(lds.buf.ptr, fx.Int32(scale_off // 4))) - local_qhead_idx = warp_id * fx.Int32(4) + rowid + local_qhead_idx = warp_id * 4 + rowid _k_tok_thread_base, _c_tok_stride_dw, _k_he_off_dw = _build_pa_k_thread_invariants( warp_id, lane16id, @@ -1133,16 +1095,16 @@ def _divide_addr(addr, dtype, width): ) work_start = copy_load(work_indptr, cu_id, copy_i32, i32_reg)[0] - work_end = copy_load(work_indptr, cu_id + fx.Int32(1), copy_i32, i32_reg)[0] + work_end = copy_load(work_indptr, cu_id + 1, copy_i32, i32_reg)[0] for work_idx in range(work_start, work_end, fx.Int32(1)): # Two naturally aligned dwordx4 loads cover the eight work fields. - info_tile = work_idx * fx.Int32(2) + info_tile = work_idx * 2 wi_lo_v = copy_load(work_info, info_tile, copy_i32x4, i32x4_reg) batch_idx = wi_lo_v[0] partial_idx = wi_lo_v[1] qo_start = wi_lo_v[2] - wi_hi_v = copy_load(work_info, info_tile + fx.Int32(1), copy_i32x4, i32x4_reg) + wi_hi_v = copy_load(work_info, info_tile + 1, copy_i32x4, i32x4_reg) kv_start = wi_hi_v[0] kv_end = wi_hi_v[1] q_head_range = wi_hi_v[3] @@ -1152,7 +1114,7 @@ def _divide_addr(addr, dtype, width): kv_page_base = copy_load(kv_indptr, batch_idx, copy_i32, i32_reg)[0] local_part_start = kv_start - kv_part_base - q_head_start = q_head_range & fx.Int32(0xFFFF) + q_head_start = q_head_range & 0xFFFF kv_h = udiv_const(q_head_start, query_group_size) context_len = copy_load(context_lengths, batch_idx, copy_i32, i32_reg)[0] @@ -1182,7 +1144,7 @@ def _divide_addr(addr, dtype, width): scale_base=scale_base, stride_ks_block=stride_ks_block, stride_ks_head=stride_ks_head, - softmax_scale=fx.Float32(softmax_scale), + softmax_scale=softmax_scale, k_scale_val=k_scale_val, v_scale_val=v_scale_val, warp_id=warp_id, @@ -1192,11 +1154,11 @@ def _divide_addr(addr, dtype, width): ) # Negative partial_idx writes final output; split rows reserve the first QL slots. - _is_direct = partial_idx < fx.Int32(0) - _po_row_base = partial_idx + fx.Int32(query_length) + _is_direct = partial_idx < 0 + _po_row_base = partial_idx + query_length num_parts_in_work = kv_end - kv_start - last_part_idx_val = num_parts_in_work - fx.Int32(1) + last_part_idx_val = num_parts_in_work - 1 _mtp_groups = math.ceil(query_length * query_group_size / 16) @@ -1238,7 +1200,7 @@ def _divide_addr(addr, dtype, width): gpu.barrier() causal_bound_per_mtp = [ - context_len + fx.Int32(1 - query_length) + qi_per_mtp[_mtp_g] for _mtp_g in range(_mtp_groups) + context_len + (1 - query_length) + qi_per_mtp[_mtp_g] for _mtp_g in range(_mtp_groups) ] state_width = 2 + VHELOOP @@ -1267,7 +1229,7 @@ def _unpack_states(flat): # Process partition zero last for online-softmax stability. rel_part = last_part_idx_val - ib lp = local_part_start + rel_part - partition_start = lp * fx.Int32(KV_COMPUTE_BLOCK) + partition_start = lp * KV_COMPUTE_BLOCK phys_block = copy_load( kv_page_indices, @@ -1275,7 +1237,7 @@ def _unpack_states(flat): copy_i32, i32_reg, )[0] - tile_token_offset = urem_const(lp, parts_per_block) * fx.Int32(KV_COMPUTE_BLOCK) + tile_token_offset = urem_const(lp, parts_per_block) * KV_COMPUTE_BLOCK k_base = _compute_block_base_dw_i64(phys_block, stride_k_block, _k_head_off) scale_scalars = _load_kv_scale_scalars(tile_token_offset, phys_block) if const_expr(per_token_kv): @@ -1344,31 +1306,31 @@ def _unpack_states(flat): running_max = fx.Float32(rmax_raw) running_sum = fx.Float32(rsum_raw) outs = [fx.Vector(out_raw) for out_raw in outs_raw] - outelems_norm = _normalize_pa_output(running_sum, outs, fx.Float32(0.0)) + outelems_norm = _normalize_pa_output(running_sum, outs) qi_val_mg = qi_per_mtp[_mtp_g] qhi_pos_mg = qhi_per_mtp[_mtp_g] - qhead = kv_h * fx.Int32(query_group_size) + qhi_pos_mg + qhead = kv_h * query_group_size + qhi_pos_mg if _is_direct: out_row = qo_start + qi_val_mg for vhe in range_constexpr(VHELOOP): - hs_base = fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + rowid * fx.Int32(4) + hs_base = vhe * NUM_WARPS * MFMA_N + warp_id * MFMA_N + rowid * 4 out_off = out_row * stride_out_seq + qhead * stride_out_head + hs_base fx.memref_store_vec(outelems_norm[vhe].to(OUTPUT_DTYPE), out_reg) - fx.copy(copy_out, out_reg, fx.slice(out_tiles, (None, out_off // fx.Int32(4)))) + fx.copy(copy_out, out_reg, fx.slice(out_tiles, (None, out_off // 4))) else: _po_row = _po_row_base + qi_val_mg for vhe in range_constexpr(VHELOOP): - hs_base = fx.Int32(vhe * NUM_WARPS * MFMA_N) + warp_id * fx.Int32(MFMA_N) + rowid * fx.Int32(4) - po_off = _po_row * stride_po_ql + qhead * fx.Int32(head_dim) + hs_base + hs_base = vhe * NUM_WARPS * MFMA_N + warp_id * MFMA_N + rowid * 4 + po_off = _po_row * stride_po_ql + qhead * head_dim + hs_base fx.memref_store_vec(outelems_norm[vhe], partial_out_reg) fx.copy( copy_f32x4, partial_out_reg, - fx.slice(partial_out_tiles, (None, po_off // fx.Int32(4))), + fx.slice(partial_out_tiles, (None, po_off // 4)), ) - safe_sum_lse = (running_sum > fx.Float32(0.0)).select(running_sum, fx.Float32(1.0)) + safe_sum_lse = (running_sum > 0.0).select(running_sum, 1.0) log_sum = fmath.log(safe_sum_lse, fastmath=arith.FastMathFlags.fast) lse_val = running_max + log_sum pl_off = _po_row * stride_pl_ql + qhead @@ -1457,7 +1419,7 @@ def compile_pa_metadata_reduce( vec_width = 2 if head_dim % 2 == 0 else 1 block_threads = head_dim // vec_width assert 0 < block_threads <= 1024, "head_dim must fit in one workgroup" - output_dtype = _PA_OUTPUT_DTYPES[output_dtype_str] + output_dtype = dtype_to_elem_type(output_dtype_str) @flyc.kernel(known_block_size=(block_threads, 1, 1)) def pa_metadata_reduce_kernel( @@ -1512,10 +1474,10 @@ def _flat_divide(tensor): reg_output = fx.make_rmem_tensor(vec_width, output_dtype) lo = copy_load(reduce_indptr, g, copy_i32, reg_i32)[0] - hi = copy_load(reduce_indptr, g + fx.Int32(1), copy_i32, reg_i32)[0] + hi = copy_load(reduce_indptr, g + 1, copy_i32, reg_i32)[0] if hi > lo: - qo_start = copy_load(reduce_final_map, g * fx.Int32(2), copy_i32, reg_i32)[0] + qo_start = copy_load(reduce_final_map, g * 2, copy_i32, reg_i32)[0] out_row = qo_start + qi for slot, st in range( @@ -1535,23 +1497,23 @@ def _flat_divide(tensor): lse = copy_load(partial_lse, prow * stride_pl_row + qhead, copy_f32, reg_f32)[0] v = copy_load( partial_output, - prow * stride_po_row + qhead * fx.Int32(head_dim) + tid * fx.Int32(vec_width), + prow * stride_po_row + qhead * head_dim + tid * vec_width, copy_f32_vec, reg_f32_vec, ) m_new = m.maximumf(lse) - scale_old = exp2_f32_fast((m - m_new) * fx.Float32(LOG2E)) - w = exp2_f32_fast((lse - m_new) * fx.Float32(LOG2E)) + scale_old = exp2_f32_fast((m - m_new) * LOG2E) + w = exp2_f32_fast((lse - m_new) * LOG2E) denom_new = denom * scale_old + w acc_new = acc * fx.Float32(scale_old) + fx.Vector(v) * fx.Float32(w) results = yield [m_new, denom_new, acc_new] denom_f = fx.Float32(results[1]) acc_f = fx.Vector(results[2]) - safe_denom = (denom_f > fx.Float32(0.0)).select(denom_f, fx.Float32(1.0)) + safe_denom = (denom_f > 0.0).select(denom_f, 1.0) out_acc = acc_f * rcp_f32(safe_denom) - out_off = out_row * stride_out_seq + qhead * stride_out_head + tid * fx.Int32(vec_width) + out_off = out_row * stride_out_seq + qhead * stride_out_head + tid * vec_width output_val = fx.Vector(out_acc).to(output_dtype) fx.memref_store_vec(output_val, reg_output) fx.copy(copy_output, reg_output, fx.slice(final_output, (None, out_off))) @@ -1591,13 +1553,6 @@ def launch_pa_metadata_reduce( return {"launch": launch_pa_metadata_reduce, "kernel": pa_metadata_reduce_kernel} -_PA_PS_REDUCE_DTYPE_STR = { - torch.float32: "f32", - torch.float16: "f16", - torch.bfloat16: "bf16", -} - - def pa_metadata_reduce( *, partial_output: torch.Tensor, @@ -1621,7 +1576,7 @@ def pa_metadata_reduce( num_groups = reduce_indptr.numel() - 1 stride_po_row = num_query_heads * head_dim stride_pl_row = num_query_heads - out_dtype_str = _PA_PS_REDUCE_DTYPE_STR[final_output.dtype] + out_dtype_str = get_dtype_str(final_output.dtype) compiled = compile_pa_metadata_reduce( query_length=int(max_seqlen_q), num_query_heads=int(num_query_heads),