diff --git a/aiter/ops/triton/_gluon_kernels/gfx1250/attention/fp8_mqa_logits.py b/aiter/ops/triton/_gluon_kernels/gfx1250/attention/fp8_mqa_logits.py index a9bfae5c4c..91ff2802e6 100644 --- a/aiter/ops/triton/_gluon_kernels/gfx1250/attention/fp8_mqa_logits.py +++ b/aiter/ops/triton/_gluon_kernels/gfx1250/attention/fp8_mqa_logits.py @@ -8,6 +8,7 @@ from aiter.ops.triton._gluon_kernels.gfx950.attention.fp8_mqa_logits import ( _weighted_sum_fma_fold, ) +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr from aiter.ops.triton.utils.common_utils import strip_annotate _MAX_PROPAGATE_NAN_ALL = gl.constexpr(PropagateNan.ALL) @@ -567,7 +568,23 @@ def mqa_logits_loop_pipelined( ) -@gluon.jit +_gluon_fp8_mqa_logits_kernel_repr = make_kernel_repr( + "_gluon_fp8_mqa_logits_kernel", + [ + "NUM_HEADS", + "HEAD_SIZE", + "BLOCK_KV", + "NUM_WARPS", + "NUM_BUFFERS", + "NUM_CHAINS", + "LOOP_VARIANT", + "USE_BUFFER_LOAD", + "USE_BUFFER_STORE", + ], +) + + +@gluon.jit(repr=_gluon_fp8_mqa_logits_kernel_repr) def _gluon_fp8_mqa_logits_kernel( Q_ptr, # fp8e4m3 [seq_len, NUM_HEADS, HEAD_SIZE] KV_ptr, # fp8e4m3 [seq_len_kv, HEAD_SIZE] diff --git a/aiter/ops/triton/_gluon_kernels/gfx950/attention/fp8_mqa_logits.py b/aiter/ops/triton/_gluon_kernels/gfx950/attention/fp8_mqa_logits.py index 2210dcfec4..573d0f3a34 100644 --- a/aiter/ops/triton/_gluon_kernels/gfx950/attention/fp8_mqa_logits.py +++ b/aiter/ops/triton/_gluon_kernels/gfx950/attention/fp8_mqa_logits.py @@ -4,6 +4,7 @@ from triton.language.core import PropagateNan from triton.language.core import _aggregate as aggregate +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr from aiter.ops.triton.utils.common_utils import strip_annotate _MAX_PROPAGATE_NAN_ALL = gl.constexpr(PropagateNan.ALL) @@ -754,7 +755,25 @@ def mqa_logits_loop_double_buf( ) -@gluon.jit +_gluon_fp8_mqa_logits_kernel_repr = make_kernel_repr( + "_gluon_fp8_mqa_logits_kernel", + [ + "NUM_HEADS", + "HEAD_SIZE", + "BLOCK_KV", + "NUM_WARPS", + "NUM_BUFFERS", + "NUM_CHAINS", + "USE_BUFFER_LOAD", + "USE_BUFFER_STORE", + "USE_PADDED_SHARED_LAYOUT", + "BLOCK_M", + "MFMA_NONK_DIM", + ], +) + + +@gluon.jit(repr=_gluon_fp8_mqa_logits_kernel_repr) def _gluon_fp8_mqa_logits_kernel( Q_ptr, # fp8e4m3 [seq_len, NUM_HEADS, HEAD_SIZE] KV_ptr, # fp8e4m3 [seq_len_kv, HEAD_SIZE] diff --git a/aiter/ops/triton/_triton_kernels/attention/block_lut.py b/aiter/ops/triton/_triton_kernels/attention/block_lut.py index b57159f5bb..acc7dba5d2 100644 --- a/aiter/ops/triton/_triton_kernels/attention/block_lut.py +++ b/aiter/ops/triton/_triton_kernels/attention/block_lut.py @@ -10,8 +10,17 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr -@triton.jit() +_block_attn_mask_to_lut_kernel_repr = make_kernel_repr( + "_block_attn_mask_to_lut_kernel", + [ + "BLOCK_KB", + ], +) + + +@triton.jit(repr=_block_attn_mask_to_lut_kernel_repr) def _block_attn_mask_to_lut_kernel( mask_ptr, lut_start_ptr, diff --git a/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention.py b/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention.py index df1fbd3e0e..59b78cbed7 100644 --- a/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention.py +++ b/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention.py @@ -4,6 +4,7 @@ from aiter.ops.triton._triton_kernels.flash_attn_triton_amd.common import ( compute_alibi_block, ) +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr def map_dims(shape, indices): @@ -1219,7 +1220,36 @@ def compute_block_masking( ) -@triton.jit +_sage_fwd_repr = make_kernel_repr( + "sage_fwd", + [ + "RETURN_LSE", + "HQ", + "HK", + "ACTUAL_BLOCK_DMODEL_QK", + "ACTUAL_BLOCK_DMODEL_V", + "IS_VARLEN", + "IS_CAUSAL", + "USE_SLIDING_WINDOW", + "WINDOW_SIZE_LEFT", + "WINDOW_SIZE_RIGHT", + "BLOCK_M", + "BLOCK_DMODEL_QK", + "BLOCK_DMODEL_V", + "BLOCK_N", + "PRE_LOAD_V", + "USE_BIAS", + "ENABLE_DROPOUT", + "RETURN_SCORES", + "USE_ALIBI", + "USE_EXP2", + "USE_SEQUSED", + "USE_BLOCK_SPARSE", + ], +) + + +@triton.jit(repr=_sage_fwd_repr) def sage_fwd( Q, K, diff --git a/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention_mxfp4.py b/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention_mxfp4.py index 63f8193cde..4f2a38fc61 100644 --- a/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention_mxfp4.py +++ b/aiter/ops/triton/_triton_kernels/attention/fav3_sage_attention_mxfp4.py @@ -1,6 +1,8 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr + @triton.jit def compute_padding_info(seqlen_k, BLOCK_N: tl.constexpr): @@ -564,7 +566,30 @@ def _sage_fwd_blocksparse_mask_mxfp4( return acc, l_i, m_i -@triton.jit +_sage_fwd_mxfp4_repr = make_kernel_repr( + "sage_fwd_mxfp4", + [ + "Q_DTYPE_STR", + "K_DTYPE_STR", + "HQ", + "HK", + "ACTUAL_BLOCK_DMODEL_QK", + "ACTUAL_BLOCK_DMODEL_V", + "IS_VARLEN", + "IS_CAUSAL", + "BLOCK_M", + "BLOCK_DMODEL_QK", + "BLOCK_DMODEL_V", + "BLOCK_N", + "PRE_LOAD_V", + "USE_BIAS", + "USE_BLOCK_SPARSE", + "RETURN_LSE", + ], +) + + +@triton.jit(repr=_sage_fwd_mxfp4_repr) def sage_fwd_mxfp4( Q, K, diff --git a/aiter/ops/triton/_triton_kernels/attention/fp8_mqa_logits.py b/aiter/ops/triton/_triton_kernels/attention/fp8_mqa_logits.py index de32438fb9..a820314a8b 100644 --- a/aiter/ops/triton/_triton_kernels/attention/fp8_mqa_logits.py +++ b/aiter/ops/triton/_triton_kernels/attention/fp8_mqa_logits.py @@ -1,8 +1,19 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr -@triton.jit +_fp8_mqa_logits_kernel_repr = make_kernel_repr( + "_fp8_mqa_logits_kernel", + [ + "NUM_HEADS", + "HEAD_SIZE", + "BLOCK_KV", + ], +) + + +@triton.jit(repr=_fp8_mqa_logits_kernel_repr) def _fp8_mqa_logits_kernel( Q_ptr, # fp8e4m3 [seq_len, H, D] KV_ptr, # fp8e4m3 [seq_len_kv, D] diff --git a/aiter/ops/triton/_triton_kernels/attention/lean_atten_paged.py b/aiter/ops/triton/_triton_kernels/attention/lean_atten_paged.py index 94398cc232..48bdb7e8b4 100644 --- a/aiter/ops/triton/_triton_kernels/attention/lean_atten_paged.py +++ b/aiter/ops/triton/_triton_kernels/attention/lean_atten_paged.py @@ -17,6 +17,8 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr + @triton.jit def find_group(x): @@ -29,7 +31,23 @@ def find_group(x): return group_id, group_size, total_blocks -@triton.jit +_la_persistent_paged_repr = make_kernel_repr( + "la_persistent_paged", + [ + "HEAD_DIM", + "BLOCK_M", + "BLOCK_N", + "batch_size", + "num_m_blocks", + "high_load_wgs", + "max_tiles_per_wg", + "tiles_per_head", + "num_splits", + ], +) + + +@triton.jit(repr=_la_persistent_paged_repr) def la_persistent_paged( Q, K, diff --git a/aiter/ops/triton/_triton_kernels/attention/pa_mqa_logits.py b/aiter/ops/triton/_triton_kernels/attention/pa_mqa_logits.py index 86228e3478..d51f2e32d6 100644 --- a/aiter/ops/triton/_triton_kernels/attention/pa_mqa_logits.py +++ b/aiter/ops/triton/_triton_kernels/attention/pa_mqa_logits.py @@ -4,13 +4,26 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr + @triton.jit def _sum_combine(a, b): return a + b -@triton.jit +_deepgemm_fp8_paged_mqa_logits_stage1_ragged_k_repr = make_kernel_repr( + "_deepgemm_fp8_paged_mqa_logits_stage1_ragged_k", + [ + "ChunkQ", + "ChunkK", + "HiddenDim", + "SplitKV", + ], +) + + +@triton.jit(repr=_deepgemm_fp8_paged_mqa_logits_stage1_ragged_k_repr) def _deepgemm_fp8_paged_mqa_logits_stage1_ragged_k( batch_size, next_n, @@ -106,7 +119,18 @@ def _deepgemm_fp8_paged_mqa_logits_stage1_ragged_k( ) -@triton.jit +_deepgemm_fp8_paged_mqa_logits_ragged_k_repr = make_kernel_repr( + "_deepgemm_fp8_paged_mqa_logits_ragged_k", + [ + "ChunkQ", + "ChunkK", + "HiddenDim", + "SplitKV", + ], +) + + +@triton.jit(repr=_deepgemm_fp8_paged_mqa_logits_ragged_k_repr) def _deepgemm_fp8_paged_mqa_logits_ragged_k( batch_size, next_n, @@ -201,7 +225,19 @@ def _deepgemm_fp8_paged_mqa_logits_ragged_k( ) -@triton.jit +_deepgemm_fp8_paged_mqa_logits_stage1_repr = make_kernel_repr( + "_deepgemm_fp8_paged_mqa_logits_stage1", + [ + "ChunkQ", + "ChunkK", + "HiddenDim", + "KVBlockSize", + "SplitKV", + ], +) + + +@triton.jit(repr=_deepgemm_fp8_paged_mqa_logits_stage1_repr) def _deepgemm_fp8_paged_mqa_logits_stage1( batch_size, next_n, @@ -311,7 +347,17 @@ def _deepgemm_fp8_paged_mqa_logits_stage1( ) -@triton.jit +_deepgemm_fp8_paged_mqa_logits_varctx_schedule_repr = make_kernel_repr( + "_deepgemm_fp8_paged_mqa_logits_varctx_schedule", + [ + "ChunkK", + "AlignedBatchSize", + "TryCount", + ], +) + + +@triton.jit(repr=_deepgemm_fp8_paged_mqa_logits_varctx_schedule_repr) def _deepgemm_fp8_paged_mqa_logits_varctx_schedule( batch_size, context_len_ptr, @@ -357,7 +403,19 @@ def _deepgemm_fp8_paged_mqa_logits_varctx_schedule( tl.store(safe_chunks_per_cta_ptr, safe_seg_lens) -@triton.jit +_deepgemm_fp8_paged_mqa_logits_repr = make_kernel_repr( + "_deepgemm_fp8_paged_mqa_logits", + [ + "ChunkQ", + "ChunkK", + "HiddenDim", + "KVBlockSize", + "SplitKV", + ], +) + + +@triton.jit(repr=_deepgemm_fp8_paged_mqa_logits_repr) def _deepgemm_fp8_paged_mqa_logits( batch_size, next_n, diff --git a/aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py b/aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py index bd86c0f3d8..99b679f276 100644 --- a/aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py +++ b/aiter/ops/triton/_triton_kernels/attention/sparse_attention_dsv4.py @@ -2,6 +2,8 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr + # ===================================================================== # Utility # ===================================================================== @@ -261,12 +263,23 @@ def _get_prefill_autotune_configs(): ] +_sparse_attn_prefill_kernel_repr = make_kernel_repr( + "_sparse_attn_prefill_kernel", + [ + "HAS_ATTN_SINK", + "BLOCK_H", + "BLOCK_D", + "BLOCK_K", + ], +) + + @triton.autotune( configs=_get_prefill_autotune_configs(), key=["num_heads", "head_dim", "HAS_ATTN_SINK"], prune_configs_by={"early_config_prune": _prefill_prune_configs}, ) -@triton.jit +@triton.jit(repr=_sparse_attn_prefill_kernel_repr) def _sparse_attn_prefill_kernel( q_ptr, # [num_queries, num_heads, head_dim] kv_ptr, # [num_kv, head_dim] diff --git a/aiter/ops/triton/_triton_kernels/attention/unified_attention.py b/aiter/ops/triton/_triton_kernels/attention/unified_attention.py index 3cd63005ce..318b9a37fb 100644 --- a/aiter/ops/triton/_triton_kernels/attention/unified_attention.py +++ b/aiter/ops/triton/_triton_kernels/attention/unified_attention.py @@ -52,7 +52,30 @@ def find_seq_idx( return left - 1 -@triton.jit +_kernel_unified_attention_2d_repr = make_kernel_repr( + "kernel_unified_attention_2d", + [ + "num_query_heads", + "num_queries_per_kv", + "BLOCK_SIZE", + "TILE_SIZE", + "HEAD_SIZE", + "HEAD_SIZE_PADDED", + "USE_ALIBI_SLOPES", + "USE_QQ_BIAS", + "USE_SOFTCAP", + "USE_SINKS", + "SLIDING_WINDOW", + "BLOCK_Q", + "BLOCK_M", + "ALL_DECODE", + "SHUFFLED_KV_CACHE", + "K_WIDTH", + ], +) + + +@triton.jit(repr=_kernel_unified_attention_2d_repr) def kernel_unified_attention_2d( output_ptr, # [num_tokens, num_query_heads, head_size] query_ptr, # [num_tokens, num_query_heads, head_size] diff --git a/aiter/ops/triton/_triton_kernels/attention/unified_attention_sparse_mla.py b/aiter/ops/triton/_triton_kernels/attention/unified_attention_sparse_mla.py index c979b28a49..f43208ce3e 100644 --- a/aiter/ops/triton/_triton_kernels/attention/unified_attention_sparse_mla.py +++ b/aiter/ops/triton/_triton_kernels/attention/unified_attention_sparse_mla.py @@ -1,6 +1,8 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr + @triton.jit def cdiv_fn(x, y): @@ -38,7 +40,23 @@ def find_seq_idx( return left - 1 -@triton.jit +_kernel_unified_attention_sparse_mla_2d_repr = make_kernel_repr( + "_kernel_unified_attention_sparse_mla_2d", + [ + "num_query_heads", + "num_queries_per_kv", + "BLOCK_SIZE", + "topk_count", + "BLOCK_M", + "ROPE_RANK", + "KV_LORA_RANK", + "TILE_SIZE", + "ALL_DECODE", + ], +) + + +@triton.jit(repr=_kernel_unified_attention_sparse_mla_2d_repr) def _kernel_unified_attention_sparse_mla_2d( output_ptr, # [num_tokens, num_query_heads, KV_LORA_RANK] query_ptr, # [num_tokens, num_query_heads, KV_LORA_RANK] diff --git a/aiter/ops/triton/attention/mla_decode.py b/aiter/ops/triton/attention/mla_decode.py index f08bf453dc..3359a79f4d 100644 --- a/aiter/ops/triton/attention/mla_decode.py +++ b/aiter/ops/triton/attention/mla_decode.py @@ -30,6 +30,8 @@ import triton import triton.language as tl +from aiter.ops.triton.utils._triton.kernel_repr import make_kernel_repr + is_hip_ = hasattr(torch.version, "hip") and torch.version.hip is not None @@ -38,7 +40,23 @@ def tanh(x): return 2 * tl.sigmoid(2 * x) - 1 -@triton.jit +_fwd_kernel_stage1_repr = make_kernel_repr( + "_fwd_kernel_stage1", + [ + "kv_group_num", + "BLOCK_DMODEL", + "BLOCK_DPE", + "BLOCK_DV", + "BLOCK_N", + "NUM_KV_SPLITS", + "PAGE_SIZE", + "Lk", + "Lv", + ], +) + + +@triton.jit(repr=_fwd_kernel_stage1_repr) def _fwd_kernel_stage1( Q, K_Buffer, @@ -277,7 +295,25 @@ def _decode_att_m_fwd( ) -@triton.jit +_fwd_grouped_kernel_stage1_repr = make_kernel_repr( + "_fwd_grouped_kernel_stage1", + [ + "kv_group_num", + "q_head_num", + "BLOCK_DMODEL", + "BLOCK_DPE", + "BLOCK_DV", + "BLOCK_N", + "BLOCK_H", + "NUM_KV_SPLITS", + "PAGE_SIZE", + "Lk", + "Lv", + ], +) + + +@triton.jit(repr=_fwd_grouped_kernel_stage1_repr) def _fwd_grouped_kernel_stage1( Q, K_Buffer, @@ -543,7 +579,17 @@ def _decode_grouped_att_m_fwd( ) -@triton.jit +_fwd_kernel_stage2_repr = make_kernel_repr( + "_fwd_kernel_stage2", + [ + "NUM_KV_SPLITS", + "BLOCK_DV", + "Lv", + ], +) + + +@triton.jit(repr=_fwd_kernel_stage2_repr) def _fwd_kernel_stage2( Mid_O, o, @@ -646,7 +692,15 @@ def _decode_softmax_reducev_fwd( ) -@triton.jit +_csr_to_dense_kernel_repr = make_kernel_repr( + "_csr_to_dense_kernel", + [ + "BLOCK_N", + ], +) + + +@triton.jit(repr=_csr_to_dense_kernel_repr) def _csr_to_dense_kernel( kv_indices, kv_indptr,