From 9bdd0cb6f879e16a90f5c740a898ba05fcd77a12 Mon Sep 17 00:00:00 2001 From: coderfeli Date: Wed, 27 May 2026 10:11:16 +0000 Subject: [PATCH 1/7] change moe --- aiter/aot/flydsl/moe.py | 15 +- .../flydsl/kernels/mixed_moe_gemm_2stage.py | 218 ++++++++---------- aiter/ops/flydsl/kernels/moe_gemm_2stage.py | 183 +++++++-------- aiter/ops/flydsl/kernels/preshuffle_gemm.py | 18 +- aiter/ops/flydsl/kernels/silu_and_mul_fq.py | 47 ++-- aiter/ops/flydsl/moe_kernels.py | 133 ++++++----- 6 files changed, 305 insertions(+), 309 deletions(-) diff --git a/aiter/aot/flydsl/moe.py b/aiter/aot/flydsl/moe.py index c7e7e95adf..655bd769ea 100644 --- a/aiter/aot/flydsl/moe.py +++ b/aiter/aot/flydsl/moe.py @@ -38,6 +38,7 @@ from aiter.jit.core import AITER_CONFIGS from aiter.ops.flydsl.moe_kernels import ( _get_compiled_silu_fused, + _ptr_view_safe, _run_compiled, _s1_args_fp4, _s1_args_std, @@ -579,13 +580,13 @@ def _make_a_user(a_dtype_user_shape): _run_compiled( silu_fused, ( - tmp_out.view(-1, inter_dim * 2), - out.view(-1).view(torch.uint8), - out_scale_sorted_flat, - sorted_token_ids, - num_valid_ids, - sorted_token_ids.view(-1), - torch.empty(0, device=dev, dtype=torch.float32), + _ptr_view_safe(tmp_out.view(-1, inter_dim * 2)), + _ptr_view_safe(out.view(-1).view(torch.uint8)), + _ptr_view_safe(out_scale_sorted_flat), + _ptr_view_safe(sorted_token_ids), + _ptr_view_safe(num_valid_ids), + _ptr_view_safe(sorted_token_ids.view(-1)), + _ptr_view_safe(torch.empty(0, device=dev, dtype=torch.float32)), tokens, sorted_token_ids.shape[0], 0, diff --git a/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py b/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py index e2bc48328b..124a2b25ce 100644 --- a/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py +++ b/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py @@ -325,6 +325,10 @@ def x_lds_elem(): kpack_bytes = 8 if is_int4 else 16 out_elem_bytes = 4 if out_is_f32 else 2 + w_elem_bytes = 2 if is_f16_b else 1 + w_elem_pack = 2 if (is_f4_b or is_int4) else 1 + w_nbytes = (experts * (2 * inter_dim) * model_dim * w_elem_bytes) // w_elem_pack + bias_nbytes = experts * (2 * inter_dim) * 4 _e_vec_s1 = min(tile_n // 32, 8) if _need_quant: @@ -431,17 +435,17 @@ def x_lds_elem(): @flyc.kernel(name=module_name) def moe_gemm1( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_num_valid_ids: fx.Tensor, - arg_bias: fx.Tensor, - arg_out_scale_sorted: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_num_valid_ids: fx.Pointer, + arg_bias: fx.Pointer, + arg_out_scale_sorted: fx.Pointer, i32_tokens_in: fx.Int32, i32_n_in: fx.Int32, i32_k_in: fx.Int32, @@ -464,6 +468,13 @@ def moe_gemm1( vec16_x = T.vec(vec16_elems, x_elem) vec2_i64 = T.vec(2, i64) + def _ptr_buffer_resource(ptr, num_records_bytes): + addr = fx.ptrtoint(ptr) + addr_i64 = arith.index_cast(T.i64, addr) + return buffer_ops.create_buffer_resource_from_addr( + addr_i64, num_records_bytes=num_records_bytes + ) + acc_init = arith.constant_vector(0.0, vec4_f32) # --- Stage1 dimension mapping --- @@ -605,18 +616,12 @@ def moe_gemm1( # X: [tokens, model_dim] x_nbytes_idx = (tokens_in * k_in * c_elem_bytes) / c_a_pack x_nbytes_i32 = arith.index_cast(T.i32, x_nbytes_idx) - x_rsrc = buffer_ops.create_buffer_resource( - arg_x, max_size=False, num_records_bytes=x_nbytes_i32 - ) + x_rsrc = _ptr_buffer_resource(arg_x, x_nbytes_i32) - w_rsrc = buffer_ops.create_buffer_resource(arg_w, max_size=False) + w_rsrc = _ptr_buffer_resource(arg_w, w_nbytes) # Out: [tokens*topk, inter_dim] - numids_rsrc = buffer_ops.create_buffer_resource( - arg_num_valid_ids, - max_size=False, - num_records_bytes=arith.constant(4, type=T.i32), - ) + numids_rsrc = _ptr_buffer_resource(arg_num_valid_ids, arith.constant(4, type=T.i32)) num_valid_i32 = buffer_ops.buffer_load( numids_rsrc, arith.constant(0, index=True), vec_width=1, dtype=T.i32 ) @@ -629,9 +634,7 @@ def moe_gemm1( kblk = k_in / c32 sx_nbytes_idx = sorted_m * kblk sx_nbytes_i32 = arith.index_cast(T.i32, sx_nbytes_idx) - sx_rsrc = buffer_ops.create_buffer_resource( - arg_scale_x, max_size=False, num_records_bytes=sx_nbytes_i32 - ) + sx_rsrc = _ptr_buffer_resource(arg_scale_x, sx_nbytes_i32) if const_expr(not is_f16_b): c32 = arith.constant(32, index=True) @@ -639,30 +642,20 @@ def moe_gemm1( mn_w = arith.constant(experts * (2 * inter_dim), index=True) sw_nbytes_idx = mn_w * kblk_w sw_nbytes_i32 = arith.index_cast(T.i32, sw_nbytes_idx) - sw_rsrc = buffer_ops.create_buffer_resource( - arg_scale_w, max_size=False, num_records_bytes=sw_nbytes_i32 - ) + sw_rsrc = _ptr_buffer_resource(arg_scale_w, sw_nbytes_i32) sorted_nbytes_idx = size_expert_ids_in * arith.constant( sort_block_m * 4, index=True ) sorted_nbytes_i32 = arith.index_cast(T.i32, sorted_nbytes_idx) - sorted_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_token_ids, - max_size=False, - num_records_bytes=sorted_nbytes_i32, - ) - sorted_w_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_weights, max_size=False, num_records_bytes=sorted_nbytes_i32 - ) + sorted_rsrc = _ptr_buffer_resource(arg_sorted_token_ids, sorted_nbytes_i32) + sorted_w_rsrc = _ptr_buffer_resource(arg_sorted_weights, sorted_nbytes_i32) eid_nbytes_idx = size_expert_ids_in * arith.constant(4, index=True) eid_nbytes_i32 = arith.index_cast(T.i32, eid_nbytes_idx) - expert_rsrc = buffer_ops.create_buffer_resource( - arg_expert_ids, max_size=False, num_records_bytes=eid_nbytes_i32 - ) + expert_rsrc = _ptr_buffer_resource(arg_expert_ids, eid_nbytes_i32) bias_rsrc = ( - buffer_ops.create_buffer_resource(arg_bias, max_size=False) + _ptr_buffer_resource(arg_bias, bias_nbytes) if enable_bias else None ) @@ -672,9 +665,21 @@ def moe_gemm1( _sorted_scale_cols_i32 = arith.constant(_sorted_scale_cols, type=T.i32) sorted_scale_rsrc = None if const_expr(_need_sort): - sorted_scale_rsrc = buffer_ops.create_buffer_resource( - arg_out_scale_sorted, max_size=False + _sort_rows_idx = size_expert_ids_in * arith.constant( + sort_block_m, index=True + ) + _sort_padded_rows = ( + (_sort_rows_idx + arith.constant(255, index=True)) + / arith.constant(256, index=True) + * arith.constant(256, index=True) + ) + _sort_padded_cols = arith.constant( + ((_sorted_scale_cols + 7) // 8) * 8, index=True ) + _sort_scale_nbytes = arith.index_cast( + T.i32, _sort_padded_rows * _sort_padded_cols + ) + sorted_scale_rsrc = _ptr_buffer_resource(arg_out_scale_sorted, _sort_scale_nbytes) # ---- persist_m loop (same pattern as stage2) ---- _PERSIST_M = persist_m @@ -2097,13 +2102,7 @@ def _act_vec4(gate_v4, up_v4): topk_i32_v = topk_i32 tokens_i32_v = tokens_i32 - from flydsl._mlir.dialects import fly as _fly - - _llvm_ptr_ty = ir.Type.parse("!llvm.ptr") - out_base_ptr = _fly.extract_aligned_pointer_as_index( - _llvm_ptr_ty, arg_out - ) - out_base_i64 = llvm.ptrtoint(T.i64, out_base_ptr) + out_base_i64 = arith.index_cast(T.i64, fx.ptrtoint(arg_out)) out_base_idx = arith.index_cast(ir.IndexType.get(), out_base_i64) if const_expr(lds_out is None): @@ -2672,17 +2671,17 @@ def store_pair(*, row_local, row, row_ctx, col_pair0, col_g0, frag): @flyc.jit def launch_mixed_moe_gemm1( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_max_token_ids: fx.Tensor, - arg_bias: fx.Tensor, - arg_out_scale_sorted: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_max_token_ids: fx.Pointer, + arg_bias: fx.Pointer, + arg_out_scale_sorted: fx.Pointer, i32_tokens_in: fx.Int32, i32_inter_in: fx.Int32, i32_k_in: fx.Int32, @@ -2874,6 +2873,10 @@ def compile_mixed_moe_gemm2( "compile_moe_gemm2(accumulate=False) only supports out_dtype in {'f16','bf16'}" ) is_int4 = b_dtype == "int4" + w_elem_bytes = 2 if is_f16_b else 1 + w_elem_pack = 2 if (is_f4_b or is_int4) else 1 + w_nbytes = (experts * model_dim * inter_dim * w_elem_bytes) // w_elem_pack + bias_nbytes = experts * model_dim * 4 # INT4 here means W4A8: A2 is int8, W is packed int4 and unpacked to int8 in-kernel. is_int8 = False @@ -3011,16 +3014,16 @@ def x_lds_elem(): @flyc.kernel(name=module_name) def moe_gemm2( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_num_valid_ids: fx.Tensor, - arg_bias: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_num_valid_ids: fx.Pointer, + arg_bias: fx.Pointer, i32_tokens_in: fx.Int32, i32_n_in: fx.Int32, i32_k_in: fx.Int32, @@ -3045,6 +3048,13 @@ def moe_gemm2( vec16_x = T.vec(vec16_elems, x_elem) vec2_i64 = T.vec(2, i64) + def _ptr_buffer_resource(ptr, num_records_bytes): + addr = fx.ptrtoint(ptr) + addr_i64 = arith.index_cast(T.i64, addr) + return buffer_ops.create_buffer_resource_from_addr( + addr_i64, num_records_bytes=num_records_bytes + ) + acc_init = ( arith.constant_vector(0, vec4_i32) if is_int8 @@ -3156,11 +3166,9 @@ def check_c_k_valid_gate(base_k): (tokens_in * c_topk) * k_in * c_elem_bytes, int(a_elem_vec_pack) ) x_nbytes_i32 = arith.index_cast(T.i32, x_nbytes_idx) - x_rsrc = buffer_ops.create_buffer_resource( - arg_x, max_size=False, num_records_bytes=x_nbytes_i32 - ) + x_rsrc = _ptr_buffer_resource(arg_x, x_nbytes_i32) - w_rsrc = buffer_ops.create_buffer_resource(arg_w, max_size=False) + w_rsrc = _ptr_buffer_resource(arg_w, w_nbytes) # OUT: [tokens, model_dim] -> clamp to descriptor max (i32 bytes) to avoid overflow on huge tokens. out_elem_bytes = 4 if out_is_f32 else 2 @@ -3175,16 +3183,10 @@ def check_c_k_valid_gate(base_k): * arith.constant(out_elem_bytes, index=True) ) out_nbytes_i32 = arith.index_cast(T.i32, out_nbytes_idx) - out_rsrc = buffer_ops.create_buffer_resource( - arg_out, max_size=False, num_records_bytes=out_nbytes_i32 - ) + out_rsrc = _ptr_buffer_resource(arg_out, out_nbytes_i32) # num_valid_ids (sorted padded MN) for scale sizing / guards. - numids_rsrc = buffer_ops.create_buffer_resource( - arg_num_valid_ids, - max_size=False, - num_records_bytes=arith.constant(4, type=T.i32), - ) + numids_rsrc = _ptr_buffer_resource(arg_num_valid_ids, arith.constant(4, type=T.i32)) num_valid_i32 = buffer_ops.buffer_load( numids_rsrc, arith.constant(0, index=True), vec_width=1, dtype=T.i32 ) @@ -3205,16 +3207,12 @@ def check_c_k_valid_gate(base_k): kblk = _div_pow2(k_in, 32) sx_nbytes_idx = num_valid_idx * kblk sx_nbytes_i32 = arith.index_cast(T.i32, sx_nbytes_idx) - sx_rsrc = buffer_ops.create_buffer_resource( - arg_scale_x, max_size=False, num_records_bytes=sx_nbytes_i32 - ) + sx_rsrc = _ptr_buffer_resource(arg_scale_x, sx_nbytes_i32) else: # scale_x (A2 scale): [tokens*topk] f32 -> bytes = tokens*topk*4 sx_nbytes_idx = (tokens_in * c_topk) * arith.constant(4, index=True) sx_nbytes_i32 = arith.index_cast(T.i32, sx_nbytes_idx) - sx_rsrc = buffer_ops.create_buffer_resource( - arg_scale_x, max_size=False, num_records_bytes=sx_nbytes_i32 - ) + sx_rsrc = _ptr_buffer_resource(arg_scale_x, sx_nbytes_i32) if const_expr(not is_f16_b): # Weight microscale buffer (packed i32 holding e8m0 bytes). @@ -3223,9 +3221,7 @@ def check_c_k_valid_gate(base_k): mn_w = arith.constant(experts * model_dim, index=True) sw_nbytes_idx = mn_w * kblk_w # bytes (e8m0) sw_nbytes_i32 = arith.index_cast(T.i32, sw_nbytes_idx) - sw_rsrc = buffer_ops.create_buffer_resource( - arg_scale_w, max_size=False, num_records_bytes=sw_nbytes_i32 - ) + sw_rsrc = _ptr_buffer_resource(arg_scale_w, sw_nbytes_i32) # sorted_token_ids / sorted_weights: [blocks*tile_m] (padded length) sorted_nbytes_idx = ( @@ -3234,14 +3230,8 @@ def check_c_k_valid_gate(base_k): * arith.constant(4, index=True) ) sorted_nbytes_i32 = arith.index_cast(T.i32, sorted_nbytes_idx) - sorted_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_token_ids, - max_size=False, - num_records_bytes=sorted_nbytes_i32, - ) - sorted_w_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_weights, max_size=False, num_records_bytes=sorted_nbytes_i32 - ) + sorted_rsrc = _ptr_buffer_resource(arg_sorted_token_ids, sorted_nbytes_i32) + sorted_w_rsrc = _ptr_buffer_resource(arg_sorted_weights, sorted_nbytes_i32) # expert ids: [sort_blocks] i32. _c_sbm = arith.constant(_sort_block_m, index=True) @@ -3252,11 +3242,9 @@ def check_c_k_valid_gate(base_k): ) eid_nbytes_idx = _sort_blocks_ub * arith.constant(4, index=True) eid_nbytes_i32 = arith.index_cast(T.i32, eid_nbytes_idx) - expert_rsrc = buffer_ops.create_buffer_resource( - arg_expert_ids, max_size=False, num_records_bytes=eid_nbytes_i32 - ) + expert_rsrc = _ptr_buffer_resource(arg_expert_ids, eid_nbytes_i32) bias_rsrc = ( - buffer_ops.create_buffer_resource(arg_bias, max_size=False) + _ptr_buffer_resource(arg_bias, bias_nbytes) if enable_bias else None ) @@ -4306,13 +4294,7 @@ def atomic_add_f16x2(val_f16x2, byte_off_i32): # Both accumulate=True (global atomic) and accumulate=False (global store) # need 64-bit addressing to avoid i32 offset overflow when # tokens * model_dim * elem_bytes > INT32_MAX (~150K tokens for model_dim=7168). - from flydsl._mlir.dialects import fly as _fly - - _llvm_ptr_ty = ir.Type.parse("!llvm.ptr") - out_base_ptr = _fly.extract_aligned_pointer_as_index( - _llvm_ptr_ty, arg_out - ) - out_base_i64 = llvm.ptrtoint(T.i64, out_base_ptr) + out_base_i64 = arith.index_cast(T.i64, fx.ptrtoint(arg_out)) out_base_idx = arith.index_cast(ir.IndexType.get(), out_base_i64) def write_row_to_lds( @@ -4513,16 +4495,16 @@ def store_pair(*, row_local, row, row_ctx, col_pair0, col_g0, frag): @flyc.jit def launch_mixed_moe_gemm2( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_num_valid_ids: fx.Tensor, - arg_bias: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_num_valid_ids: fx.Pointer, + arg_bias: fx.Pointer, i32_tokens_in: fx.Int32, i32_n_in: fx.Int32, i32_k_in: fx.Int32, diff --git a/aiter/ops/flydsl/kernels/moe_gemm_2stage.py b/aiter/ops/flydsl/kernels/moe_gemm_2stage.py index 08ab740ac6..25995f4a1e 100644 --- a/aiter/ops/flydsl/kernels/moe_gemm_2stage.py +++ b/aiter/ops/flydsl/kernels/moe_gemm_2stage.py @@ -249,10 +249,15 @@ def out_mlir(): ir.ShapedType.get_dynamic_size() # W is packed int4 for W4A8/W4A16/W4A_FP8: 2 values per byte. - ( + w_nbytes = ( (experts * (2 * inter_dim) * model_dim) // 2 if w_is_int4 - else (experts * (2 * inter_dim) * model_dim) + else (experts * (2 * inter_dim) * model_dim * elem_bytes) + ) + sw_nbytes = ( + experts * (2 * inter_dim) * num_groups * (2 if _scale_is_bf16 else 4) + if needs_scale_w + else 0 ) total_threads = 256 @@ -334,15 +339,15 @@ def out_mlir(): @flyc.kernel def moe_gemm1( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_max_token_ids: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_max_token_ids: fx.Pointer, i32_tokens_in: fx.Int32, i32_inter_in: fx.Int32, i32_k_in: fx.Int32, @@ -376,6 +381,13 @@ def moe_gemm1( vec8_x = T.vec(vec8_elems, x_elem) vec16_x = T.vec(vec16_elems, x_elem) + def _ptr_buffer_resource(ptr, num_records_bytes): + addr = fx.ptrtoint(ptr) + addr_i64 = arith.index_cast(T.i64, addr) + return buffer_ops.create_buffer_resource_from_addr( + addr_i64, num_records_bytes=num_records_bytes + ) + def silu(x): # device fast path: # emu = exp(-x) ~= exp2(log2e * (-x)) -> v_exp_f32 @@ -437,11 +449,7 @@ def silu(x): # Block validity: compute as early as possible so invalid blocks skip all buffer-resource # setup, LDS pointer math, and gmem prefetch work. bx_m = bx * fx.Index(tile_m) - maxids_rsrc = buffer_ops.create_buffer_resource( - arg_max_token_ids, - max_size=False, - num_records_bytes=fx.Index(4), - ) + maxids_rsrc = _ptr_buffer_resource(arg_max_token_ids, fx.Index(4)) max_token_id_i32 = buffer_ops.buffer_load( maxids_rsrc, fx.Index(0), vec_width=1, dtype=T.i32 ) @@ -494,11 +502,9 @@ def silu(x): # X: [tokens, k] bytes = tokens*k*elem_bytes x_rows = tokens_in * (c_topk if x_is_token_slot else fx.Index(1)) x_nbytes_idx = x_rows * k_in * arith.index(int(elem_bytes)) - x_rsrc = buffer_ops.create_buffer_resource( - arg_x, max_size=False, num_records_bytes=x_nbytes_idx - ) + x_rsrc = _ptr_buffer_resource(arg_x, x_nbytes_idx) - w_rsrc = buffer_ops.create_buffer_resource(arg_w, max_size=False) + w_rsrc = _ptr_buffer_resource(arg_w, w_nbytes) # OUT: normal=[tokens, topk, inter] f16/bf16, # split-K=[tokens*topk, 2*inter] f32 (or bf16 for bf16 split-K) @@ -511,9 +517,7 @@ def silu(x): out_nbytes_idx = ( tokens_in * c_topk * inter_in * fx.Index(out_elem_bytes) ) - out_rsrc = buffer_ops.create_buffer_resource( - arg_out, max_size=False, num_records_bytes=out_nbytes_idx - ) + out_rsrc = _ptr_buffer_resource(arg_out, out_nbytes_idx) # scale_x: fp16/bf16 path ignores (implicit scale=1.0); int4_bf16 also uses 1.0. if const_expr(is_f16_or_bf16): @@ -521,30 +525,19 @@ def silu(x): else: sx_rows = tokens_in * (c_topk if x_is_token_slot else fx.Index(1)) sx_nbytes_idx = sx_rows * fx.Index(4) - sx_rsrc = buffer_ops.create_buffer_resource( - arg_scale_x, max_size=False, num_records_bytes=sx_nbytes_idx - ) + sx_rsrc = _ptr_buffer_resource(arg_scale_x, sx_nbytes_idx) # scale_w: fp16/bf16 (non-int4) path ignores; int4_bf16 needs dequant scale. if const_expr(not needs_scale_w): sw_rsrc = None else: - sw_rsrc = buffer_ops.create_buffer_resource( - arg_scale_w, max_size=False - ) + sw_rsrc = _ptr_buffer_resource(arg_scale_w, sw_nbytes) - sorted_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_token_ids, max_size=False - ) - sorted_w_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_weights, max_size=False - ) + sorted_nbytes_idx = size_expert_ids_in * fx.Index(tile_m) * fx.Index(4) + sorted_rsrc = _ptr_buffer_resource(arg_sorted_token_ids, sorted_nbytes_idx) + sorted_w_rsrc = _ptr_buffer_resource(arg_sorted_weights, sorted_nbytes_idx) # expert ids: [blocks] i32 -> bytes = size_expert_ids_in*4 - expert_rsrc = buffer_ops.create_buffer_resource( - arg_expert_ids, - max_size=False, - num_records_bytes=(size_expert_ids_in * fx.Index(4)), - ) + expert_rsrc = _ptr_buffer_resource(arg_expert_ids, size_expert_ids_in * fx.Index(4)) # Expert id for this M tile (keep address math in `index`) expert_i32 = buffer_ops.buffer_load( @@ -1481,7 +1474,7 @@ def _unflatten_b_tile(vals): _splitk_use_bf16 and not _has_buffer_atomic_bf16_s1 ) - out_base_idx = buffer_ops.extract_base_index(arg_out) + out_base_idx = arith.index_cast(T.index, fx.ptrtoint(arg_out)) _split_k_out_row_stride = ( inter_dim * 2 * out_elem_bytes ) # bytes per row @@ -1953,15 +1946,15 @@ def _stage1_store_row(*, mi: int, ii: int, row_in_tile, row): # ── Host launcher (flyc.jit + .launch) ──────────────────────────────── @flyc.jit def launch_moe_gemm1( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_max_token_ids: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_max_token_ids: fx.Pointer, i32_tokens_in: fx.Int32, i32_inter_in: fx.Int32, i32_k_in: fx.Int32, @@ -2131,10 +2124,15 @@ def compile_moe_gemm2( ir.ShapedType.get_dynamic_size() # W is packed int4 for W4A8/W4A16/W4A_FP8: 2 values per byte. - ( + w_nbytes = ( (experts * model_dim * inter_dim) // 2 if w_is_int4 - else (experts * model_dim * inter_dim) + else (experts * model_dim * inter_dim * elem_bytes) + ) + sw_nbytes = ( + experts * model_dim * num_groups * (2 if _scale_is_bf16 else 4) + if needs_scale_w + else 0 ) total_threads = 256 @@ -2249,15 +2247,15 @@ def out_elem(): @flyc.kernel def moe_gemm2( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_num_valid_ids: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_num_valid_ids: fx.Pointer, i32_tokens_in: fx.Int32, i32_n_in: fx.Int32, i32_k_in: fx.Int32, @@ -2290,6 +2288,13 @@ def moe_gemm2( vec8_x = T.vec(vec8_elems, x_elem) vec16_x = T.vec(vec16_elems, x_elem) + def _ptr_buffer_resource(ptr, num_records_bytes): + addr = fx.ptrtoint(ptr) + addr_i64 = arith.index_cast(T.i64, addr) + return buffer_ops.create_buffer_resource_from_addr( + addr_i64, num_records_bytes=num_records_bytes + ) + acc_init = ( arith.constant_vector(0, T.i32x4) if is_int8 @@ -2368,11 +2373,9 @@ def moe_gemm2( # X(A2): [tokens*topk, inter_dim] bytes = tokens*topk*k*elem_bytes x_nbytes_idx = (tokens_in * c_topk) * k_in * arith.index(int(elem_bytes)) - x_rsrc = buffer_ops.create_buffer_resource( - arg_x, max_size=False, num_records_bytes=x_nbytes_idx - ) + x_rsrc = _ptr_buffer_resource(arg_x, x_nbytes_idx) - w_rsrc = buffer_ops.create_buffer_resource(arg_w, max_size=False) + w_rsrc = _ptr_buffer_resource(arg_w, w_nbytes) # OUT: [tokens, model_dim] -> clamp to descriptor max (i32 bytes) to avoid overflow on huge tokens. out_elem_bytes = 4 if out_is_f32 else 2 @@ -2381,50 +2384,34 @@ def moe_gemm2( out_nbytes_idx = ( tokens_in * fx.Index(topk) * n_in * fx.Index(out_elem_bytes) ) - out_rsrc = buffer_ops.create_buffer_resource( - arg_out, max_size=False, num_records_bytes=out_nbytes_idx - ) + out_rsrc = _ptr_buffer_resource(arg_out, out_nbytes_idx) # scale_x: fp16/bf16 path ignores (implicit scale=1.0); int4_bf16 also uses 1.0. if const_expr(is_f16_or_bf16): sx_rsrc = None else: # scale_x (A2 scale): [tokens*topk] f32 -> bytes = tokens*topk*4 sx_nbytes_idx = (tokens_in * c_topk) * fx.Index(4) - sx_rsrc = buffer_ops.create_buffer_resource( - arg_scale_x, max_size=False, num_records_bytes=sx_nbytes_idx - ) + sx_rsrc = _ptr_buffer_resource(arg_scale_x, sx_nbytes_idx) # scale_w: fp16/bf16 (non-int4) path ignores; int4_bf16 needs dequant scale. if const_expr(not needs_scale_w): sw_rsrc = None else: # scale_w: [experts*model_dim] f32 (static shape in practice) - sw_rsrc = buffer_ops.create_buffer_resource(arg_scale_w, max_size=False) + sw_rsrc = _ptr_buffer_resource(arg_scale_w, sw_nbytes) # sorted_token_ids / sorted_weights: [blocks*tile_m] (CK-style padded length) sorted_nbytes_idx = size_expert_ids_in * fx.Index(tile_m) * fx.Index(4) - sorted_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_token_ids, - max_size=False, - num_records_bytes=sorted_nbytes_idx, - ) - sorted_w_rsrc = buffer_ops.create_buffer_resource( - arg_sorted_weights, max_size=False, num_records_bytes=sorted_nbytes_idx - ) + sorted_rsrc = _ptr_buffer_resource(arg_sorted_token_ids, sorted_nbytes_idx) + sorted_w_rsrc = _ptr_buffer_resource(arg_sorted_weights, sorted_nbytes_idx) # expert ids: [blocks] i32 -> bytes = size_expert_ids_in*4 eid_nbytes_idx = size_expert_ids_in * fx.Index(4) - expert_rsrc = buffer_ops.create_buffer_resource( - arg_expert_ids, max_size=False, num_records_bytes=eid_nbytes_idx - ) + expert_rsrc = _ptr_buffer_resource(arg_expert_ids, eid_nbytes_idx) bx_m = bx * fx.Index(tile_m) # Early-exit guard (as in 2ce65fb): some routing paths can produce extra/garbage # expert blocks beyond `num_valid_ids`. Skip those blocks entirely to avoid OOB. - numids_rsrc = buffer_ops.create_buffer_resource( - arg_num_valid_ids, - max_size=False, - num_records_bytes=fx.Index(4), - ) + numids_rsrc = _ptr_buffer_resource(arg_num_valid_ids, fx.Index(4)) num_valid_i32 = buffer_ops.buffer_load( numids_rsrc, fx.Index(0), vec_width=1, dtype=T.i32 ) @@ -3384,7 +3371,7 @@ def _stage2_row_atomic(*, mi: int, ii: int, row_in_tile, row): # gfx950+ has buffer_atomic_pk_add_bf16, so bf16 uses buffer atomics there. out_base_idx = None if const_expr(_needs_global_atomic_bf16): - out_base_idx = buffer_ops.extract_base_index(arg_out) + out_base_idx = arith.index_cast(T.index, fx.ptrtoint(arg_out)) def write_row_to_lds( *, @@ -3544,15 +3531,15 @@ def store_pair(*, row_local, row, row_ctx, col_pair0, col_g0, frag): # ── Host launcher (flyc.jit + .launch) ──────────────────────────────── @flyc.jit def launch_moe_gemm2( - arg_out: fx.Tensor, - arg_x: fx.Tensor, - arg_w: fx.Tensor, - arg_scale_x: fx.Tensor, - arg_scale_w: fx.Tensor, - arg_sorted_token_ids: fx.Tensor, - arg_expert_ids: fx.Tensor, - arg_sorted_weights: fx.Tensor, - arg_num_valid_ids: fx.Tensor, + arg_out: fx.Pointer, + arg_x: fx.Pointer, + arg_w: fx.Pointer, + arg_scale_x: fx.Pointer, + arg_scale_w: fx.Pointer, + arg_sorted_token_ids: fx.Pointer, + arg_expert_ids: fx.Pointer, + arg_sorted_weights: fx.Pointer, + arg_num_valid_ids: fx.Pointer, i32_tokens_in: fx.Int32, i32_n_in: fx.Int32, i32_k_in: fx.Int32, diff --git a/aiter/ops/flydsl/kernels/preshuffle_gemm.py b/aiter/ops/flydsl/kernels/preshuffle_gemm.py index 13f6d8c36d..38333e97a1 100644 --- a/aiter/ops/flydsl/kernels/preshuffle_gemm.py +++ b/aiter/ops/flydsl/kernels/preshuffle_gemm.py @@ -456,11 +456,19 @@ def kernel_gemm( arg_c, max_size=False, num_records_bytes=_c_nrec ) _needs_per_token_scale = not is_f16_or_bf16 and not is_fp4 - scale_a_rsrc = ( - None - if (is_f16_or_bf16) - else buffer_ops.create_buffer_resource(arg_scale_a, max_size=False) - ) + scale_a_rsrc = None + if const_expr(not is_f16_or_bf16): + if const_expr(is_fp4): + _scale_a_rows = (c_m + fx.Index(31)) // fx.Index(32) + _scale_a_stride_elems = fx.Index((K // (32 * 4 * 2)) * 64) + _scale_a_nrec = fx.Int64( + _scale_a_rows * _scale_a_stride_elems * fx.Index(4) + ) + else: + _scale_a_nrec = fx.Int64(c_m * fx.Index(4)) + scale_a_rsrc = buffer_ops.create_buffer_resource( + arg_scale_a, max_size=False, num_records_bytes=_scale_a_nrec + ) # ---- Bias buffer resource (for fused epilogue) ---- # Use max_size=True so the buffer descriptor's size is taken from the diff --git a/aiter/ops/flydsl/kernels/silu_and_mul_fq.py b/aiter/ops/flydsl/kernels/silu_and_mul_fq.py index 2cc90924a9..93c67ac826 100644 --- a/aiter/ops/flydsl/kernels/silu_and_mul_fq.py +++ b/aiter/ops/flydsl/kernels/silu_and_mul_fq.py @@ -93,13 +93,13 @@ def build_silu_and_mul_fq_module( @flyc.kernel def silu_and_mul_fq_kernel( - x: fx.Tensor, - out_buf: fx.Tensor, - out_scale_sorted: fx.Tensor, - sorted_ids: fx.Tensor, - num_valid_ids: fx.Tensor, - topk_ids: fx.Tensor, - bias: fx.Tensor, + x: fx.Pointer, + out_buf: fx.Pointer, + out_scale_sorted: fx.Pointer, + sorted_ids: fx.Pointer, + num_valid_ids: fx.Pointer, + topk_ids: fx.Pointer, + bias: fx.Pointer, token_num: Int32, ): bid = fx.block_idx.x @@ -142,14 +142,19 @@ def silu_and_mul_fq_kernel( topk_i32 = arith.constant(topk, type=i32) n32_sort = scale_cols_i32 * c32_i32 - in_rsrc = buffer_ops.create_buffer_resource(x, max_size=True) - out_rsrc = buffer_ops.create_buffer_resource(out_buf, max_size=True) - scale_rsrc = buffer_ops.create_buffer_resource(out_scale_sorted, max_size=True) - tid_rsrc = buffer_ops.create_buffer_resource(sorted_ids, max_size=True) - nv_rsrc = buffer_ops.create_buffer_resource(num_valid_ids, max_size=True) + def _ptr_buffer_resource(ptr): + addr = fx.ptrtoint(ptr) + addr_i64 = arith.index_cast(T.i64, addr) + return buffer_ops.create_buffer_resource_from_addr(addr_i64) + + in_rsrc = _ptr_buffer_resource(x) + out_rsrc = _ptr_buffer_resource(out_buf) + scale_rsrc = _ptr_buffer_resource(out_scale_sorted) + tid_rsrc = _ptr_buffer_resource(sorted_ids) + nv_rsrc = _ptr_buffer_resource(num_valid_ids) if enable_bias: - topk_rsrc = buffer_ops.create_buffer_resource(topk_ids, max_size=True) - bias_rsrc = buffer_ops.create_buffer_resource(bias, max_size=True) + topk_rsrc = _ptr_buffer_resource(topk_ids) + bias_rsrc = _ptr_buffer_resource(bias) def _load_bias_scalar(offset): return buffer_ops.buffer_load(bias_rsrc, offset, vec_width=1, dtype=f32) @@ -551,13 +556,13 @@ def _f32_to_e2m1(qx_f32): @flyc.jit def launch_silu_and_mul_fq( - x: fx.Tensor, - out_buf: fx.Tensor, - out_scale_sorted: fx.Tensor, - sorted_ids: fx.Tensor, - num_valid_ids: fx.Tensor, - topk_ids: fx.Tensor, - bias: fx.Tensor, + x: fx.Pointer, + out_buf: fx.Pointer, + out_scale_sorted: fx.Pointer, + sorted_ids: fx.Pointer, + num_valid_ids: fx.Pointer, + topk_ids: fx.Pointer, + bias: fx.Pointer, token_num: fx.Int32, num_sorted_rows: fx.Int32, stream: fx.Stream = fx.Stream(None), diff --git a/aiter/ops/flydsl/moe_kernels.py b/aiter/ops/flydsl/moe_kernels.py index 357bffdcf0..306fcb416c 100644 --- a/aiter/ops/flydsl/moe_kernels.py +++ b/aiter/ops/flydsl/moe_kernels.py @@ -449,6 +449,19 @@ def _view_safe(t: torch.Tensor) -> torch.Tensor: ) +def _ptr_view_safe(t: torch.Tensor): + """Pass only the device data pointer; shape is carried by explicit args.""" + import flydsl.compiler as flyc + import flydsl.expr as fx + + view = _view_safe(t) + type_name = type(view).__name__ + module_name = type(view).__module__ + if type_name == "FakeTensor" or "fake_tensor" in module_name: + return flyc.from_c_void_p(fx.Uint8, 0) + return flyc.from_c_void_p(fx.Uint8, view.data_ptr()) + + def _s1_args_fp4( out, a, @@ -473,17 +486,17 @@ def _s1_args_fp4( if stream is None: stream = torch.cuda.current_stream() return ( - _view_safe(out), - _view_safe(a), - _view_safe(w), - _view_safe(a_scale), - _view_safe(w_scale), - sorted_ids, - sorted_expert_ids, - sorted_weights, - num_valid_ids, - _bias, - out_scale_sorted, + _ptr_view_safe(out), + _ptr_view_safe(a), + _ptr_view_safe(w), + _ptr_view_safe(a_scale), + _ptr_view_safe(w_scale), + _ptr_view_safe(sorted_ids), + _ptr_view_safe(sorted_expert_ids), + _ptr_view_safe(sorted_weights), + _ptr_view_safe(num_valid_ids), + _ptr_view_safe(_bias), + _ptr_view_safe(out_scale_sorted), token_num, n_in, k_in, @@ -511,15 +524,15 @@ def _s1_args_std( if stream is None: stream = torch.cuda.current_stream() return ( - out, - a, - w, - a_scale, - w_scale, - sorted_ids, - sorted_expert_ids, - sorted_weights, - num_valid_ids, + _ptr_view_safe(out), + _ptr_view_safe(a), + _ptr_view_safe(w), + _ptr_view_safe(a_scale), + _ptr_view_safe(w_scale), + _ptr_view_safe(sorted_ids), + _ptr_view_safe(sorted_expert_ids), + _ptr_view_safe(sorted_weights), + _ptr_view_safe(num_valid_ids), token_num, n_in, k_in, @@ -554,16 +567,16 @@ def _s2_args_fp4( if stream is None: stream = torch.cuda.current_stream() return ( - _view_safe(target), - _view_safe(a), - _view_safe(w), - _view_safe(a_scale), - _view_safe(w_scale), - sorted_ids, - sorted_expert_ids, - sorted_weights, - num_valid_ids, - _bias, + _ptr_view_safe(target), + _ptr_view_safe(a), + _ptr_view_safe(w), + _ptr_view_safe(a_scale), + _ptr_view_safe(w_scale), + _ptr_view_safe(sorted_ids), + _ptr_view_safe(sorted_expert_ids), + _ptr_view_safe(sorted_weights), + _ptr_view_safe(num_valid_ids), + _ptr_view_safe(_bias), token_num, n_in, k_in, @@ -591,15 +604,15 @@ def _s2_args_std( if stream is None: stream = torch.cuda.current_stream() return ( - target, - a, - w, - a_scale, - w_scale, - sorted_ids, - sorted_expert_ids, - sorted_weights, - num_valid_ids, + _ptr_view_safe(target), + _ptr_view_safe(a), + _ptr_view_safe(w), + _ptr_view_safe(a_scale), + _ptr_view_safe(w_scale), + _ptr_view_safe(sorted_ids), + _ptr_view_safe(sorted_expert_ids), + _ptr_view_safe(sorted_weights), + _ptr_view_safe(num_valid_ids), token_num, n_in, k_in, @@ -919,13 +932,13 @@ def flydsl_moe_stage1( _run_compiled( _silu_fused_k, ( - tmp_out.view(-1, inter_dim * 2), - out.view(-1).view(torch.uint8), - out_scale_sorted_flat, - sorted_token_ids, - num_valid_ids, - topk_ids_arg, - bias_arg, + _ptr_view_safe(tmp_out.view(-1, inter_dim * 2)), + _ptr_view_safe(out.view(-1).view(torch.uint8)), + _ptr_view_safe(out_scale_sorted_flat), + _ptr_view_safe(sorted_token_ids), + _ptr_view_safe(num_valid_ids), + _ptr_view_safe(topk_ids_arg), + _ptr_view_safe(bias_arg), token_num, num_sorted_rows, torch.cuda.current_stream(), @@ -944,13 +957,13 @@ def flydsl_moe_stage1( _run_compiled( _silu_fused_k, ( - tmp_out.view(-1, inter_dim * 2), - out.view(-1).view(torch.uint8), - out_scale_sorted_flat, - sorted_token_ids, - num_valid_ids, - topk_ids_arg, - bias_arg, + _ptr_view_safe(tmp_out.view(-1, inter_dim * 2)), + _ptr_view_safe(out.view(-1).view(torch.uint8)), + _ptr_view_safe(out_scale_sorted_flat), + _ptr_view_safe(sorted_token_ids), + _ptr_view_safe(num_valid_ids), + _ptr_view_safe(topk_ids_arg), + _ptr_view_safe(bias_arg), token_num, num_sorted_rows, torch.cuda.current_stream(), @@ -967,13 +980,13 @@ def flydsl_moe_stage1( _run_compiled( _silu_fused_k, ( - tmp_out.view(-1, inter_dim * 2), - out.view(-1).view(torch.uint8), - out_scale_sorted_flat, - sorted_token_ids, - num_valid_ids, - topk_ids_arg, - bias_arg, + _ptr_view_safe(tmp_out.view(-1, inter_dim * 2)), + _ptr_view_safe(out.view(-1).view(torch.uint8)), + _ptr_view_safe(out_scale_sorted_flat), + _ptr_view_safe(sorted_token_ids), + _ptr_view_safe(num_valid_ids), + _ptr_view_safe(topk_ids_arg), + _ptr_view_safe(bias_arg), token_num, num_sorted_rows, torch.cuda.current_stream(), From fa0776650ed939e8f326023eebec46aedc33b67d Mon Sep 17 00:00:00 2001 From: coderfeli Date: Wed, 27 May 2026 10:11:16 +0000 Subject: [PATCH 2/7] change gemm --- aiter/aot/flydsl/gemm.py | 29 ++++++++++- aiter/ops/flydsl/gemm_kernels.py | 33 ++++++++----- aiter/ops/flydsl/kernels/preshuffle_gemm.py | 53 +++++++++++---------- aiter/ops/flydsl/kernels/small_m_hgemm.py | 30 ++++++------ aiter/ops/flydsl/kernels/splitk_hgemm.py | 30 ++++++------ aiter/ops/flydsl/kernels/tensor_shim.py | 19 +++++--- 6 files changed, 118 insertions(+), 76 deletions(-) diff --git a/aiter/aot/flydsl/gemm.py b/aiter/aot/flydsl/gemm.py index d77ed02104..dfdf3bcf32 100644 --- a/aiter/aot/flydsl/gemm.py +++ b/aiter/aot/flydsl/gemm.py @@ -197,6 +197,12 @@ def _compile_executable_to_cache(exe, *args) -> None: exe(*args) +def _ptr_view_safe(t): + from aiter.ops.flydsl.gemm_kernels import _ptr_view_safe as _wrap + + return _wrap(t) + + def _compile_hgemm_to_cache( *, m: int, @@ -271,7 +277,15 @@ def _compile_hgemm_to_cache( # optional bias and split-K sync tensors. launch_bias = bias if has_bias else b _compile_executable_to_cache( - exe, out, a, b, launch_bias, m, semaphore, signal, stream + exe, + _ptr_view_safe(out), + _ptr_view_safe(a), + _ptr_view_safe(b), + _ptr_view_safe(launch_bias), + m, + _ptr_view_safe(semaphore), + _ptr_view_safe(signal), + stream, ) @@ -322,7 +336,18 @@ def _compile_preshuffle_to_cache( waves_per_eu=None if waves_per_eu <= 0 else waves_per_eu, xcd_swizzle=xcd_swizzle, ) - _compile_executable_to_cache(exe, out, a, b, scale_a, scale_b, bias, m, n, stream) + _compile_executable_to_cache( + exe, + _ptr_view_safe(out), + _ptr_view_safe(a), + _ptr_view_safe(b), + _ptr_view_safe(scale_a), + _ptr_view_safe(scale_b), + _ptr_view_safe(bias), + m, + n, + stream, + ) def compile_one_config( diff --git a/aiter/ops/flydsl/gemm_kernels.py b/aiter/ops/flydsl/gemm_kernels.py index bcbd4755f3..9c61ab8c77 100644 --- a/aiter/ops/flydsl/gemm_kernels.py +++ b/aiter/ops/flydsl/gemm_kernels.py @@ -14,6 +14,7 @@ from torch import Tensor import flydsl.expr as fx +import flydsl.compiler as flyc from aiter import logger from flydsl.runtime.device import get_rocm_arch @@ -67,6 +68,14 @@ def _get_dtypes(): SPLIT_K_GLOBAL_SEMAPHORE: dict[SplitKStreamKey, torch.Tensor] = {} SPLIT_K_GLOBAL_SIGNAL: dict[SplitKStreamKey, torch.Tensor] = {} + +def _ptr_view_safe(t: torch.Tensor): + type_name = type(t).__name__ + module_name = type(t).__module__ + if type_name == "FakeTensor" or "fake_tensor" in module_name: + return flyc.from_c_void_p(fx.Uint8, 0) + return flyc.from_c_void_p(fx.Uint8, t.data_ptr()) + # Keep the generic auto-generated catalog aligned with the upstream FlyDSL # reference tuning space. The wider local one-off search space introduced # gfx950-faulting candidates (for example tile_k=160 and tile_n=160/192), @@ -800,13 +809,13 @@ def launcher( semaphore, signal = _get_split_k_tensors(a.device, launch_stream) return _run_compiled( kernel, - out, - a, - b, - launch_bias, + _ptr_view_safe(out), + _ptr_view_safe(a), + _ptr_view_safe(b), + _ptr_view_safe(launch_bias), runtime_m, - semaphore, - signal, + _ptr_view_safe(semaphore), + _ptr_view_safe(signal), fx.Stream(launch_stream), ) @@ -1004,12 +1013,12 @@ def _as_i8(t): _dummy_bias = torch.empty(0, dtype=Out.dtype, device=Out.device) _run_compiled( exe, - out_contig.view(-1), - _as_i8(XQ.contiguous()).view(-1), - _as_i8(WQ.contiguous()).view(-1), - x_scale.contiguous().view(-1), - w_scale.contiguous().view(-1), - _dummy_bias, + _ptr_view_safe(out_contig.view(-1)), + _ptr_view_safe(_as_i8(XQ.contiguous()).view(-1)), + _ptr_view_safe(_as_i8(WQ.contiguous()).view(-1)), + _ptr_view_safe(x_scale.contiguous().view(-1)), + _ptr_view_safe(w_scale.contiguous().view(-1)), + _ptr_view_safe(_dummy_bias), m, n, fx.Stream(torch.cuda.current_stream()), diff --git a/aiter/ops/flydsl/kernels/preshuffle_gemm.py b/aiter/ops/flydsl/kernels/preshuffle_gemm.py index 38333e97a1..e1de7714ef 100644 --- a/aiter/ops/flydsl/kernels/preshuffle_gemm.py +++ b/aiter/ops/flydsl/kernels/preshuffle_gemm.py @@ -7,6 +7,7 @@ import flydsl.expr as fx from flydsl.compiler.kernel_function import CompilationContext from flydsl.expr import buffer_ops, const_expr, gpu, math, range_constexpr, rocdl +from flydsl.expr.typing import T from flydsl.runtime.device import get_rocm_arch as get_hip_arch from flydsl.utils.smem_allocator import SmemAllocator, SmemPtr from .mfma_epilogues import mfma_epilog @@ -344,12 +345,12 @@ def _out_elem(): @flyc.kernel def kernel_gemm( - arg_c: fx.Tensor, - arg_a: fx.Tensor, - arg_b: fx.Tensor, - arg_scale_a: fx.Tensor, - arg_scale_b: fx.Tensor, - arg_bias: fx.Tensor, + arg_c: fx.Pointer, + arg_a: fx.Pointer, + arg_b: fx.Pointer, + arg_scale_a: fx.Pointer, + arg_scale_b: fx.Pointer, + arg_bias: fx.Pointer, i32_m: fx.Int32, i32_n: fx.Int32, ): @@ -449,12 +450,18 @@ def kernel_gemm( # ---- Buffer resources (runtime byte sizes for OOB protection) ---- _a_nrec = fx.Int64(c_m * (K * elem_bytes // a_elem_vec_pack)) _c_nrec = fx.Int64(c_m * c_n * 2) - a_rsrc = buffer_ops.create_buffer_resource( - arg_a, max_size=False, num_records_bytes=_a_nrec - ) - c_rsrc = buffer_ops.create_buffer_resource( - arg_c, max_size=False, num_records_bytes=_c_nrec - ) + + def _ptr_buffer_resource(ptr, num_records_bytes=None): + addr = fx.ptrtoint(ptr) + addr_i64 = fx.arith.index_cast(T.i64, addr) + if num_records_bytes is None: + return buffer_ops.create_buffer_resource_from_addr(addr_i64) + return buffer_ops.create_buffer_resource_from_addr( + addr_i64, num_records_bytes=num_records_bytes + ) + + a_rsrc = _ptr_buffer_resource(arg_a, _a_nrec) + c_rsrc = _ptr_buffer_resource(arg_c, _c_nrec) _needs_per_token_scale = not is_f16_or_bf16 and not is_fp4 scale_a_rsrc = None if const_expr(not is_f16_or_bf16): @@ -466,9 +473,7 @@ def kernel_gemm( ) else: _scale_a_nrec = fx.Int64(c_m * fx.Index(4)) - scale_a_rsrc = buffer_ops.create_buffer_resource( - arg_scale_a, max_size=False, num_records_bytes=_scale_a_nrec - ) + scale_a_rsrc = _ptr_buffer_resource(arg_scale_a, _scale_a_nrec) # ---- Bias buffer resource (for fused epilogue) ---- # Use max_size=True so the buffer descriptor's size is taken from the @@ -476,12 +481,12 @@ def kernel_gemm( # size (was c_n * 2, which broke if out_dtype became fp32 etc.). bias_rsrc = None if const_expr(_has_bias): - bias_rsrc = buffer_ops.create_buffer_resource(arg_bias, max_size=True) - b_rsrc = buffer_ops.create_buffer_resource(arg_b, max_size=True) + bias_rsrc = _ptr_buffer_resource(arg_bias) + b_rsrc = _ptr_buffer_resource(arg_b) scale_b_rsrc = ( None if (is_f16_or_bf16) - else buffer_ops.create_buffer_resource(arg_scale_b, max_size=True) + else _ptr_buffer_resource(arg_scale_b) ) bx_m = bx * tile_m @@ -2139,12 +2144,12 @@ def prefetch_a0_pack( # ── Host launcher ────────────────────────────────────────────────────── @flyc.jit def launch_gemm( - arg_c: fx.Tensor, - arg_a: fx.Tensor, - arg_b: fx.Tensor, - arg_scale_a: fx.Tensor, - arg_scale_b: fx.Tensor, - arg_bias: fx.Tensor, + arg_c: fx.Pointer, + arg_a: fx.Pointer, + arg_b: fx.Pointer, + arg_scale_a: fx.Pointer, + arg_scale_b: fx.Pointer, + arg_bias: fx.Pointer, i32_m: fx.Int32, i32_n: fx.Int32, stream: fx.Stream, diff --git a/aiter/ops/flydsl/kernels/small_m_hgemm.py b/aiter/ops/flydsl/kernels/small_m_hgemm.py index 4ac7e83f6f..9f12464d61 100644 --- a/aiter/ops/flydsl/kernels/small_m_hgemm.py +++ b/aiter/ops/flydsl/kernels/small_m_hgemm.py @@ -538,13 +538,13 @@ def compile_small_m_hgemm_kernel( @flyc.kernel def small_m_hgemm_kernel( - C: fx.Tensor, - A: fx.Tensor, - B: fx.Tensor, - BIAS: fx.Tensor, + C: fx.Pointer, + A: fx.Pointer, + B: fx.Pointer, + BIAS: fx.Pointer, m: fx.Int32, - semaphore: fx.Tensor, - signal: fx.Tensor, + semaphore: fx.Pointer, + signal: fx.Pointer, ): dtype_ = get_dtype_in_kernel(dtype) _ptr_type = ir.Type.parse("!llvm.ptr<1>") @@ -657,8 +657,7 @@ def zero_c_tile(c_g, bias_g, tile_n_offset): scf.YieldOp([]) def get_llvm_ptr(ptr, offset, dtype_bytes): - base_ptr = fly.extract_aligned_pointer_as_index(_ptr_type, ptr) - base_ptr = llvm.PtrToIntOp(_i64_type, base_ptr).result + base_ptr = arith.index_cast(_i64_type, fx.ptrtoint(ptr)) byte_offset = arith.index_cast( T.i64, fx.Index(offset) * fx.Index(dtype_bytes) ) @@ -934,8 +933,7 @@ def block_mma_sync(a_frags, b_frags, c_frags): def store_split_k_tile(c_tensor, c_g, c_s, tile_n_offset): out_raw = c_tensor - out_base_ptr = fly.extract_aligned_pointer_as_index(_ptr_type, out_raw) - out_base_int = llvm.PtrToIntOp(_i64_type, out_base_ptr).result + out_base_int = arith.index_cast(_i64_type, fx.ptrtoint(out_raw)) for i in range_constexpr(LDG_REG_C_COUNT): global_tid = BLOCK_THREADS * i + tid m_local_idx = fx.Index(global_tid // LDG_C_X_THREADS) @@ -1353,13 +1351,13 @@ def hot_loop_scheduler(): @flyc.jit def launch_small_m_hgemm_kernel( - C: fx.Tensor, - A: fx.Tensor, - B: fx.Tensor, - BIAS: fx.Tensor, + C: fx.Pointer, + A: fx.Pointer, + B: fx.Pointer, + BIAS: fx.Pointer, m: fx.Int32, - semaphore: fx.Tensor, - signal: fx.Tensor, + semaphore: fx.Pointer, + signal: fx.Pointer, stream: fx.Stream = fx.Stream(None), ): allocator.finalized = False diff --git a/aiter/ops/flydsl/kernels/splitk_hgemm.py b/aiter/ops/flydsl/kernels/splitk_hgemm.py index 5621753497..30fb23a757 100644 --- a/aiter/ops/flydsl/kernels/splitk_hgemm.py +++ b/aiter/ops/flydsl/kernels/splitk_hgemm.py @@ -215,13 +215,13 @@ def compile_hgemm_kernel( @flyc.kernel(known_block_size=[BLOCK_THREADS, 1, 1]) def hgemm_kernel( - C: fx.Tensor, - A: fx.Tensor, - B: fx.Tensor, - BIAS: fx.Tensor, + C: fx.Pointer, + A: fx.Pointer, + B: fx.Pointer, + BIAS: fx.Pointer, m: fx.Int32, - semaphore: fx.Tensor, - signal: fx.Tensor, + semaphore: fx.Pointer, + signal: fx.Pointer, ): dtype_ = get_dtype_in_kernel(dtype) _ptr_type = ir.Type.parse("!llvm.ptr<1>") @@ -282,8 +282,7 @@ def swizzle_for_cache_reuse(pid): c_frags = [acc_init] * C_FRAGS_LEN def get_llvm_ptr(ptr, offset, dtype_bytes): - base_ptr = fly.extract_aligned_pointer_as_index(_ptr_type, ptr) - base_ptr = llvm.PtrToIntOp(_i64_type, base_ptr).result + base_ptr = arith.index_cast(_i64_type, fx.ptrtoint(ptr)) byte_offset = arith.index_cast( T.i64, fx.Index(offset) * fx.Index(dtype_bytes) ) @@ -863,8 +862,7 @@ def hot_loop_scheduler(): if const_expr(IS_SPLIT_K): split_k_barrier() out_raw = C - out_base_ptr = fly.extract_aligned_pointer_as_index(_ptr_type, out_raw) - out_base_int = llvm.PtrToIntOp(_i64_type, out_base_ptr).result + out_base_int = arith.index_cast(_i64_type, fx.ptrtoint(out_raw)) for i in range_constexpr(LDG_REG_C_COUNT): global_tid = BLOCK_THREADS * i + tid m_local_idx = fx.Index(global_tid // LDG_C_X_THREADS) @@ -943,13 +941,13 @@ def hot_loop_scheduler(): @flyc.jit def launch_hgemm_kernel( - C: fx.Tensor, - A: fx.Tensor, - B: fx.Tensor, - BIAS: fx.Tensor, + C: fx.Pointer, + A: fx.Pointer, + B: fx.Pointer, + BIAS: fx.Pointer, m: fx.Int32, - semaphore: fx.Tensor, - signal: fx.Tensor, + semaphore: fx.Pointer, + signal: fx.Pointer, stream: fx.Stream = fx.Stream(None), ): allocator.finalized = False diff --git a/aiter/ops/flydsl/kernels/tensor_shim.py b/aiter/ops/flydsl/kernels/tensor_shim.py index 571a7403ae..18cd7e5c83 100644 --- a/aiter/ops/flydsl/kernels/tensor_shim.py +++ b/aiter/ops/flydsl/kernels/tensor_shim.py @@ -12,7 +12,7 @@ from flydsl._mlir import ir from flydsl.expr.typing import T -from flydsl.expr import buffer_ops, range_constexpr, vector, arith +from flydsl.expr import buffer_ops, range_constexpr, vector, arith, ptrtoint def _run_compiled(exe, *args): @@ -293,8 +293,13 @@ def __init__( static_bytes_offset_i64=None, ): super().__init__(dtype, shape, stride, base_offset) + raw = extract_to_ir_values(memref)[0] if static_bytes_offset_i64 is None: - self.rsrc = buffer_ops.create_buffer_resource(memref, max_size=True) + if str(raw.type).startswith("!fly.ptr"): + base_i64 = arith.index_cast(T.i64, ptrtoint(memref)) + self.rsrc = buffer_ops.create_buffer_resource_from_addr(base_i64) + else: + self.rsrc = buffer_ops.create_buffer_resource(memref, max_size=True) else: array_base_i64 = self.get_llvm_ptr(memref, (static_bytes_offset_i64)) self.rsrc = buffer_ops.create_buffer_resource_from_addr(array_base_i64) @@ -313,10 +318,12 @@ def store(self, offset, value, vec_size=1): def get_llvm_ptr(self, ptr, bytes_offset_i64, ptr_type="!llvm.ptr<1>"): bytes_offset_i64 = arith.index_cast(T.i64, bytes_offset_i64) _ptr_type = ir.Type.parse(ptr_type) - base_ptr = fly.extract_aligned_pointer_as_index( - _ptr_type, extract_to_ir_values(ptr)[0] - ) - base_ptr = llvm.PtrToIntOp(T.i64, base_ptr).result + raw = extract_to_ir_values(ptr)[0] + if str(raw.type).startswith("!fly.ptr"): + base_ptr = arith.index_cast(T.i64, ptrtoint(ptr)) + else: + base_ptr = fly.extract_aligned_pointer_as_index(_ptr_type, raw) + base_ptr = llvm.PtrToIntOp(T.i64, base_ptr).result llvm_ptr = llvm.AddOp( base_ptr, bytes_offset_i64, llvm.IntegerOverflowFlags(0) ).result From 17c45b855eb14e27c143c081d2dd16a55e5a633b Mon Sep 17 00:00:00 2001 From: coderfeli Date: Wed, 27 May 2026 10:11:16 +0000 Subject: [PATCH 3/7] port to new fx.ptr. update flydsl version --- .github/requirements/triton-test.txt | 2 +- .../flydsl/kernels/flash_attn_func_gfx1201.py | 58 +++++--- aiter/ops/flydsl/kernels/moe_gemm_2stage.py | 133 +++++++----------- .../ops/flydsl/kernels/qk_norm_rope_quant.py | 43 +++--- pyproject.toml | 2 +- requirements.txt | 2 +- setup.py | 2 +- 7 files changed, 119 insertions(+), 123 deletions(-) diff --git a/.github/requirements/triton-test.txt b/.github/requirements/triton-test.txt index 3a4e74c349..15883a02f7 100644 --- a/.github/requirements/triton-test.txt +++ b/.github/requirements/triton-test.txt @@ -8,7 +8,7 @@ pybind11==3.0.1 ninja==1.11.1.4 psutil packaging -flydsl==0.1.8 +flydsl==v0.1.9.dev599 # Test deps. pandas==2.2.3 diff --git a/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py b/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py index 3d0bb48dee..adab48ad79 100644 --- a/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py +++ b/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py @@ -52,7 +52,6 @@ from flydsl.utils.smem_allocator import SmemAllocator, SmemPtr from flydsl._mlir import ir from flydsl._mlir.dialects import ( - fly as _fly, llvm as _llvm, memref as _memref, ) @@ -72,9 +71,10 @@ def _llvm_ptr_ty(): return ir.Type.parse("!llvm.ptr") -def _extract_aligned_pointer(tensor) -> ir.Value: - """Extract the aligned LLVM pointer from a FlyDSL tensor/memref.""" - return _fly.extract_aligned_pointer_as_index(_llvm_ptr_ty(), _llvm_value(tensor)) +def _pointer_to_llvm_ptr(ptr) -> ir.Value: + """Convert a FlyDSL pointer argument to the LLVM pointer used by raw loads.""" + ptr_i64 = arith.index_cast(T.i64, fx.ptrtoint(ptr)) + return _llvm.IntToPtrOp(_llvm_ptr_ty(), ptr_i64).result def _pointer_load(result_type: ir.Type, ptr: ir.Value) -> ir.Value: @@ -202,18 +202,18 @@ def build_flash_attn_func_module_primary( @flyc.kernel(known_block_size=[BLOCK_SIZE, 1, 1]) def flash_attn_func_kernel( - Q: fx.Tensor, - K: fx.Tensor, - V: fx.Tensor, - O: fx.Tensor, # noqa: E741 + Q: fx.Pointer, + K: fx.Pointer, + V: fx.Pointer, + O: fx.Pointer, # noqa: E741 seq_len: fx.Int32, ): elem_type = dtype_to_elem_type(dtype_str) elem_dtype = elem_numeric_cls - q_ptr = _extract_aligned_pointer(Q) - k_ptr = _extract_aligned_pointer(K) - v_ptr = _extract_aligned_pointer(V) - o_ptr = _extract_aligned_pointer(O) + q_ptr = _pointer_to_llvm_ptr(Q) + k_ptr = _pointer_to_llvm_ptr(K) + v_ptr = _pointer_to_llvm_ptr(V) + o_ptr = _pointer_to_llvm_ptr(O) fm_fast = arith.FastMathFlags.fast # Local fast-math arithmetic helpers — preserve fastmath flag while using @@ -682,10 +682,10 @@ def _load_v_rowmajor(st_kv_base_val, pks_val, dc_val): @flyc.jit def launch_flash_attn_func( - Q: fx.Tensor, - K: fx.Tensor, - V: fx.Tensor, - O: fx.Tensor, # noqa: E741 + Q: fx.Pointer, + K: fx.Pointer, + V: fx.Pointer, + O: fx.Pointer, # noqa: E741 batch_size: fx.Int32, seq_len: fx.Int32, stream: fx.Stream = fx.Stream(None), @@ -756,7 +756,25 @@ def launch_flash_attn_func( "llvm_options": {"enable-post-misched": False, "lsr-drop-solution": True}, } + def _ptr_arg(t): + if hasattr(t, "data_ptr"): + type_name = type(t).__name__ + module_name = type(t).__module__ + ptr = 0 if type_name == "FakeTensor" or "fake_tensor" in module_name else t.data_ptr() + return flyc.from_c_void_p(fx.Uint8, ptr) + return t + + def _wrap_qkvo(args, kwargs): + args = list(args) + for idx in range(min(4, len(args))): + args[idx] = _ptr_arg(args[idx]) + for name in ("Q", "K", "V", "O"): + if name in kwargs: + kwargs[name] = _ptr_arg(kwargs[name]) + return tuple(args), kwargs + def _launch(*args, **kwargs): + args, kwargs = _wrap_qkvo(args, kwargs) with CompilationContext.compile_hints(_fmha_compile_hints): return launch_flash_attn_func(*args, **kwargs) @@ -764,10 +782,10 @@ def _compile(Q, K, V, O, batch_size, seq_len, stream=None): # noqa: E741 with CompilationContext.compile_hints(_fmha_compile_hints): return flyc.compile( launch_flash_attn_func, - Q, - K, - V, - O, + _ptr_arg(Q), + _ptr_arg(K), + _ptr_arg(V), + _ptr_arg(O), batch_size, seq_len, fx.Stream(stream), diff --git a/aiter/ops/flydsl/kernels/moe_gemm_2stage.py b/aiter/ops/flydsl/kernels/moe_gemm_2stage.py index 25995f4a1e..17784acbb2 100644 --- a/aiter/ops/flydsl/kernels/moe_gemm_2stage.py +++ b/aiter/ops/flydsl/kernels/moe_gemm_2stage.py @@ -3635,9 +3635,9 @@ def elem_type(): @flyc.kernel def moe_reduction_kernel( - X: fx.Tensor, - Y: fx.Tensor, - valid_mask: fx.Tensor, + X: fx.Pointer, + Y: fx.Pointer, + valid_mask: fx.Pointer, i32_m_tokens: fx.Int32, ): m_tokens = fx.Index(i32_m_tokens) @@ -3647,15 +3647,21 @@ def moe_reduction_kernel( elem_bits = 32 if dtype_str == "f32" else 16 copy_vec_width = 128 // elem_bits # 8 for f16/bf16, 4 for f32 n_sub = VEC_WIDTH // copy_vec_width # 1 for f16/bf16, 2 for f32 - # Buffer-backed tensors via layout API (all dtypes) - X_buf = fx.rocdl.make_buffer_tensor(X) - Y_buf = fx.rocdl.make_buffer_tensor(Y) - # Scalar buffer resources for tail path and mask - x_rsrc = buffer_ops.create_buffer_resource(X, max_size=True) - y_rsrc = buffer_ops.create_buffer_resource(Y, max_size=True) - mask_rsrc = buffer_ops.create_buffer_resource( - valid_mask, max_size=False, num_records_bytes=mask_nbytes_idx - ) + elem_nbytes_idx = fx.Index(4 if dtype_str == "f32" else 2) + + def _ptr_buffer_resource(ptr, num_records_bytes): + addr = fx.ptrtoint(ptr) + addr_i64 = arith.index_cast(T.i64, addr) + return buffer_ops.create_buffer_resource_from_addr( + addr_i64, num_records_bytes=num_records_bytes + ) + + x_nbytes = fx.Int64(m_tokens * c_topk * c_model_dim * elem_nbytes_idx) + y_nbytes = fx.Int64(m_tokens * c_model_dim * elem_nbytes_idx) + mask_nbytes = fx.Int64(mask_nbytes_idx) + x_rsrc = _ptr_buffer_resource(X, x_nbytes) + y_rsrc = _ptr_buffer_resource(Y, y_nbytes) + mask_rsrc = _ptr_buffer_resource(valid_mask, mask_nbytes) token_idx = gpu.block_id("x") tile_idx = gpu.block_id("y") @@ -3679,12 +3685,6 @@ def moe_reduction_kernel( end_ok = col_base + c_vecw <= c_model_dim _if_full = scf.IfOp(end_ok, has_else=True) with _if_then(_if_full): - # ── Vector path via layout API (all dtypes) ── - # fx.copy auto-iterates when atom width < VEC_WIDTH - # (e.g. f32: BufferCopy128b handles 4, fx.copy issues 2 calls for 8) - copy_atom = fx.make_copy_atom( - fx.rocdl.BufferCopy128b(), elem_bits - ) vec_type_c = T.vec(copy_vec_width, compute_type()) vec_type_e = T.vec(copy_vec_width, elem_type()) @@ -3692,28 +3692,8 @@ def moe_reduction_kernel( vector.broadcast(vec_type_c, fx.Float32(0.0).ir_value()) for _ in range(n_sub) ] - reg_ty = fx.MemRefType.get( - elem_type(), - fx.LayoutType.get(copy_vec_width, 1), - fx.AddressSpace.Register, - ) - reg_lay = fx.make_layout(copy_vec_width, 1) - - tok_i32 = fx.Int32(token_idx) - tile_i32 = fx.Int32(tile_idx) - tid_i32 = fx.Int32(tid) for k in range_constexpr(topk): - # X[token, k, :] → tile → thread's VEC_WIDTH slice - x_row = X_buf[tok_i32, fx.Int32(k), None] - x_tiled = fx.logical_divide( - x_row, fx.make_layout(tile_cols, 1) - ) - x_div = fx.logical_divide( - x_tiled[None, tile_i32], fx.make_layout(VEC_WIDTH, 1) - ) - x_thread = x_div[None, tid_i32] - if const_expr(use_mask): m_idx_i32 = fx.Int32(token_idx * c_topk + fx.Index(k)) mv = buffer_ops.buffer_load( @@ -3721,19 +3701,18 @@ def moe_reduction_kernel( ) mv_ok = mv != fx.Int8(0) - if const_expr(n_sub > 1): - x_inner = fx.logical_divide( - x_thread, fx.make_layout(copy_vec_width, 1) - ) for si in range_constexpr(n_sub): - src = ( - x_inner[None, fx.Int32(si)] - if n_sub > 1 - else x_thread + x_idx_i32 = fx.Int32( + (token_idx * c_topk + fx.Index(k)) * c_model_dim + + col_base + + fx.Index(si * copy_vec_width) + ) + vec_e = buffer_ops.buffer_load( + x_rsrc, + x_idx_i32, + vec_width=copy_vec_width, + dtype=elem_type(), ) - r = fx.memref_alloca(reg_ty, reg_lay) - fx.copy_atom_call(copy_atom, src, r) - vec_e = fx.memref_load_vec(r) if const_expr(use_mask): zero_e = vector.broadcast( @@ -3749,39 +3728,17 @@ def moe_reduction_kernel( acc_vecs[si] = acc_vecs[si] + vec_c # ── Store results ── - if const_expr(n_sub > 1): - y_row = Y_buf[tok_i32, None] - y_tiled = fx.logical_divide( - y_row, fx.make_layout(tile_cols, 1) - ) - y_div = fx.logical_divide( - y_tiled[None, tile_i32], fx.make_layout(VEC_WIDTH, 1) - ) - y_inner = fx.logical_divide( - y_div[None, tid_i32], fx.make_layout(copy_vec_width, 1) - ) - for si in range_constexpr(n_sub): out_vec = acc_vecs[si] if const_expr(elem_bits < 32): out_vec = out_vec.truncf(vec_type_e) - if const_expr(n_sub > 1): - dst = y_inner[None, fx.Int32(si)] - else: - y_row = Y_buf[tok_i32, None] - y_tiled = fx.logical_divide( - y_row, fx.make_layout(tile_cols, 1) - ) - y_div = fx.logical_divide( - y_tiled[None, tile_i32], - fx.make_layout(VEC_WIDTH, 1), - ) - dst = y_div[None, tid_i32] - - r_out = fx.memref_alloca(reg_ty, reg_lay) - fx.memref_store_vec(out_vec, r_out) - fx.copy_atom_call(copy_atom, r_out, dst) + y_idx_i32 = fx.Int32( + token_idx * c_model_dim + + col_base + + fx.Index(si * copy_vec_width) + ) + buffer_ops.buffer_store(out_vec, y_rsrc, y_idx_i32) with _if_else(_if_full): # Tail path: scalar load/store per lane. @@ -3837,9 +3794,9 @@ def moe_reduction_kernel( @flyc.jit def launch_moe_reduction( - X: fx.Tensor, - Y: fx.Tensor, - valid_mask: fx.Tensor, + X: fx.Pointer, + Y: fx.Pointer, + valid_mask: fx.Pointer, i32_m_tokens: fx.Int32, stream: fx.Stream, ): @@ -3962,7 +3919,21 @@ def __call__( valid_mask = torch.empty( (0, self._topk), device=arg_out.device, dtype=torch.uint8 ) - self._reduce_exe(X, Y, valid_mask, tokens_in, stream) + + def _ptr_arg(t): + type_name = type(t).__name__ + module_name = type(t).__module__ + if type_name == "FakeTensor" or "fake_tensor" in module_name: + return flyc.from_c_void_p(fx.Uint8, 0) + return flyc.from_c_void_p(fx.Uint8, t.data_ptr()) + + self._reduce_exe( + _ptr_arg(X), + _ptr_arg(Y), + _ptr_arg(valid_mask), + tokens_in, + stream, + ) @property def mode(self) -> str: diff --git a/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py b/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py index 0073641e6f..349ef40c37 100644 --- a/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py +++ b/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py @@ -278,17 +278,17 @@ def _build_kernel( @flyc.kernel(name=_kname) def kernel( - q_in: fx.Tensor, # [T, H, D] bf16, contig (H, D) - kv_in: fx.Tensor, # [T, D] bf16, may be strided - q_weight: fx.Tensor, # [D] bf16 (dummy when not q_weighted) - kv_weight: fx.Tensor, # [D] bf16 - cos_cache: fx.Tensor, # [max_pos, RD/2] bf16 - sin_cache: fx.Tensor, # [max_pos, RD/2] bf16 - positions: fx.Tensor, # [T] i64 - q_out: fx.Tensor, # [T, H, D] bf16 or fp8 - kv_out: fx.Tensor, # [T, D] bf16 or fp8 - q_scale: fx.Tensor, # [T, H, NG] f32 or uint8 (e8m0) - kv_scale: fx.Tensor, # [T, NG] f32 or uint8 (e8m0) + q_in: fx.Pointer, # [T, H, D] bf16, contig (H, D) + kv_in: fx.Pointer, # [T, D] bf16, may be strided + q_weight: fx.Pointer, # [D] bf16 (dummy when not q_weighted) + kv_weight: fx.Pointer, # [D] bf16 + cos_cache: fx.Pointer, # [max_pos, RD/2] bf16 + sin_cache: fx.Pointer, # [max_pos, RD/2] bf16 + positions: fx.Pointer, # [T] i64 + q_out: fx.Pointer, # [T, H, D] bf16 or fp8 + kv_out: fx.Pointer, # [T, D] bf16 or fp8 + q_scale: fx.Pointer, # [T, H, NG] f32 or uint8 (e8m0) + kv_scale: fx.Pointer, # [T, NG] f32 or uint8 (e8m0) kv_in_row_stride: Int32, # KV row stride in bf16 elements ): f32 = T.f32 @@ -310,19 +310,26 @@ def load_vec( bid_x = fx.block_idx.x # 0..H-1 (Q head) or H (KV) bid_t = fx.block_idx.y # token id (chunked at MAX_GRID_Y per launch) tid = fx.thread_idx.x + bid_t_idx = arith.index_cast(T.index, _to_raw(bid_t)) + + def _ptr_buffer_resource(ptr, num_records_bytes=None): + addr = fx.ptrtoint(ptr) + addr_i64 = arith.index_cast(T.i64, addr) + if num_records_bytes is None: + return buffer_ops.create_buffer_resource_from_addr(addr_i64) + return buffer_ops.create_buffer_resource_from_addr( + addr_i64, num_records_bytes=num_records_bytes + ) # --- shared: load position (i64 -> i32) --- - pos_rsrc = buffer_ops.create_buffer_resource(positions, max_size=True) + pos_nbytes = fx.Int64((bid_t_idx + fx.Index(1)) * fx.Index(8)) + pos_rsrc = _ptr_buffer_resource(positions, pos_nbytes) pos_val_i64 = buffer_ops.buffer_load(pos_rsrc, bid_t, vec_width=1, dtype=T.i64) pos_i32 = arith.trunci(i32, pos_val_i64) # --- shared: cos/sin buffer tensors (used by rope-threads only) --- - cos_buf = fx.rocdl.make_buffer_tensor(cos_cache) - sin_buf = fx.rocdl.make_buffer_tensor(sin_cache) - cos_row = fx.slice(cos_buf, (pos_i32, None)) - sin_row = fx.slice(sin_buf, (pos_i32, None)) - cos_div = fx.logical_divide(cos_row, rope_lay) - sin_div = fx.logical_divide(sin_row, rope_lay) + cos_g = GTensor(cos_cache, dtype=T.bf16, shape=(-1, RD // 2)) + sin_g = GTensor(sin_cache, dtype=T.bf16, shape=(-1, RD // 2)) def wave_reduce_add(x): w = _to_raw(x) diff --git a/pyproject.toml b/pyproject.toml index 21d0ea1f36..46da109ffd 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,7 @@ requires = [ "psutil", "ninja", "pandas", - "flydsl==0.1.8" + "flydsl==v0.1.9.dev599" ] [tool.setuptools_scm] diff --git a/requirements.txt b/requirements.txt index 40e2257b49..700c1302f2 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,4 +6,4 @@ pyyaml einops pybind11>=3.0.1 ninja -flydsl==0.1.8 +flydsl==v0.1.9.dev599 diff --git a/setup.py b/setup.py index c98f7d75a9..302e896637 100644 --- a/setup.py +++ b/setup.py @@ -13,7 +13,7 @@ OPT_COMPILER_CONFIG = os.path.join(this_dir, "aiter", "jit", "optCompilerConfig.json") PACKAGE_NAME = "amd-aiter" -FLYDSL_VERSION = "flydsl==0.1.8" +FLYDSL_VERSION = "flydsl==v0.1.9.dev599" BUILD_TARGET = os.environ.get("BUILD_TARGET", "auto") PREBUILD_KERNELS = int(os.environ.get("PREBUILD_KERNELS", 0)) From 86a2b86a4b9a1e61c0964d13d941265006de2657 Mon Sep 17 00:00:00 2001 From: root Date: Wed, 27 May 2026 10:18:22 +0000 Subject: [PATCH 4/7] change rope --- .../ops/flydsl/kernels/qk_norm_rope_quant.py | 70 +++++++++++-------- 1 file changed, 40 insertions(+), 30 deletions(-) diff --git a/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py b/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py index 349ef40c37..e3f76eec62 100644 --- a/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py +++ b/aiter/ops/flydsl/kernels/qk_norm_rope_quant.py @@ -280,10 +280,10 @@ def _build_kernel( def kernel( q_in: fx.Pointer, # [T, H, D] bf16, contig (H, D) kv_in: fx.Pointer, # [T, D] bf16, may be strided - q_weight: fx.Pointer, # [D] bf16 (dummy when not q_weighted) - kv_weight: fx.Pointer, # [D] bf16 - cos_cache: fx.Pointer, # [max_pos, RD/2] bf16 - sin_cache: fx.Pointer, # [max_pos, RD/2] bf16 + q_weight: fx.Tensor, # [D] bf16 (dummy when not q_weighted) + kv_weight: fx.Tensor, # [D] bf16 + cos_cache: fx.Tensor, # [max_pos, RD/2] bf16 + sin_cache: fx.Tensor, # [max_pos, RD/2] bf16 positions: fx.Pointer, # [T] i64 q_out: fx.Pointer, # [T, H, D] bf16 or fp8 kv_out: fx.Pointer, # [T, D] bf16 or fp8 @@ -322,14 +322,17 @@ def _ptr_buffer_resource(ptr, num_records_bytes=None): ) # --- shared: load position (i64 -> i32) --- - pos_nbytes = fx.Int64((bid_t_idx + fx.Index(1)) * fx.Index(8)) - pos_rsrc = _ptr_buffer_resource(positions, pos_nbytes) + pos_rsrc = _ptr_buffer_resource(positions) pos_val_i64 = buffer_ops.buffer_load(pos_rsrc, bid_t, vec_width=1, dtype=T.i64) pos_i32 = arith.trunci(i32, pos_val_i64) # --- shared: cos/sin buffer tensors (used by rope-threads only) --- - cos_g = GTensor(cos_cache, dtype=T.bf16, shape=(-1, RD // 2)) - sin_g = GTensor(sin_cache, dtype=T.bf16, shape=(-1, RD // 2)) + cos_buf = fx.rocdl.make_buffer_tensor(cos_cache) + sin_buf = fx.rocdl.make_buffer_tensor(sin_cache) + cos_row = fx.slice(cos_buf, (pos_i32, None)) + sin_row = fx.slice(sin_buf, (pos_i32, None)) + cos_div = fx.logical_divide(cos_row, rope_lay) + sin_div = fx.logical_divide(sin_row, rope_lay) def wave_reduce_add(x): w = _to_raw(x) @@ -525,7 +528,6 @@ def emit_body( # the input is index-typed. Doing the math in index avoids large # H*D configs (e.g. H=128 D=512 → 128 KB/token, max offset 8.6 GiB # at bid_t=65534) silently producing garbage if we feed i64. - bid_t_idx = arith.index_cast(T.index, _to_raw(bid_t)) q_tok_off_bytes = arith.MulIOp( bid_t_idx, arith.constant(H * D * 2, type=T.index) ).result @@ -579,7 +581,7 @@ def emit_body( qo_rsrc = qo_g_tmp.rsrc # row_base_bytes is now token-relative (head_idx * D bytes for fp8). row_base_bytes = ArithValue(head_idx) * arith.constant(D, type=i32) - qs_rsrc = buffer_ops.create_buffer_resource(q_scale, max_size=True) + qs_rsrc = _ptr_buffer_resource(q_scale) # q_scale layout (T, H, NG) flat: bid_t * H*NG + head_idx * NG. # Per-lane adds group_idx inside emit_body. scale_base_off_q = ArithValue(bid_t) * arith.constant( @@ -622,7 +624,7 @@ def emit_body( # buffer_ops with the explicit kv_in_row_stride argument, then # round-trip through an rmem tensor to get a Fly-wrapped vec that # the rest of emit_body (.to/.reduce/[i]) expects. - kv_rsrc = buffer_ops.create_buffer_resource(kv_in, max_size=True) + kv_rsrc = _ptr_buffer_resource(kv_in) kv_off_elems = ArithValue(bid_t) * ArithValue( kv_in_row_stride ) + ArithValue(tid) * arith.constant(VEC, type=i32) @@ -655,7 +657,7 @@ def emit_body( ) kvo_rsrc = kvo_g_tmp.rsrc row_base_bytes = arith.constant(0, type=i32) # already at token base - kvs_rsrc = buffer_ops.create_buffer_resource(kv_scale, max_size=True) + kvs_rsrc = _ptr_buffer_resource(kv_scale) # kv_scale layout (T, NG) flat: bid_t * NG. Per-lane adds # group_idx inside emit_body. scale_base_off_kv = ArithValue(bid_t) * arith.constant(NG, type=i32) @@ -697,17 +699,17 @@ def emit_body( # @flyc.jit function in the codebase. @flyc.jit def launch_qk_norm_rope_quant( - q_in: fx.Tensor, - kv_in: fx.Tensor, + q_in: fx.Pointer, + kv_in: fx.Pointer, q_weight: fx.Tensor, kv_weight: fx.Tensor, cos_cache: fx.Tensor, sin_cache: fx.Tensor, - positions: fx.Tensor, - q_out: fx.Tensor, - kv_out: fx.Tensor, - q_scale: fx.Tensor, - kv_scale: fx.Tensor, + positions: fx.Pointer, + q_out: fx.Pointer, + kv_out: fx.Pointer, + q_scale: fx.Pointer, + kv_scale: fx.Pointer, kv_in_row_stride: fx.Int32, num_tokens: fx.Int32, stream: fx.Stream = fx.Stream(None), @@ -970,6 +972,14 @@ def flydsl_qk_norm_rope_quant( stream = torch.cuda.current_stream() fx_stream = Stream(stream) + def _ptr_arg(t): + return flyc.from_c_void_p(fx.Uint8, t.data_ptr()) + + q_weight_static = flyc.from_dlpack(q_weight_arg) + kv_weight_static = flyc.from_dlpack(kv_weight) + cos_static = flyc.from_dlpack(cos_2d) + sin_static = flyc.from_dlpack(sin_2d) + # HW grid Y is a 16-bit field on AMD HIP → cap 65535 blocks/launch. The # kernel uses per-token GTensor base-shift so each chunk's resource span # is small (just the chunk's tokens), but the grid Y dim itself is HW- @@ -985,17 +995,17 @@ def flydsl_qk_norm_rope_quant( n = min(MAX_GRID_Y, T_tok - start) end = start + n launcher( - q_view[start:end], - kv[start:end], - q_weight_arg, - kv_weight, - cos_2d, - sin_2d, - positions[start:end], - q_out[start:end], - kv_out[start:end], - q_scale_arg[start:end] if quant else q_scale_arg, - kv_scale_arg[start:end] if quant else kv_scale_arg, + _ptr_arg(q_view[start:end]), + _ptr_arg(kv[start:end]), + q_weight_static, + kv_weight_static, + cos_static, + sin_static, + _ptr_arg(positions[start:end]), + _ptr_arg(q_out[start:end]), + _ptr_arg(kv_out[start:end]), + _ptr_arg(q_scale_arg[start:end] if quant else q_scale_arg), + _ptr_arg(kv_scale_arg[start:end] if quant else kv_scale_arg), kv.stride(0), n, stream=fx_stream, From 0fd5fb525fadf38aa26f6533f527bf0e3841e675 Mon Sep 17 00:00:00 2001 From: coderfeli Date: Wed, 27 May 2026 10:41:44 +0000 Subject: [PATCH 5/7] style: black format --- aiter/ops/flydsl/gemm_kernels.py | 1 + .../flydsl/kernels/flash_attn_func_gfx1201.py | 6 +++++- .../flydsl/kernels/mixed_moe_gemm_2stage.py | 20 ++++++++++--------- aiter/ops/flydsl/kernels/moe_gemm_2stage.py | 12 ++++++++--- aiter/ops/flydsl/kernels/preshuffle_gemm.py | 6 +----- 5 files changed, 27 insertions(+), 18 deletions(-) diff --git a/aiter/ops/flydsl/gemm_kernels.py b/aiter/ops/flydsl/gemm_kernels.py index 9c61ab8c77..0923a543fe 100644 --- a/aiter/ops/flydsl/gemm_kernels.py +++ b/aiter/ops/flydsl/gemm_kernels.py @@ -76,6 +76,7 @@ def _ptr_view_safe(t: torch.Tensor): return flyc.from_c_void_p(fx.Uint8, 0) return flyc.from_c_void_p(fx.Uint8, t.data_ptr()) + # Keep the generic auto-generated catalog aligned with the upstream FlyDSL # reference tuning space. The wider local one-off search space introduced # gfx950-faulting candidates (for example tile_k=160 and tile_n=160/192), diff --git a/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py b/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py index adab48ad79..85cf28eff6 100644 --- a/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py +++ b/aiter/ops/flydsl/kernels/flash_attn_func_gfx1201.py @@ -760,7 +760,11 @@ def _ptr_arg(t): if hasattr(t, "data_ptr"): type_name = type(t).__name__ module_name = type(t).__module__ - ptr = 0 if type_name == "FakeTensor" or "fake_tensor" in module_name else t.data_ptr() + ptr = ( + 0 + if type_name == "FakeTensor" or "fake_tensor" in module_name + else t.data_ptr() + ) return flyc.from_c_void_p(fx.Uint8, ptr) return t diff --git a/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py b/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py index 124a2b25ce..a17ab82ca6 100644 --- a/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py +++ b/aiter/ops/flydsl/kernels/mixed_moe_gemm_2stage.py @@ -621,7 +621,9 @@ def _ptr_buffer_resource(ptr, num_records_bytes): w_rsrc = _ptr_buffer_resource(arg_w, w_nbytes) # Out: [tokens*topk, inter_dim] - numids_rsrc = _ptr_buffer_resource(arg_num_valid_ids, arith.constant(4, type=T.i32)) + numids_rsrc = _ptr_buffer_resource( + arg_num_valid_ids, arith.constant(4, type=T.i32) + ) num_valid_i32 = buffer_ops.buffer_load( numids_rsrc, arith.constant(0, index=True), vec_width=1, dtype=T.i32 ) @@ -655,9 +657,7 @@ def _ptr_buffer_resource(ptr, num_records_bytes): eid_nbytes_i32 = arith.index_cast(T.i32, eid_nbytes_idx) expert_rsrc = _ptr_buffer_resource(arg_expert_ids, eid_nbytes_i32) bias_rsrc = ( - _ptr_buffer_resource(arg_bias, bias_nbytes) - if enable_bias - else None + _ptr_buffer_resource(arg_bias, bias_nbytes) if enable_bias else None ) # Sorted-scale buffer resource for fused mxfp4 quantization @@ -679,7 +679,9 @@ def _ptr_buffer_resource(ptr, num_records_bytes): _sort_scale_nbytes = arith.index_cast( T.i32, _sort_padded_rows * _sort_padded_cols ) - sorted_scale_rsrc = _ptr_buffer_resource(arg_out_scale_sorted, _sort_scale_nbytes) + sorted_scale_rsrc = _ptr_buffer_resource( + arg_out_scale_sorted, _sort_scale_nbytes + ) # ---- persist_m loop (same pattern as stage2) ---- _PERSIST_M = persist_m @@ -3186,7 +3188,9 @@ def check_c_k_valid_gate(base_k): out_rsrc = _ptr_buffer_resource(arg_out, out_nbytes_i32) # num_valid_ids (sorted padded MN) for scale sizing / guards. - numids_rsrc = _ptr_buffer_resource(arg_num_valid_ids, arith.constant(4, type=T.i32)) + numids_rsrc = _ptr_buffer_resource( + arg_num_valid_ids, arith.constant(4, type=T.i32) + ) num_valid_i32 = buffer_ops.buffer_load( numids_rsrc, arith.constant(0, index=True), vec_width=1, dtype=T.i32 ) @@ -3244,9 +3248,7 @@ def check_c_k_valid_gate(base_k): eid_nbytes_i32 = arith.index_cast(T.i32, eid_nbytes_idx) expert_rsrc = _ptr_buffer_resource(arg_expert_ids, eid_nbytes_i32) bias_rsrc = ( - _ptr_buffer_resource(arg_bias, bias_nbytes) - if enable_bias - else None + _ptr_buffer_resource(arg_bias, bias_nbytes) if enable_bias else None ) # ---- persist loop ---- diff --git a/aiter/ops/flydsl/kernels/moe_gemm_2stage.py b/aiter/ops/flydsl/kernels/moe_gemm_2stage.py index 17784acbb2..13ad672317 100644 --- a/aiter/ops/flydsl/kernels/moe_gemm_2stage.py +++ b/aiter/ops/flydsl/kernels/moe_gemm_2stage.py @@ -533,11 +533,17 @@ def silu(x): sw_rsrc = _ptr_buffer_resource(arg_scale_w, sw_nbytes) sorted_nbytes_idx = size_expert_ids_in * fx.Index(tile_m) * fx.Index(4) - sorted_rsrc = _ptr_buffer_resource(arg_sorted_token_ids, sorted_nbytes_idx) - sorted_w_rsrc = _ptr_buffer_resource(arg_sorted_weights, sorted_nbytes_idx) + sorted_rsrc = _ptr_buffer_resource( + arg_sorted_token_ids, sorted_nbytes_idx + ) + sorted_w_rsrc = _ptr_buffer_resource( + arg_sorted_weights, sorted_nbytes_idx + ) # expert ids: [blocks] i32 -> bytes = size_expert_ids_in*4 - expert_rsrc = _ptr_buffer_resource(arg_expert_ids, size_expert_ids_in * fx.Index(4)) + expert_rsrc = _ptr_buffer_resource( + arg_expert_ids, size_expert_ids_in * fx.Index(4) + ) # Expert id for this M tile (keep address math in `index`) expert_i32 = buffer_ops.buffer_load( diff --git a/aiter/ops/flydsl/kernels/preshuffle_gemm.py b/aiter/ops/flydsl/kernels/preshuffle_gemm.py index e1de7714ef..60c468b78b 100644 --- a/aiter/ops/flydsl/kernels/preshuffle_gemm.py +++ b/aiter/ops/flydsl/kernels/preshuffle_gemm.py @@ -483,11 +483,7 @@ def _ptr_buffer_resource(ptr, num_records_bytes=None): if const_expr(_has_bias): bias_rsrc = _ptr_buffer_resource(arg_bias) b_rsrc = _ptr_buffer_resource(arg_b) - scale_b_rsrc = ( - None - if (is_f16_or_bf16) - else _ptr_buffer_resource(arg_scale_b) - ) + scale_b_rsrc = None if (is_f16_or_bf16) else _ptr_buffer_resource(arg_scale_b) bx_m = bx * tile_m by_n = by * tile_n From 90bdbef48c79fe0a8310985efa45849f67ed229b Mon Sep 17 00:00:00 2001 From: coderfeli Date: Wed, 27 May 2026 14:30:52 +0000 Subject: [PATCH 6/7] Normalize FlyDSL dependency version for CI builds Co-authored-by: Cursor --- .github/requirements/triton-test.txt | 2 +- pyproject.toml | 2 +- requirements.txt | 2 +- setup.py | 5 +++-- 4 files changed, 6 insertions(+), 5 deletions(-) diff --git a/.github/requirements/triton-test.txt b/.github/requirements/triton-test.txt index 15883a02f7..51ffb7c05b 100644 --- a/.github/requirements/triton-test.txt +++ b/.github/requirements/triton-test.txt @@ -8,7 +8,7 @@ pybind11==3.0.1 ninja==1.11.1.4 psutil packaging -flydsl==v0.1.9.dev599 +flydsl==0.1.9.dev599 # Test deps. pandas==2.2.3 diff --git a/pyproject.toml b/pyproject.toml index 46da109ffd..67a11c1c5c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -8,7 +8,7 @@ requires = [ "psutil", "ninja", "pandas", - "flydsl==v0.1.9.dev599" + "flydsl==0.1.9.dev599" ] [tool.setuptools_scm] diff --git a/requirements.txt b/requirements.txt index 700c1302f2..173753473e 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,4 +6,4 @@ pyyaml einops pybind11>=3.0.1 ninja -flydsl==v0.1.9.dev599 +flydsl==0.1.9.dev599 diff --git a/setup.py b/setup.py index 302e896637..d3bf647356 100644 --- a/setup.py +++ b/setup.py @@ -13,7 +13,7 @@ OPT_COMPILER_CONFIG = os.path.join(this_dir, "aiter", "jit", "optCompilerConfig.json") PACKAGE_NAME = "amd-aiter" -FLYDSL_VERSION = "flydsl==v0.1.9.dev599" +FLYDSL_VERSION = "flydsl==0.1.9.dev599" BUILD_TARGET = os.environ.get("BUILD_TARGET", "auto") PREBUILD_KERNELS = int(os.environ.get("PREBUILD_KERNELS", 0)) @@ -57,8 +57,9 @@ def is_develop_mode(): if not IS_WINDOWS and is_develop_mode(): try: from importlib.metadata import version as pkg_version + from packaging.version import Version - if pkg_version("flydsl") != FLYDSL_VERSION.split("==")[1]: + if Version(pkg_version("flydsl")) != Version(FLYDSL_VERSION.split("==")[1]): raise ImportError("version mismatch") except Exception: subprocess.check_call( From eb0ec849c048bf195c2826679d15eebcb8268022 Mon Sep 17 00:00:00 2001 From: coderfeli Date: Wed, 27 May 2026 14:48:34 +0000 Subject: [PATCH 7/7] Fix ATOM image checkout for target commits Co-authored-by: Cursor --- .github/workflows/atom-test.yaml | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/.github/workflows/atom-test.yaml b/.github/workflows/atom-test.yaml index 32ee1ca5da..1701e64ed5 100644 --- a/.github/workflows/atom-test.yaml +++ b/.github/workflows/atom-test.yaml @@ -184,9 +184,9 @@ jobs: RUN pip install --upgrade "pybind11>=3.0.1" RUN pip show pybind11 RUN rm -rf /app/aiter-test - RUN git clone ${{ env.GITHUB_REPO_URL }} /app/aiter-test && \\ + RUN git clone --no-checkout ${{ env.GITHUB_REPO_URL }} /app/aiter-test && \\ cd /app/aiter-test && \\ - git checkout ${{ env.GITHUB_COMMIT_SHA }} && \\ + git checkout --force ${{ env.GITHUB_COMMIT_SHA }} && \\ git submodule sync && git submodule update --init --recursive && \\ MAX_JOBS=64 PREBUILD_KERNELS=0 GPU_ARCHS=gfx950 pip install -e . && \\ ./.github/scripts/install_triton.sh