From 46fb01349093497f739973b94d018d512571c6fc Mon Sep 17 00:00:00 2001 From: lcy-seso Date: Sun, 27 Sep 2026 22:04:45 +0800 Subject: [PATCH 1/5] [Feat][Manifest] Add spec-only entries for the 200-operator release plan Adds 24 spec-only entries: nine quantization ops, six sampling ops in a new sampling family, seven attention and KV-cache ops, a fused Q/K RMSNorm plus RoPE op, and Kimi Delta Attention. Each entry states its reference contract, workload rows from its issue and a roofline; the counts that depend on metadata values are functions in tileops.perf.formulas. The DeepSeek sparse attention rows need per-request top-k positions padded with -1, which no existing generator yields, so attn.sparse_topk_positions is added with its test. A spec-only entry may have no class yet, so the public-surface test requires every implemented entry to be reachable rather than every entry. --- src/tileops/__init__.py | 2 + src/tileops/manifest/primitives.py | 23 ++ src/tileops/manifest/spec/attention.yaml | 240 ++++++++++++++++++ .../manifest/spec/linear_attention.yaml | 69 +++++ src/tileops/manifest/spec/norm.yaml | 34 +++ src/tileops/manifest/spec/quantization.yaml | 203 ++++++++++++++- src/tileops/manifest/spec/sampling.yaml | 167 ++++++++++++ src/tileops/perf/formulas.py | 188 ++++++++++++++ src/tileops/sampling.py | 3 + tests/test_manifest_generators.py | 16 ++ tests/test_public_api.py | 9 +- 11 files changed, 948 insertions(+), 6 deletions(-) create mode 100644 src/tileops/manifest/spec/sampling.yaml create mode 100644 src/tileops/sampling.py diff --git a/src/tileops/__init__.py b/src/tileops/__init__.py index dcecc266d..182c1a350 100644 --- a/src/tileops/__init__.py +++ b/src/tileops/__init__.py @@ -42,6 +42,7 @@ quantization, reduction, rope, + sampling, sequence_modeling, ) from .ops.op_base import Op @@ -57,6 +58,7 @@ "convolution", "fft", "moe", + "sampling", "rope", "attention", "linear_attention", diff --git a/src/tileops/manifest/primitives.py b/src/tileops/manifest/primitives.py index cf491fc47..4db576b90 100644 --- a/src/tileops/manifest/primitives.py +++ b/src/tileops/manifest/primitives.py @@ -449,6 +449,24 @@ def causal_topk_indices(rng, batch, seq, heads_kv, k, extent, start, stride): return rows +def sparse_topk_positions(rng, lengths, queries, k): + """Per request of length `n` and query `s`, up to `k` distinct positions in `[0, n - queries + s]`, padded with -1.""" + if queries <= 0 or k <= 0 or any(n < queries for n in lengths): + raise ValueError( + f"attn.sparse_topk_positions needs S_q > 0, K > 0 and every length >= S_q, " + f"got lengths={lengths}, S_q={queries}, K={k}" + ) + rows = [] + for n in lengths: + per_query = [] + for s in range(queries): + visible = n - queries + s + 1 + picked = rng.sample(range(visible), min(k, visible)) + per_query.append(picked + [-1] * (k - len(picked))) + rows.append(per_query) + return rows + + def key_windows(lengths, first, count, side): """Per query position, the first key of its sequence (`start`) or one past itself (`end`).""" if not lengths or any(n <= 0 for n in lengths): @@ -523,6 +541,7 @@ def moe_layout_metadata(layout, rows, experts): "sample_indices": sample_indices, "moe.layout_metadata": moe_layout_metadata, "causal_topk_indices": causal_topk_indices, + "attn.sparse_topk_positions": sparse_topk_positions, "key_windows": key_windows, "full": full, } @@ -543,6 +562,7 @@ def moe_layout_metadata(layout, rows, experts): "sample_indices": (("Int", "Int"), "Value"), "moe.layout_metadata": (("ADT", "Int", "Int"), "Value"), "causal_topk_indices": (("Int", "Int", "Int", "Int", "Int", "Int", "Int"), "Value"), + "attn.sparse_topk_positions": (("Seq[Int]", "Int", "Int"), "Value"), "key_windows": (("Seq[Int]", "Int", "Int", "'start' | 'end'"), "Value"), "full": (("Seq[Int]", "Int"), "Value"), } @@ -563,6 +583,7 @@ def moe_layout_metadata(layout, rows, experts): "sample_indices": 1, "moe.layout_metadata": 1, "causal_topk_indices": 4, + "attn.sparse_topk_positions": 3, "key_windows": 1, } # The shape of each generator's result, from its arguments. @@ -589,6 +610,7 @@ def moe_layout_metadata(layout, rows, experts): heads, k, ), + "attn.sparse_topk_positions": lambda L, queries, k: (len(L), queries, k), "key_windows": lambda L, first, count, side: (count,), "full": lambda shape, value: tuple(shape), } @@ -600,6 +622,7 @@ def moe_layout_metadata(layout, rows, experts): "topk_ids", "sample_indices", "causal_topk_indices", + "attn.sparse_topk_positions", } ) diff --git a/src/tileops/manifest/spec/attention.yaml b/src/tileops/manifest/spec/attention.yaml index c62b43c02..78a59e2b3 100644 --- a/src/tileops/manifest/spec/attention.yaml +++ b/src/tileops/manifest/spec/attention.yaml @@ -493,3 +493,243 @@ NSAVarlenFwdOp: roofline: # The blocks the selection keeps decide the scores and the key/value rows read. func: "tileops.perf.formulas.nsa_fwd_varlen_roofline" + +# --------------------------------------------------------------------------- +# attention — paged KV cache and latent-cache operators +# --------------------------------------------------------------------------- +# A paged cache is [NP, PS, ...]: NP pages of PS token rows. Slot s addresses row +# s % PS of page s // PS. The Multi-Head Latent Attention (MLA) latent cache is one +# [NP, PS, R] tensor with no head axis: every query head reads the same latent row. + +MultiHeadLatentAttentionPagedFwdOp: + # Absorbed-form MLA decode over a paged latent cache: key = the whole cache row + # (kv_c ‖ k_pe), value = its first kv_lora_rank columns. The S_q query tokens are the + # last S_q positions of each request's cache_seqlens. sm_scale defaults to DK ** -0.5. + family: attention + status: spec-only + signature: + forall: {B: Dim, S_q: Dim, H: Dim, DK: Dim, NP: Dim, PS: Dim, W: Dim, T: "DType[float16 | bfloat16]", KV: "DType[float16 | bfloat16 | float8_e4m3fn]", cache_lens: "Seq[Int]"} + params: + kv_lora_rank: {type: int} + is_causal: {type: bool, default: true} + sm_scale: {type: "float | None", default: null} + inputs: + q: {dtype: T, shape: "[B, S_q, H, DK]"} + kv_cache: {dtype: KV, shape: "[NP, PS, DK]", contiguous: true} + block_table: {dtype: int32, shape: "[B, W]", values: "paged_block_table(B, W, NP)", requires: ["in_range(0, NP)"]} + cache_seqlens: {dtype: int32, shape: "[B]", values: "as_tensor(cache_lens)", requires: ["in_range(S_q, W * PS + 1)"]} + # Per-tensor dequantization scale of an FP8 cache. + kv_scale: {dtype: float32, shape: "[1]", optional: true} + outputs: + o: {dtype: T, shape: "[B, S_q, H, kv_lora_rank]"} + lse: {dtype: float32, shape: "[B, S_q, H]"} + dtype_combos: + - {T: float16, KV: float16} + - {T: bfloat16, KV: bfloat16} + - {T: float16, KV: float8_e4m3fn} + - {T: bfloat16, KV: float8_e4m3fn} + shape_rules: + - "S_q > 0 and PS > 0" + - "0 < kv_lora_rank < DK" + - "(KV == 'float8_e4m3fn') == present(kv_scale)" + workloads: + - {S_q: 1, H: 128, DK: 576, NP: 4096, PS: 64, W: 64, kv_lora_rank: 512, cache_lens: "repeat(4096, 64)", dtype_cases: [{T: bfloat16, KV: bfloat16}, {T: float16, KV: float16}], label: ds-v3-4k} + - {S_q: 1, H: 128, DK: 576, NP: 32768, PS: 64, W: 512, kv_lora_rank: 512, cache_lens: "repeat(32768, 64)", dtype_cases: [{T: bfloat16, KV: bfloat16}], label: ds-v3-32k} + - {S_q: 2, H: 128, DK: 576, NP: 4096, PS: 64, W: 128, kv_lora_rank: 512, cache_lens: "repeat(8192, 32)", dtype_cases: [{T: bfloat16, KV: bfloat16}], label: ds-v3-mtp} + - {S_q: 1, H: 64, DK: 576, NP: 32768, PS: 64, W: 256, kv_lora_rank: 512, cache_lens: "repeat(16384, 128)", some: [kv_scale], dtype_cases: [{T: bfloat16, KV: float8_e4m3fn}], label: kimi-k2-fp8} + roofline: + # The scores each query row sees, and the cache rows the block table reaches. + func: "tileops.perf.formulas.mla_paged_fwd_roofline" + +MultiHeadLatentAttentionVarlenFwdOp: + # MLA prefill over packed requests after the latent is decompressed: the key of head h + # is k_nope[:, h] ‖ k_pe, with k_pe shared by every head; queries and keys are the same + # tokens. The request length bound follows from cu_seqlens. sm_scale defaults to + # (DN + PE) ** -0.5. + family: attention + status: spec-only + signature: + forall: {B: Dim, T_q: Dim, H: Dim, DN: Dim, PE: Dim, DV: Dim, T: "DType[float16 | bfloat16]", seq_lens: "Seq[Int]"} + params: + is_causal: {type: bool, default: true} + sm_scale: {type: "float | None", default: null} + inputs: + q: {dtype: T, shape: "[T_q, H, DN + PE]"} + k_nope: {dtype: T, shape: "[T_q, H, DN]"} + k_pe: {dtype: T, shape: "[T_q, PE]"} + v: {dtype: T, shape: "[T_q, H, DV]"} + cu_seqlens: {dtype: int32, shape: "[B + 1]", values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(T_q)"]} + outputs: + o: {dtype: T, shape: "[T_q, H, DV]"} + lse: {dtype: float32, shape: "[T_q, H]"} + shape_rules: + - "DN + PE > 0" + workloads: + - {T_q: 32768, H: 128, DN: 128, PE: 64, DV: 128, seq_lens: "repeat(4096, 8)", dtype_cases: [{T: bfloat16}, {T: float16}], label: ds-v3-8x4k} + - {T_q: 32256, H: 128, DN: 128, PE: 64, DV: 128, seq_lens: [512, 1024, 2048, 4096, 8192, 16384], dtype_cases: [{T: bfloat16}], label: ds-v3-mixed} + - {T_q: 32768, H: 64, DN: 128, PE: 64, DV: 128, seq_lens: "repeat(8192, 4)", dtype_cases: [{T: bfloat16}], label: kimi-k2-4x8k} + roofline: + # The scores each request sees under its mask; every tensor moves once. + func: "tileops.perf.formulas.mla_varlen_fwd_roofline" + +PagedKVCacheWriteFwdOp: + # Scatters token n's k and v rows into slot slot_mapping[n] of k_pages and v_pages; -1 + # skips the token, and distinct non-negative slots are the caller's obligation. An FP8 + # cache stores (k / k_scale) and (v / v_scale) with a saturating cast. + family: attention + status: spec-only + signature: + forall: {N: Dim, H_kv: Dim, D: Dim, NP: Dim, PS: Dim, T: "DType[float16 | bfloat16]", KV: "DType[float16 | bfloat16 | float8_e4m3fn]"} + inputs: + k: {dtype: T, shape: "[N, H_kv, D]"} + v: {dtype: T, shape: "[N, H_kv, D]"} + k_pages: {dtype: KV, shape: "[NP, PS, H_kv, D]", mutated: true, contiguous: true} + v_pages: {dtype: KV, shape: "[NP, PS, H_kv, D]", mutated: true, contiguous: true} + slot_mapping: {dtype: int64, shape: "[N]", values: "sample_indices(N, NP * PS)", requires: ["in_range(-1, NP * PS)"]} + k_scale: {dtype: float32, shape: "[1]", optional: true} + v_scale: {dtype: float32, shape: "[1]", optional: "present(k_scale)"} + outputs: {} + dtype_combos: + - {T: float16, KV: float16} + - {T: bfloat16, KV: bfloat16} + - {T: float16, KV: float8_e4m3fn} + - {T: bfloat16, KV: float8_e4m3fn} + shape_rules: + - "PS > 0" + - "(KV == 'float8_e4m3fn') == present(k_scale)" + workloads: + - {N: 64, H_kv: 8, D: 128, NP: 4096, PS: 16, dtype_cases: [{T: bfloat16, KV: bfloat16}, {T: float16, KV: float16}], label: llama-8b-decode} + - {N: 8192, H_kv: 8, D: 128, NP: 4096, PS: 16, dtype_cases: [{T: bfloat16, KV: bfloat16}], label: llama-8b-prefill} + - {N: 4096, H_kv: 4, D: 128, NP: 4096, PS: 16, some: [k_scale], dtype_cases: [{T: bfloat16, KV: float8_e4m3fn}], label: qwen3-235b-fp8} + - {N: 256, H_kv: 8, D: 64, NP: 4096, PS: 16, dtype_cases: [{T: bfloat16, KV: bfloat16}], label: gpt-oss-120b} + roofline: + # Only the tokens with a slot are read and written. + func: "tileops.perf.formulas.paged_kv_cache_write_roofline" + +MultiHeadLatentAttentionKVCacheWriteFwdOp: + # Writes kv_c ‖ k_pe of token n into slot slot_mapping[n] of the latent cache; -1 skips + # the token, and distinct non-negative slots are the caller's obligation. With + # fuse_rope, k_pe is rotated at positions[n] by cos_sin_cache (cos ‖ sin halves) as it + # is written; k_pe itself is not modified. An FP8 cache stores the row divided by scale + # with a saturating cast. + family: attention + status: spec-only + signature: + forall: {N: Dim, DC: Dim, PE: Dim, NP: Dim, PS: Dim, P: Dim, T: "DType[float16 | bfloat16]", KV: "DType[float16 | bfloat16 | float8_e4m3fn]", C: "DType[float16 | bfloat16 | float32]", seq_lens: "Seq[Int]"} + params: + fuse_rope: {type: bool, default: false} + rope_layout: {type: "'neox' | 'interleaved'", default: neox} + inputs: + kv_c: {dtype: T, shape: "[N, DC]"} + k_pe: {dtype: T, shape: "[N, PE]"} + kv_cache: {dtype: KV, shape: "[NP, PS, DC + PE]", mutated: true, contiguous: true} + slot_mapping: {dtype: int64, shape: "[N]", values: "sample_indices(N, NP * PS)", requires: ["in_range(-1, NP * PS)"]} + scale: {dtype: float32, shape: "[1]", optional: true} + positions: {dtype: int64, shape: "[N]", optional: fuse_rope, values: "packed_positions(seq_lens)", requires: ["in_range(0, P)"]} + cos_sin_cache: {dtype: C, shape: "[P, PE]", optional: fuse_rope} + outputs: {} + dtype_combos: + - {T: float16, KV: float16} + - {T: bfloat16, KV: bfloat16} + - {T: float16, KV: float8_e4m3fn} + - {T: bfloat16, KV: float8_e4m3fn} + shape_rules: + - "PS > 0" + - "not fuse_rope or PE % 2 == 0" + - "(KV == 'float8_e4m3fn') == present(scale)" + workloads: + - {N: 64, DC: 512, PE: 64, NP: 2048, PS: 64, dtype_cases: [{T: bfloat16, KV: bfloat16}], label: ds-v3-decode} + - {N: 16384, DC: 512, PE: 64, NP: 2048, PS: 64, dtype_cases: [{T: bfloat16, KV: bfloat16}, {T: float16, KV: float16}], label: ds-v3-prefill} + - {DC: 512, PE: 64, NP: 2048, PS: 64, P: 163840, seq_lens: "repeat(1024, 4)", fuse_rope: true, rope_layout: interleaved, some: [scale], dtype_cases: [{T: bfloat16, KV: float8_e4m3fn, C: float32}], label: ds-v3-fp8-rope} + - {DC: 512, PE: 64, NP: 2048, PS: 64, P: 131072, seq_lens: [128], fuse_rope: true, rope_layout: interleaved, dtype_cases: [{T: bfloat16, KV: bfloat16, C: bfloat16}], label: kimi-k2-rope} + roofline: + # Only the tokens with a slot are read and written; the rotation reads the cos/sin + # rows their positions name. + func: "tileops.perf.formulas.mla_kv_cache_write_roofline" + +MergeAttentionStatesFwdOp: + # Merges two partial attention results over disjoint key sets by their log-sum-exp: + # v = (v_a * e^s_a + v_b * e^s_b) / (e^s_a + e^s_b), s = log(e^s_a + e^s_b), computed + # against max(s_a, s_b) (FlashInfer merge_state). A -inf state contributes nothing; two + # merge to s = -inf and v = 0. + family: attention + status: spec-only + signature: + forall: {T_q: Dim, H: Dim, D: Dim, T: "DType[float16 | bfloat16]"} + inputs: + v_a: {dtype: T, shape: "[T_q, H, D]"} + s_a: {dtype: float32, shape: "[T_q, H]"} + v_b: {dtype: T, shape: "[T_q, H, D]"} + s_b: {dtype: float32, shape: "[T_q, H]"} + outputs: + v: {dtype: T, shape: "[T_q, H, D]"} + s: {dtype: float32, shape: "[T_q, H]"} + workloads: + - {T_q: 8192, H: 128, D: 128, dtype_cases: [{T: bfloat16}, {T: float16}], label: ds-v3-chunked-prefill} + - {T_q: 4096, H: 64, D: 128, dtype_cases: [{T: bfloat16}], label: llama-70b-cascade} + - {T_q: 256, H: 64, D: 128, dtype_cases: [{T: bfloat16}], label: qwen3-235b-dcp} + roofline: + # Per row: the max, two subtracts, two exps, the sum, log and add, and the two weights' + # divides; per element: two multiplies and an add. + flops: "T_q * H * (3 * D + 10)" + +DeepSeekSparseAttentionPagedFwdOp: + # DeepSeek Sparse Attention: MLA decode over the key positions indices names, read from + # a paged cache in the FlashMLA DeepSeek-V3.2 FP8 row format. A 656-byte row holds 512 + # float8_e4m3fn latent values, 4 float32 scales (one per 128 of them) and 64 bfloat16 + # rope values; the value is the dequantized latent. -1 pads indices; every other entry + # below cache_seqlens[b] is the caller's obligation. sm_scale defaults to 576 ** -0.5. + family: attention + status: spec-only + signature: + forall: {B: Dim, S_q: Dim, H: Dim, K: Dim, NP: Dim, PS: Dim, W: Dim, cache_lens: "Seq[Int]"} + params: + sm_scale: {type: "float | None", default: null} + inputs: + q: {dtype: bfloat16, shape: "[B, S_q, H, 576]"} + kv_cache: {dtype: uint8, shape: "[NP, PS, 656]", contiguous: true} + block_table: {dtype: int32, shape: "[B, W]", values: "paged_block_table(B, W, NP)", requires: ["in_range(0, NP)"]} + cache_seqlens: {dtype: int32, shape: "[B]", values: "as_tensor(cache_lens)", requires: ["in_range(S_q, W * PS + 1)"]} + indices: {dtype: int32, shape: "[B, S_q, K]", values: "attn.sparse_topk_positions(cache_lens, S_q, K)", requires: ["in_range(-1, W * PS)"]} + outputs: + o: {dtype: bfloat16, shape: "[B, S_q, H, 512]"} + lse: {dtype: float32, shape: "[B, S_q, H]"} + shape_rules: + - "S_q > 0 and PS > 0" + workloads: + - {S_q: 1, H: 128, K: 2048, NP: 32768, PS: 64, W: 512, cache_lens: "repeat(32768, 64)", label: ds-v32-32k} + - {S_q: 2, H: 128, K: 2048, NP: 32768, PS: 64, W: 2048, cache_lens: "repeat(131072, 16)", label: ds-v32-mtp-128k} + roofline: + # The selected scores, and the cache rows they resolve to through the block table. + func: "tileops.perf.formulas.dsa_paged_fwd_roofline" + +PagedKVCacheGatherFwdOp: + # Gathers request b's cache positions [start_b, start_b + len_b) into rows + # cu_seq_lens[b] .. cu_seq_lens[b + 1] of dst, start_b = seq_starts[b] or 0; an FP8 + # cache is dequantized by scale into out_dtype. dst is passed because its length is + # cu_seq_lens[-1], a value no shape carries. + family: attention + status: spec-only + signature: + forall: {B: Dim, T_q: Dim, NP: Dim, PS: Dim, W: Dim, E: Shape, KV: "DType[float16 | bfloat16 | float8_e4m3fn]", seq_lens: "Seq[Int]", starts: "Seq[Int]"} + params: + out_dtype: {type: "float16 | bfloat16 | None", default: null} + inputs: + dst: {dtype: "coalesce_dtype(out_dtype, KV)", shape: "[T_q, *E]", mutated: true, write_only: true} + cache: {dtype: KV, shape: "[NP, PS, *E]", contiguous: true} + block_table: {dtype: int32, shape: "[B, W]", values: "paged_block_table(B, W, NP)", requires: ["in_range(0, NP)"]} + cu_seq_lens: {dtype: int32, shape: "[B + 1]", values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(T_q)", "max_segment(W * PS)"]} + seq_starts: {dtype: int32, shape: "[B]", optional: true, values: "as_tensor(starts)", requires: ["in_range(0, W * PS)", "attn.paged_fits(cu_seq_lens, W * PS)"]} + scale: {dtype: float32, shape: "[1]", optional: true} + outputs: {} + shape_rules: + - "PS > 0" + - "(KV == 'float8_e4m3fn') == present(scale)" + - "present(out_dtype) if KV == 'float8_e4m3fn' else (not present(out_dtype) or out_dtype.value == KV)" + workloads: + - {T_q: 131072, NP: 2048, PS: 64, W: 256, E: [576], seq_lens: "repeat(16384, 8)", out_dtype: bfloat16, some: [scale], dtype_cases: [{KV: float8_e4m3fn}], label: ds-v3-chunk-fp8} + - {T_q: 129024, NP: 6144, PS: 64, W: 1024, E: [576], seq_lens: [2048, 4096, 8192, 16384, 32768, 65536], dtype_cases: [{KV: bfloat16}], label: ds-v3-mixed} + - {T_q: 32768, NP: 4096, PS: 16, W: 512, E: [8, 128], seq_lens: "repeat(4096, 8)", starts: "repeat(2048, 8)", some: [seq_starts], dtype_cases: [{KV: bfloat16}], label: llama-70b-cp} + roofline: + # The cache rows each request's range reaches through the block table. + func: "tileops.perf.formulas.paged_kv_cache_gather_roofline" diff --git a/src/tileops/manifest/spec/linear_attention.yaml b/src/tileops/manifest/spec/linear_attention.yaml index b0ed018f4..aaf0594bd 100644 --- a/src/tileops/manifest/spec/linear_attention.yaml +++ b/src/tileops/manifest/spec/linear_attention.yaml @@ -99,6 +99,75 @@ GatedDeltaNetFwdOp: roofline: flops: "B * S * HV * 7 * K * V" +KimiDeltaAttentionFwdOp: + # Kimi Delta Attention (KDA): the gated delta rule with a per-key-channel decay g + # (log space), one contract for equal-length, packed-varlen and single-token calls, with + # GatedDeltaNetFwdOp's state layout; final_state is always returned. With + # use_gate_in_kernel, g is the raw gate input and the decay is + # -exp(A_log) * softplus(g + dt_bias), or lower_bound * sigmoid(exp(A_log) * (g + dt_bias)) + # when lower_bound is passed (FLA fused_recurrent_kda). + family: linear_attention + status: spec-only + signature: + types: + State: + params: {packed: Bool, vf: Bool, B: Dim, N: Dim, HV: Dim, K: Dim, V: Dim} + match: [packed, vf] + cases: + - {when: [false, false], is: "[B, HV, K, V]"} + - {when: [false, true], is: "[B, HV, V, K]"} + - {when: [true, false], is: "[N, HV, K, V]"} + - {when: [true, true], is: "[N, HV, V, K]"} + forall: {B: Dim, S: Dim, H: Dim, HV: Dim, K: Dim, V: Dim, N: Dim, T: "DType[float16 | bfloat16]", seq_lens: "Seq[Int]"} + params: + scale: {type: "float | None", default: null} + use_qk_l2norm_in_kernel: {type: bool, default: false} + use_beta_sigmoid_in_kernel: {type: bool, default: false} + allow_neg_eigval: {type: bool, default: false} + state_v_first: {type: bool, default: false} + use_gate_in_kernel: {type: bool, default: false} + lower_bound: {type: "float | None", default: null} + inputs: + q: {dtype: T, shape: "[B, S, H, K]"} + k: {dtype: T, shape: "[B, S, H, K]"} + v: {dtype: T, shape: "[B, S, HV, V]"} + g: {dtype: T, shape: "[B, S, HV, K]"} + beta: {dtype: T, shape: "[B, S, HV]"} + initial_state: {dtype: float32, shape: "State[present(cu_seqlens), state_v_first, B, N, HV, K, V]", optional: true} + cu_seqlens: {dtype: int64, shape: "[N + 1]", optional: true, values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(S)"]} + cu_seqlens_cpu: {dtype: int64, shape: "[N + 1]", optional: true, device: cpu, values: "prefix_sum(seq_lens)", requires: ["prefix_offsets(S)"]} + A_log: {dtype: float32, shape: "[HV]", optional: use_gate_in_kernel} + dt_bias: {dtype: float32, shape: "[HV * K]", optional: true} + outputs: + o: {dtype: T, shape: "[B, S, HV, V]"} + final_state: {dtype: float32, shape: "State[present(cu_seqlens), state_v_first, B, N, HV, K, V]"} + shape_rules: + - "H > 0 and HV % H == 0" + - "K > 0 and V > 0" + - "not present(cu_seqlens_cpu) or present(cu_seqlens)" + - "not present(cu_seqlens) or B == 1" + - "not allow_neg_eigval or use_beta_sigmoid_in_kernel" + - "not present(dt_bias) or use_gate_in_kernel" + - "not present(lower_bound) or use_gate_in_kernel" + workloads: + - {B: 1, S: 4096, H: 32, HV: 32, K: 128, V: 128, seq_lens: "repeat(1024, 4)", use_qk_l2norm_in_kernel: true, some: [cu_seqlens], dtype_cases: [{T: bfloat16}, {T: float16}], label: kimi-linear-4k} + - {B: 1, S: 32768, H: 32, HV: 32, K: 128, V: 128, seq_lens: "repeat(8192, 4)", use_qk_l2norm_in_kernel: true, some: [cu_seqlens], dtype_cases: [{T: bfloat16}], label: kimi-linear-32k} + - {B: 1, S: 1, H: 32, HV: 32, K: 128, V: 128, use_qk_l2norm_in_kernel: true, some: [initial_state], dtype_cases: [{T: bfloat16}], label: kimi-linear-decode-b1} + - {B: 64, S: 1, H: 32, HV: 32, K: 128, V: 128, use_qk_l2norm_in_kernel: true, some: [initial_state], dtype_cases: [{T: bfloat16}], label: kimi-linear-decode-b64} + roofline: + # The per-token recurrence (the signature carries no chunk size). Per token and value + # head: exp(g), the decay, k^T S, v - k^T S and beta, the rank-1 update and the q^T S + # readout; per query head the q scale, shared by its value heads. Flag terms: the q/k + # l2 norms; the in-kernel gate per channel (softplus form 3, sigmoid form 6, the bias + # add 1) and its per-head factor (-exp(A_log) 2, exp(A_log) 1); the beta sigmoid and its + # doubling. + flops: >- + B * S * (HV * (7 * K * V + K + 2 * V) + H * K + + (H * (6 * K + 4) if use_qk_l2norm_in_kernel else 0) + + (HV * K * ((6 if present(lower_bound) else 3) + (1 if present(dt_bias) else 0)) if use_gate_in_kernel else 0) + + (HV * (4 + (1 if allow_neg_eigval else 0)) if use_beta_sigmoid_in_kernel else 0)) + + ((1 if present(lower_bound) else 2) * HV if use_gate_in_kernel else 0) + DeltaNetDecodeFwdOp: # Single-token ungated DeltaNet decode. State ownership is functional: state # is read and new_state is returned. diff --git a/src/tileops/manifest/spec/norm.yaml b/src/tileops/manifest/spec/norm.yaml index f1d19f422..518d1d97f 100644 --- a/src/tileops/manifest/spec/norm.yaml +++ b/src/tileops/manifest/spec/norm.yaml @@ -276,3 +276,37 @@ InstanceNormFwdOp: # Input statistics: mean (1), centered variance (3), normalize the centered value (1); # running statistics: normalize (2). One more per present affine tensor. flops: "((5 if use_input_stats else 2) + (1 if present(weight) else 0) + (1 if present(bias) else 0)) * B * C * prod(L)" + +FusedQKNormRopeFwdOp: + # In place on a packed projection qkv = [q heads | k heads | v heads] of width D each: + # RMSNorm over each q head (q_weight) and each k head (k_weight), then RoPE at + # positions[n] on the first R columns of each q and k head, from cos_sin_cache (cos ‖ sin + # halves). The v columns are untouched. + family: norm + status: spec-only + signature: + forall: {N: Dim, D: Dim, P: Dim, R: Dim, T: "DType[float16 | bfloat16]", C: "DType[float16 | bfloat16 | float32]", seq_lens: "Seq[Int]"} + params: + num_heads: {type: int} + num_kv_heads: {type: int} + eps: {type: float, default: 1.0e-6} + rope_layout: {type: "'neox' | 'interleaved'", default: neox} + inputs: + qkv: {dtype: T, shape: "[N, (num_heads + 2 * num_kv_heads) * D]", mutated: true} + q_weight: {dtype: T, shape: "[D]"} + k_weight: {dtype: T, shape: "[D]"} + cos_sin_cache: {dtype: C, shape: "[P, R]"} + positions: {dtype: int64, shape: "[N]", values: "packed_positions(seq_lens)", requires: ["in_range(0, P)"]} + outputs: {} + shape_rules: + - "num_heads > 0 and num_kv_heads > 0" + - "R > 0 and R % 2 == 0 and R <= D" + workloads: + - {D: 128, P: 40960, R: 128, num_heads: 32, num_kv_heads: 8, seq_lens: [64], dtype_cases: [{T: bfloat16, C: bfloat16}, {T: float16, C: float32}], label: qwen3-8b-t64} + - {D: 128, P: 40960, R: 128, num_heads: 32, num_kv_heads: 8, seq_lens: "repeat(1024, 8)", dtype_cases: [{T: bfloat16, C: bfloat16}], label: qwen3-8b-t8192} + - {D: 128, P: 40960, R: 128, num_heads: 64, num_kv_heads: 4, seq_lens: "repeat(1024, 4)", dtype_cases: [{T: bfloat16, C: bfloat16}], label: qwen3-235b-t4096} + - {D: 128, P: 131072, R: 64, num_heads: 96, num_kv_heads: 8, seq_lens: "repeat(1024, 4)", dtype_cases: [{T: bfloat16, C: float32}], label: glm-4.5-partial-rope} + roofline: + # The q and k columns are read and written, the v columns not touched, and only the + # cos/sin rows the positions name are read. + func: "tileops.perf.formulas.fused_qk_norm_rope_roofline" diff --git a/src/tileops/manifest/spec/quantization.yaml b/src/tileops/manifest/spec/quantization.yaml index f3a00512c..da1d7d8da 100644 --- a/src/tileops/manifest/spec/quantization.yaml +++ b/src/tileops/manifest/spec/quantization.yaml @@ -1,7 +1,8 @@ -# Narrow quantization helper operators. +# Quantization and dequantization operators. # -# These entries describe existing helper surfaces only. They do not claim the -# full INT8/INT4/NF4/FP8 quantization family from the older release plan. +# INT8 is symmetric in [-127, 127]; scales are float32; rounding is half-to-even +# (torch.round); an all-zero group gets scale 1.0, so dequantization never divides +# by zero. A scale is the dequantization multiplier: x ~= q * scale. FP8QuantFwdOp: family: quantization @@ -25,3 +26,199 @@ FP8QuantFwdOp: # Per element: abs, the row-max comparison, the scaling and the clamp; per row: the amax # floor and the scale. flops: "4 * batch * seq_len_kv * kv_group * index_dim + 2 * batch * seq_len_kv * kv_group" + +INT8QuantPerTensorFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + x: {dtype: T, shape: "[M, K]"} + outputs: + q: {dtype: int8, shape: "[M, K]"} + # amax(|x|) / 127 + scale: {dtype: float32, shape: "[1]"} + shape_rules: + - "M > 0 and K > 0" + workloads: + - {M: 4096, K: 4096, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-prefill} + - {M: 32, K: 4096, dtype_cases: [{T: bfloat16}], label: llama-8b-decode} + - {M: 8192, K: 5120, dtype_cases: [{T: bfloat16}, {T: float32}], label: qwen3-32b-prefill} + roofline: + # Per element: abs, the max, the scaling, round and clamp; once: the scale and its zero guard. + flops: "5 * M * K + 2" + +INT8QuantPerChannelFwdOp: + family: quantization + status: spec-only + signature: + forall: {N: Dim, K: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + w: {dtype: T, shape: "[N, K]"} + outputs: + q: {dtype: int8, shape: "[N, K]"} + # One per output channel (row): amax(|w[n, :]|) / 127. + scale: {dtype: float32, shape: "[N]"} + shape_rules: + - "K > 0" + workloads: + - {N: 14336, K: 4096, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-mlp-up} + - {N: 8192, K: 28672, dtype_cases: [{T: bfloat16}], label: llama-70b-mlp-down} + - {N: 5120, K: 8192, dtype_cases: [{T: bfloat16}, {T: float32}], label: qwen3-32b-o-proj} + roofline: + # Per element: abs, the max, the scaling, round and clamp; per row: the scale and its zero guard. + flops: "5 * N * K + 2 * N" + +INT8QuantPerBlockFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + x: {dtype: T, shape: "[M, K]"} + outputs: + q: {dtype: int8, shape: "[M, K]"} + # One per 128 contiguous elements along K; the last block of a row may be partial. + scale: {dtype: float32, shape: "[M, ceil_div(K, 128)]"} + workloads: + - {M: 4096, K: 7168, dtype_cases: [{T: float16}, {T: bfloat16}], label: ds-v3-prefill} + - {M: 64, K: 7168, dtype_cases: [{T: bfloat16}], label: ds-v3-decode} + - {M: 4096, K: 4096, dtype_cases: [{T: bfloat16}, {T: float32}], label: qwen3-235b-prefill} + roofline: + # Per element: abs, the max, the scaling, round and clamp; per block: the scale and its zero guard. + flops: "5 * M * K + 2 * M * ceil_div(K, 128)" + +INT4QuantPerGroupFwdOp: + # Asymmetric INT4 with a scale and a zero point per group of group_size elements along + # K: group_size == K is per channel. The outputs are GemmW4A16FwdOp's weight operands. + family: quantization + status: spec-only + signature: + forall: {N: Dim, K: Dim, T: "DType[float16]"} + params: + group_size: {type: int, default: 128} + inputs: + w: {dtype: T, shape: "[N, K]"} + outputs: + # Two INT4 values per byte, in the order GemmW4A16FwdOp.repack produces. + packed_weight: {dtype: uint8, shape: "[N, K // 2]"} + # A constant group gets a nonzero scale. + weight_scale: {dtype: T, shape: "[N, K // group_size]"} + weight_zero: {dtype: uint8, shape: "[N, K // group_size]"} + shape_rules: + - "group_size > 0" + - "K % 2 == 0 and K % group_size == 0" + workloads: + - {N: 14336, K: 4096, dtype_cases: [{T: float16}], label: llama-8b-mlp-up} + - {N: 14336, K: 4096, group_size: 4096, dtype_cases: [{T: float16}], label: llama-8b-mlp-up-chan} + - {N: 8192, K: 28672, dtype_cases: [{T: float16}], label: llama-70b-mlp-down} + - {N: 5120, K: 8192, dtype_cases: [{T: float16}], label: qwen3-32b-o-proj} + roofline: + # Per element: the min, the max, the scaling, the zero-point add, round and clamp; per + # group: the range, its divide and zero guard, and the zero point (negate, divide, round, clamp). + flops: "6 * N * K + 7 * N * (K // group_size)" + +SmoothQuantFwdOp: + # q = quantize_per_row(x / smooth): per-channel smoothing, then symmetric INT8 per row. + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim, T: "DType[float16 | bfloat16]"} + inputs: + x: {dtype: T, shape: "[M, K]"} + smooth: {dtype: float32, shape: "[K]"} + outputs: + q: {dtype: int8, shape: "[M, K]"} + scale: {dtype: float32, shape: "[M]"} + shape_rules: + - "K > 0" + workloads: + - {M: 4096, K: 4096, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-prefill} + - {M: 32, K: 4096, dtype_cases: [{T: bfloat16}], label: llama-8b-decode} + - {M: 4096, K: 8192, dtype_cases: [{T: bfloat16}], label: llama-70b-prefill} + roofline: + # Per element: the smoothing divide, abs, the max, the scaling, round and clamp; per row: + # the scale and its zero guard. + flops: "6 * M * K + 2 * M" + +INT8DequantPerTensorFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim} + params: + out_dtype: {type: "float16 | bfloat16 | float32"} + inputs: + q: {dtype: int8, shape: "[M, K]"} + scale: {dtype: float32, shape: "[1]"} + outputs: + x: {dtype: out_dtype, shape: "[M, K]"} + workloads: + - {M: 4096, K: 4096, out_dtype: bfloat16, label: llama-8b-prefill} + - {M: 32, K: 4096, out_dtype: bfloat16, label: llama-8b-decode} + - {M: 8192, K: 5120, out_dtype: float16, label: qwen3-32b-prefill} + - {M: 8192, K: 5120, out_dtype: float32, label: qwen3-32b-prefill} + roofline: + # One multiply per element, in float32; the cast is not counted. + flops: "M * K" + +INT8DequantPerChannelFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim} + params: + out_dtype: {type: "float16 | bfloat16 | float32"} + inputs: + q: {dtype: int8, shape: "[M, K]"} + scale: {dtype: float32, shape: "[M]"} + outputs: + x: {dtype: out_dtype, shape: "[M, K]"} + workloads: + - {M: 14336, K: 4096, out_dtype: bfloat16, label: llama-8b-mlp-up} + - {M: 14336, K: 4096, out_dtype: float16, label: llama-8b-mlp-up} + - {M: 8192, K: 28672, out_dtype: bfloat16, label: llama-70b-mlp-down} + roofline: + flops: "M * K" + +INT8DequantPerBlockFwdOp: + family: quantization + status: spec-only + signature: + forall: {M: Dim, K: Dim} + params: + out_dtype: {type: "float16 | bfloat16 | float32"} + inputs: + q: {dtype: int8, shape: "[M, K]"} + # One per 128 contiguous elements along K. + scale: {dtype: float32, shape: "[M, ceil_div(K, 128)]"} + outputs: + x: {dtype: out_dtype, shape: "[M, K]"} + workloads: + - {M: 4096, K: 7168, out_dtype: bfloat16, label: ds-v3-prefill} + - {M: 64, K: 7168, out_dtype: bfloat16, label: ds-v3-decode} + - {M: 4096, K: 4096, out_dtype: float16, label: qwen3-235b-prefill} + roofline: + flops: "M * K" + +FP8QuantPerBlockFwdOp: + # The 128x128 block-scaled FP8 weight format of DeepSeek-V3 checkpoints; scale is the + # dequantization multiplier stored as weight_scale_inv. Edge tiles may be partial. + family: quantization + status: spec-only + signature: + forall: {N: Dim, K: Dim, T: "DType[bfloat16 | float16 | float32]"} + inputs: + w: {dtype: T, shape: "[N, K]"} + outputs: + q: {dtype: float8_e4m3fn, shape: "[N, K]"} + # amax(|tile|) / 448 per 128x128 tile. + scale: {dtype: float32, shape: "[ceil_div(N, 128), ceil_div(K, 128)]"} + workloads: + - {N: 7168, K: 18432, dtype_cases: [{T: bfloat16}], label: ds-v3-mlp-down} + - {N: 4096, K: 7168, dtype_cases: [{T: bfloat16}, {T: float32}], label: ds-v3-expert-gate-up} + - {N: 9216, K: 4096, dtype_cases: [{T: bfloat16}, {T: float16}], label: qwen3-235b-qkv} + roofline: + # Per element: abs, the max, the scaling and the saturating cast; per tile: the scale and + # its zero guard. + flops: "4 * N * K + 2 * ceil_div(N, 128) * ceil_div(K, 128)" diff --git a/src/tileops/manifest/spec/sampling.yaml b/src/tileops/manifest/spec/sampling.yaml new file mode 100644 index 000000000..6cf300d2f --- /dev/null +++ b/src/tileops/manifest/spec/sampling.yaml @@ -0,0 +1,167 @@ +# Logit filters and token draws of a decode step. +# +# Logits are [B, V]; a filter returns logits of the same shape and dtype with the removed +# entries set to -inf, so filters compose. Sampling parameters are per-row tensors. The +# random ops take their Philox state as (seed, offset) tensors, so a fixed pair gives a +# fixed result. + +TopKMaskFwdOp: + # Keeps the k[b] largest logits of row b, ties at the threshold kept; a row with + # k[b] >= V is unchanged (FlashInfer top_k_mask_logits). + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]", k_list: "Seq[Int]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + k: {dtype: int32, shape: "[B]", values: "as_tensor(k_list)", requires: ["in_range(1, 2147483648)"]} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + workloads: + - {V: 128256, k_list: "repeat(50, 1)", dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {V: 128256, k_list: "repeat(50, 64)", dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {V: 128256, k_list: "repeat(50, 256)", dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {V: 151936, k_list: "repeat(20, 128)", dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + - {V: 201088, k_list: "repeat(50, 64)", dtype_cases: [{T: bfloat16}], label: gpt-oss-120b-b64} + roofline: + # The threshold search and the mask on each row k selects; a row k leaves whole is the identity. + func: "tileops.perf.formulas.top_k_mask_roofline" + +MinPMaskFwdOp: + # Masks every logit whose probability is below min_p[b] * max(prob[b]), i.e. every + # logit below max_logit + log(min_p[b]). 0 < min_p <= 1 is the caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + min_p: {dtype: float32, shape: "[B]"} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + shape_rules: + - "V > 0" + workloads: + - {B: 1, V: 128256, dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {B: 64, V: 128256, dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {B: 256, V: 128256, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {B: 128, V: 151936, dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + roofline: + # Per element: the row max and the masking compare-select; per row: the threshold (log, add). + flops: "2 * B * V + 2 * B" + +TopPMaskFwdOp: + # Keeps the smallest set of highest-probability tokens whose cumulative probability + # reaches p[b]: a token survives while the exclusive cumulative sum before it is below + # p[b]. 0 < p < 1 is the caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + p: {dtype: float32, shape: "[B]"} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + shape_rules: + - "V > 0" + workloads: + - {B: 1, V: 128256, dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {B: 64, V: 128256, dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {B: 256, V: 128256, dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {B: 128, V: 151936, dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + - {B: 64, V: 201088, dtype_cases: [{T: bfloat16}], label: gpt-oss-120b-b64} + roofline: + # Per element: the softmax numerator (max, subtract, exp, sum), a weighted selection of + # the threshold (compare, accumulate) and the mask; per row: p times the sum. + flops: "7 * B * V + B" + +TopKTopPMaskFwdOp: + # TopPMaskFwdOp(TopKMaskFwdOp(logits, k), p): top-p's probabilities are renormalized + # over the tokens top-k keeps. 0 < p < 1 is the caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim, T: "DType[float16 | bfloat16 | float32]", k_list: "Seq[Int]"} + inputs: + logits: {dtype: T, shape: "[B, V]"} + k: {dtype: int32, shape: "[B]", values: "as_tensor(k_list)", requires: ["in_range(1, 2147483648)"]} + p: {dtype: float32, shape: "[B]"} + outputs: + masked_logits: {dtype: T, shape: "[B, V]"} + shape_rules: + - "V > 0" + workloads: + - {V: 128256, k_list: "repeat(50, 1)", dtype_cases: [{T: bfloat16}, {T: float32}], label: llama-8b-b1} + - {V: 128256, k_list: "repeat(50, 64)", dtype_cases: [{T: bfloat16}], label: llama-8b-b64} + - {V: 128256, k_list: "repeat(50, 256)", dtype_cases: [{T: float16}, {T: bfloat16}], label: llama-8b-b256} + - {V: 151936, k_list: "repeat(20, 128)", dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} + roofline: + # Top-k on the rows k restricts, top-p over each row's survivors, the final mask. + func: "tileops.perf.formulas.top_k_top_p_mask_roofline" + +SamplingFromProbsFwdOp: + # Draws one index per row with probability proportional to probs[b]; a zero-weight + # token is never drawn. Each row finite, non-negative and with a positive total is the + # caller's obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, V: Dim} + inputs: + probs: {dtype: float32, shape: "[B, V]"} + seed: {dtype: int64, shape: "[1]"} + offset: {dtype: int64, shape: "[1]"} + outputs: + samples: {dtype: int32, shape: "[B]"} + workloads: + - {B: 1, V: 128256, label: llama-8b-b1} + - {B: 64, V: 128256, label: llama-8b-b64} + - {B: 256, V: 128256, label: llama-8b-b256} + - {B: 128, V: 151936, label: qwen3-235b-b128} + - {B: 64, V: 201088, label: gpt-oss-120b-b64} + roofline: + # Per element: the row total, then the running sum and its compare against the draw; + # per row: the draw scaled by the total. + flops: "3 * B * V + B" + +ChainSpeculativeSamplingFwdOp: + # Verifies N draft tokens per request. Draft token i is accepted with probability + # min(1, target / draft); at the first rejection a token is drawn from + # normalize(max(0, target - draft)) and the rest of the row is -1; when every draft is + # accepted a bonus token is drawn from target row N. num_accepted is the accepted + # prefix length, the bonus excluded. Every probability row finite, non-negative and + # normalized, and each draft token of positive draft probability, is the caller's + # obligation. + family: sampling + status: spec-only + signature: + forall: {B: Dim, N: Dim, V: Dim} + inputs: + draft_probs: {dtype: float32, shape: "[B, N, V]"} + draft_token_ids: {dtype: int32, shape: "[B, N]", values: "topk_ids(B, N, V)", requires: ["in_range(0, V)"]} + target_probs: {dtype: float32, shape: "[B, N + 1, V]"} + seed: {dtype: int64, shape: "[1]"} + offset: {dtype: int64, shape: "[1]"} + outputs: + output_token_ids: {dtype: int32, shape: "[B, N + 1]"} + num_accepted: {dtype: int32, shape: "[B]"} + shape_rules: + - "N > 0" + workloads: + - {B: 16, N: 4, V: 128256, label: llama-8b-b16-n4} + - {B: 64, N: 4, V: 128256, label: llama-8b-b64-n4} + - {B: 64, N: 1, V: 129280, label: ds-v3-mtp-n1} + - {B: 64, N: 3, V: 129280, label: ds-v3-mtp-n3} + - {B: 32, N: 5, V: 151936, label: qwen3-235b-eagle-n5} + roofline: + # The work depends on where the chain stops, which the random draws decide, so both + # counts are the cheaper of the two extreme outcomes. All accepted: N ratio tests (a + # multiply and a compare), a draw from target row N (3 per token, 1 per row), reading + # the N token ids, 2 N probabilities and one target row. Rejected at the first: one + # test, the residual max(0, t - d) (2 per token) and a draw from it, reading one id and + # the two rows. + flops: "B * (2 * N + 3 * V + 1 if 2 * N + 3 * V + 1 < 5 * V + 3 else 5 * V + 3)" + bytes: >- + bytes(seed) + bytes(offset) + bytes(output_token_ids) + bytes(num_accepted) + + B * (12 * N + 4 * V if 12 * N + 4 * V < 4 + 8 * V else 4 + 8 * V) diff --git a/src/tileops/perf/formulas.py b/src/tileops/perf/formulas.py index 3c4800fad..b9980d600 100644 --- a/src/tileops/perf/formulas.py +++ b/src/tileops/perf/formulas.py @@ -22,11 +22,13 @@ "conv_roofline", "dsa_decode_roofline", "dsa_distinct_kv_rows", + "dsa_paged_fwd_roofline", "dsa_selected_keys", "fft_c2c_roofline", "fp8_lightning_indexer_roofline", "fused_moe_fwd_roofline", "fused_moe_shared_expert_fwd_roofline", + "fused_qk_norm_rope_roofline", "gqa_dense_fwd_roofline", "gqa_paged_cache_rows", "gqa_paged_fwd_roofline", @@ -34,6 +36,9 @@ "gqa_prefill_paged_with_kv_cache_fwd_roofline", "gqa_varlen_fwd_roofline", "lightning_indexer_scored_keys", + "mla_kv_cache_write_roofline", + "mla_paged_fwd_roofline", + "mla_varlen_fwd_roofline", "moe_expert_mlp_roofline", "moe_grouped_gemm_roofline", "moe_layout_active_experts", @@ -47,8 +52,12 @@ "nsa_topk_varlen_roofline", "paged_decode_cache_rows", "paged_decode_roofline", + "paged_kv_cache_gather_roofline", + "paged_kv_cache_write_roofline", "paged_rows", "pool_roofline", + "top_k_mask_roofline", + "top_k_top_p_mask_roofline", "topk_selector_roofline", "topk_selector_window_scores", "visible_score_rows", @@ -759,3 +768,182 @@ def gqa_prefill_paged_cache_rows(call: "CallView") -> int: # A request with no new token attends nothing and reads none of its cache. read_lens = [c if q else 0 for q, c in zip(q_lens, cache_lens, strict=True)] return paged_rows(call.values("block_table"), read_lens, call.ix["page_size"])[0] + + +# ---------------------------------------------------------------- paged caches and MLA + + +_FP8 = "float8_e4m3fn" + + +def _elem_bytes(call: "CallView", name: str) -> int: + """Bytes of one element of tensor *name*.""" + return call.bytes(name) // max(1, prod(call.tensors[name][0])) + + +def _paged_cache_read_bytes( + call: "CallView", cache: str, table: str, ends: list, starts: "list | None" = None +) -> int: + """Bytes of the distinct rows of the paged *cache* ``[NP, PS, ...]`` that the requests' + ranges ``[starts[b], ends[b])`` reach through *table*, plus the table entries consulted.""" + shape = call.tensors[cache][0] + rows, consulted = paged_rows(call.values(table), ends, shape[1], starts) + row_bytes = call.bytes(cache) // max(1, shape[0] * shape[1]) + return rows * row_bytes + consulted * _elem_bytes(call, table) + + +def mla_paged_fwd_roofline(call: "CallView") -> tuple[int, int]: + """Paged MLA decode: QK over the whole latent row, PV over its first ``kv_lora_rank`` + columns, for the scores each query row sees in its request's cache. + + An FP8 cache's per-tensor scale folds into the score scale and, once per query row and + head, into the softmax denominator. The cache is read at the rows the block table reaches; + every other tensor moves once. + """ + ix = call.ix + lengths = call.values("cache_seqlens") + pairs = [visible_score_rows(ix["S_q"], c, ix["is_causal"], -1, -1) for c in lengths] + scores, rows = sum(p[0] for p in pairs), sum(p[1] for p in pairs) + flops = attention_flops(ix["H"], scores, rows, ix["DK"], ix["kv_lora_rank"]) + if call.tensors["kv_cache"][1] == _FP8: + flops += ix["H"] * rows + moved = _derived_bytes(call) - call.bytes("kv_cache") - call.bytes("block_table") + moved += _paged_cache_read_bytes(call, "kv_cache", "block_table", lengths) + return flops, moved + + +def mla_varlen_fwd_roofline(call: "CallView") -> tuple[int, int]: + """Packed MLA prefill: QK over ``DN + PE``, PV over ``DV``, for the scores each request + sees under its mask; every tensor moves once.""" + ix = call.ix + pairs = [ + visible_score_rows(n, n, ix["is_causal"], -1, -1) for n in _segments(call, "cu_seqlens") + ] + scores, rows = sum(p[0] for p in pairs), sum(p[1] for p in pairs) + flops = attention_flops(ix["H"], scores, rows, ix["DN"] + ix["PE"], ix["DV"]) + return flops, _derived_bytes(call) + + +# FlashMLA DeepSeek-V3.2 FP8 cache row: the QK width and the value (latent) width. +_DSA_QK_DIM, _DSA_V_DIM = 576, 512 + + +def dsa_paged_fwd_roofline(call: "CallView") -> tuple[int, int]: + """Paged DeepSeek sparse attention over an FP8 latent cache. + + Every index slot ``0 <= j < cache_seqlens[b]`` is a score, a repeated slot once per + occurrence. The latent of each distinct cache row is dequantized once (a multiply per + value) and read once; the block table is read at the pages the slots resolve to; ``q`` is + read at the query rows with a score; every other tensor moves once. + """ + ix = call.ix + lengths, table = call.values("cache_seqlens"), call.values("block_table") + page_size = call.tensors["kv_cache"][0][1] + scores = rows = 0 + distinct: set = set() + pages: set = set() + for b, per_query in enumerate(call.values("indices")): + for slots in per_query: + valid = [j for j in slots if 0 <= j < lengths[b]] + scores += len(valid) + rows += bool(valid) + for j in valid: + pages.add((b, j // page_size)) + distinct.add((table[b][j // page_size], j % page_size)) + flops = attention_flops(ix["H"], scores, rows, _DSA_QK_DIM, _DSA_V_DIM) + flops += _DSA_V_DIM * len(distinct) + shape = call.tensors["kv_cache"][0] + moved = _derived_bytes(call) - call.bytes("q") - call.bytes("kv_cache") + moved -= call.bytes("block_table") + moved += rows * ix["H"] * _DSA_QK_DIM * _elem_bytes(call, "q") + moved += len(distinct) * shape[2] + len(pages) * _elem_bytes(call, "block_table") + return flops, moved + + +def _written_slots(call: "CallView") -> int: + """Tokens a cache write stores: the ``slot_mapping`` entries other than -1.""" + return sum(1 for s in call.values("slot_mapping") if s >= 0) + + +def paged_kv_cache_write_roofline(call: "CallView") -> tuple[int, int]: + """Paged K/V cache write: each token with a slot reads its k and v rows and writes them + into both caches; an FP8 cache scales and saturates each value (2 per value).""" + shape = call.tensors["k"][0] + tokens, row = _written_slots(call), prod(shape[1:]) + flops = 4 * tokens * row if call.tensors["k_pages"][1] == _FP8 else 0 + moved = 2 * tokens * row * (_elem_bytes(call, "k") + _elem_bytes(call, "k_pages")) + moved += call.bytes("slot_mapping") + if tokens * row: + moved += sum(call.bytes(s) for s in ("k_scale", "v_scale") if call.present(s)) + return flops, moved + + +def mla_kv_cache_write_roofline(call: "CallView") -> tuple[int, int]: + """MLA latent-cache write: each token with a slot reads kv_c and k_pe and writes the + concatenated row; an FP8 cache scales and saturates each value (2 per value), and a fused + RoPE rotates k_pe (3 per value) from the cos/sin rows the tokens' positions name.""" + ix = call.ix + width, pe = ix["DC"] + ix["PE"], ix["PE"] + slots = call.values("slot_mapping") + tokens = _written_slots(call) + flops = 2 * tokens * width if call.tensors["kv_cache"][1] == _FP8 else 0 + moved = tokens * width * (_elem_bytes(call, "kv_c") + _elem_bytes(call, "kv_cache")) + moved += call.bytes("slot_mapping") + moved += call.bytes("scale") if tokens * width and call.present("scale") else 0 + if ix["fuse_rope"]: + flops += 3 * tokens * pe + positions = [p for p, s in zip(call.values("positions"), slots, strict=True) if s >= 0] + moved += len(positions) * _elem_bytes(call, "positions") + moved += len(set(positions)) * pe * _elem_bytes(call, "cos_sin_cache") + return flops, moved + + +def paged_kv_cache_gather_roofline(call: "CallView") -> tuple[int, int]: + """Paged cache gather: the cache rows each request's range reaches are read once and + written to ``dst``; an FP8 cache is dequantized with a multiply per value.""" + lengths = _segments(call, "cu_seq_lens") + starts = call.values("seq_starts") if call.present("seq_starts") else [0] * len(lengths) + ends = [s + n for s, n in zip(starts, lengths, strict=True)] + fp8 = call.tensors["cache"][1] == _FP8 + flops = prod(call.tensors["dst"][0]) if fp8 else 0 + moved = _derived_bytes(call) - call.bytes("cache") - call.bytes("block_table") + moved += _paged_cache_read_bytes(call, "cache", "block_table", ends, starts) + if call.present("scale") and not prod(call.tensors["dst"][0]): + moved -= call.bytes("scale") # nothing is dequantized + return flops, moved + + +def fused_qk_norm_rope_roofline(call: "CallView") -> tuple[int, int]: + """Q/K RMSNorm and RoPE in place: per q and k value, RMSNorm with its weight (4), and 3 + per rotated value. The q and k columns are read and written, the v columns untouched, + and the cos/sin rows the positions name read once.""" + ix = call.ix + heads = ix["num_heads"] + ix["num_kv_heads"] + values = ix["N"] * heads * ix["D"] + flops = ix["N"] * heads * (4 * ix["D"] + 3 * ix["R"]) + moved = 2 * values * _elem_bytes(call, "qkv") + moved += sum(call.bytes(t) for t in ("q_weight", "k_weight", "positions")) + moved += len(set(call.values("positions"))) * ix["R"] * _elem_bytes(call, "cos_sin_cache") + return flops, moved + + +# ---------------------------------------------------------------- sampling + + +def top_k_mask_roofline(call: "CallView") -> tuple[int, int]: + """Top-k logit mask: a threshold selection (1 per logit) and the mask (1 per logit) on + each row ``k`` restricts; a row with ``k >= V`` is the identity. Every tensor moves once.""" + vocab = call.ix["V"] + return 2 * vocab * sum(1 for k in call.values("k") if k < vocab), _derived_bytes(call) + + +def top_k_top_p_mask_roofline(call: "CallView") -> tuple[int, int]: + """Top-k then top-p logit mask, per row: the top-k threshold (1 per logit) where ``k`` + restricts the row; over the ``min(k, V)`` survivors the softmax numerator (max, + subtract, exp, sum) and the weighted threshold selection (compare, accumulate); ``p`` + times the sum; the final mask (1 per logit). Every tensor moves once.""" + vocab = call.ix["V"] + flops = sum( + (vocab if k < vocab else 0) + 6 * min(k, vocab) + 1 + vocab for k in call.values("k") + ) + return flops, _derived_bytes(call) diff --git a/src/tileops/sampling.py b/src/tileops/sampling.py new file mode 100644 index 000000000..2803fe20e --- /dev/null +++ b/src/tileops/sampling.py @@ -0,0 +1,3 @@ +"""The sampling ops, at the public path ``tileops.sampling``.""" + +__all__: list[str] = [] diff --git a/tests/test_manifest_generators.py b/tests/test_manifest_generators.py index 1994bbe2e..09e8fef38 100644 --- a/tests/test_manifest_generators.py +++ b/tests/test_manifest_generators.py @@ -102,3 +102,19 @@ def test_sample_indices_draws_distinct_values_in_range(): assert sorted(GENERATORS["sample_indices"](random.Random(0), 5, 5)) == list(range(5)) with pytest.raises(ValueError): GENERATORS["sample_indices"](random.Random(0), 3, 2) + + +def test_sparse_topk_positions_draw_visible_positions_and_pad(): + # Request lengths 3 and 5, two queries: query s of a length-n request sees n - 1 + s positions. + rows = GENERATORS["attn.sparse_topk_positions"](random.Random(0), [3, 5], 2, 4) + assert (len(rows), len(rows[0]), len(rows[0][0])) == (2, 2, 4) + for n, per_query in zip([3, 5], rows, strict=True): + for s, slots in enumerate(per_query): + visible = n - 1 + s + picked = [j for j in slots if j != -1] + assert len(picked) == len(set(picked)) == min(4, visible) + assert all(0 <= j < visible for j in picked) and slots[len(picked) :] == [-1] * ( + 4 - len(picked) + ) + with pytest.raises(ValueError): + GENERATORS["attn.sparse_topk_positions"](random.Random(0), [1], 2, 4) diff --git a/tests/test_public_api.py b/tests/test_public_api.py index f19f01797..4e6fc01dd 100644 --- a/tests/test_public_api.py +++ b/tests/test_public_api.py @@ -112,10 +112,13 @@ def test_each_public_name_has_exactly_one_family(): @pytest.mark.smoke def test_public_surface_is_the_manifest(): - """An op reachable as `tileops..` has a manifest entry, and every entry - is reachable. Abstract bases are not public at all.""" + """An op reachable as `tileops..` has a manifest entry, and every + implemented entry is reachable; a spec-only entry may have no class yet. Abstract + bases are not public at all.""" public = {name for family in FAMILIES for name in _family_module(family).__all__} - assert public == set(load_manifest()) + manifest = load_manifest() + assert public <= set(manifest) + assert {name for name, entry in manifest.items() if entry["status"] == "implemented"} <= public @pytest.mark.smoke From 678b45933a3d4dce2e5060820629be7f91459698 Mon Sep 17 00:00:00 2001 From: lcy-seso Date: Sun, 27 Sep 2026 23:02:55 +0800 Subject: [PATCH 2/5] [Chore][Manifest] Trim comments on the new entries --- src/tileops/manifest/spec/attention.yaml | 7 +++---- src/tileops/manifest/spec/sampling.yaml | 2 +- 2 files changed, 4 insertions(+), 5 deletions(-) diff --git a/src/tileops/manifest/spec/attention.yaml b/src/tileops/manifest/spec/attention.yaml index 78a59e2b3..12f62576f 100644 --- a/src/tileops/manifest/spec/attention.yaml +++ b/src/tileops/manifest/spec/attention.yaml @@ -544,8 +544,7 @@ MultiHeadLatentAttentionPagedFwdOp: MultiHeadLatentAttentionVarlenFwdOp: # MLA prefill over packed requests after the latent is decompressed: the key of head h # is k_nope[:, h] ‖ k_pe, with k_pe shared by every head; queries and keys are the same - # tokens. The request length bound follows from cu_seqlens. sm_scale defaults to - # (DN + PE) ** -0.5. + # tokens. sm_scale defaults to (DN + PE) ** -0.5. family: attention status: spec-only signature: @@ -706,8 +705,8 @@ DeepSeekSparseAttentionPagedFwdOp: PagedKVCacheGatherFwdOp: # Gathers request b's cache positions [start_b, start_b + len_b) into rows # cu_seq_lens[b] .. cu_seq_lens[b + 1] of dst, start_b = seq_starts[b] or 0; an FP8 - # cache is dequantized by scale into out_dtype. dst is passed because its length is - # cu_seq_lens[-1], a value no shape carries. + # cache is dequantized by scale into out_dtype. dst carries the length cu_seq_lens[-1], + # which no input shape states. family: attention status: spec-only signature: diff --git a/src/tileops/manifest/spec/sampling.yaml b/src/tileops/manifest/spec/sampling.yaml index 6cf300d2f..431fd72cc 100644 --- a/src/tileops/manifest/spec/sampling.yaml +++ b/src/tileops/manifest/spec/sampling.yaml @@ -24,7 +24,7 @@ TopKMaskFwdOp: - {V: 151936, k_list: "repeat(20, 128)", dtype_cases: [{T: bfloat16}], label: qwen3-235b-b128} - {V: 201088, k_list: "repeat(50, 64)", dtype_cases: [{T: bfloat16}], label: gpt-oss-120b-b64} roofline: - # The threshold search and the mask on each row k selects; a row k leaves whole is the identity. + # The threshold search and the mask on each row k restricts; a row with k >= V is the identity. func: "tileops.perf.formulas.top_k_mask_roofline" MinPMaskFwdOp: From b30b37b4b01109c5ce169d8b74cbf8770b2c27bf Mon Sep 17 00:00:00 2001 From: lcy-seso Date: Sun, 27 Sep 2026 23:46:37 +0800 Subject: [PATCH 3/5] [Refactor] Route dtype logic through the registry and hoist module constants The dtype registry now records each dtype's category next to its bits, and derives FLOAT8_DTYPES from both. The roofline formulas that price an 8-bit float cache test membership in it rather than naming float8_e4m3fn, the category primitive reads the registry, and integer promotion asks for the category instead of listing the integer names. The elementwise workload's signed-integer branch uses torch's own dtype properties. Module-level constants defined after the first function or class move above it, unchanged. An assignment whose value reads a later module-level name, or calls anything but a pure builtin constructor, stays where it is. Every manifest row's roofline value and the validator output are unchanged. --- benchmarks/conftest.py | 7 +- benchmarks/hardware/memory/hbm_bandwidth.py | 6 +- benchmarks/ops/bench_fused_gated.py | 35 ++-- benchmarks/ops/bench_pool.py | 29 ++- benchmarks/tests/test_benchmark_boundaries.py | 20 +-- scripts/lint/tilelang_idioms_lint.py | 36 ++-- scripts/nightly_report.py | 38 ++-- scripts/validate_manifest.py | 38 ++-- scripts/validate_roofline_bytes.py | 12 +- src/tileops/backend/protocol.py | 29 ++- src/tileops/kernels/attention/call_spec.py | 18 +- src/tileops/kernels/attention/gqa_dense.py | 51 +++--- src/tileops/kernels/attention/gqa_fwd.py | 16 +- src/tileops/kernels/elementwise/_broadcast.py | 7 +- src/tileops/kernels/elementwise/_dtype.py | 18 +- src/tileops/kernels/fft.py | 168 +++++++++--------- src/tileops/kernels/gemm/dense.py | 8 +- src/tileops/kernels/gemm/heuristics.py | 10 +- src/tileops/kernels/gemm/w4a16.py | 8 +- .../kernels/grouped_gemm/heuristics.py | 12 +- src/tileops/kernels/norm/fused_add_norm.py | 7 +- src/tileops/kernels/pool/common.py | 9 +- src/tileops/kernels/reduction/_primitives.py | 10 +- src/tileops/kernels/reduction/call_spec.py | 12 +- .../kernels/reduction/logical_reduce.py | 18 +- src/tileops/kernels/reduction/reduce.py | 6 +- src/tileops/manifest/dtype_rules.py | 49 ++--- src/tileops/manifest/expr.py | 19 +- src/tileops/manifest/primitives.py | 46 +++-- src/tileops/ops/_signature_codegen.py | 6 +- src/tileops/ops/elementwise/_base.py | 21 +-- src/tileops/ops/elementwise/arithmetic.py | 5 +- src/tileops/ops/op_base.py | 9 +- src/tileops/ops/pool.py | 22 +-- src/tileops/perf/formulas.py | 39 ++-- src/tileops/perf/profile.py | 24 +-- src/tileops/trace/record.py | 26 ++- src/tileops/trace/ui.py | 108 ++++++----- src/tileops/utils/utils.py | 22 +-- workloads/elementwise.py | 2 +- workloads/gemm.py | 6 +- 41 files changed, 491 insertions(+), 541 deletions(-) diff --git a/benchmarks/conftest.py b/benchmarks/conftest.py index d55b0a0c8..d3c1044d1 100644 --- a/benchmarks/conftest.py +++ b/benchmarks/conftest.py @@ -8,6 +8,9 @@ import benchmarks.baselines # noqa: F401 from benchmarks.report import BenchmarkReport, _bench_results +# What a row carries besides its measurements. +_NOT_A_MEASUREMENT = frozenset({"tag", "op", "op_module", "ops", "params", "run_config", "result"}) + def pytest_make_parametrize_id(config, val, argname): """Render the values pytest would otherwise collect as `shape0`, `dtype0`. @@ -30,10 +33,6 @@ def pytest_make_parametrize_id(config, val, argname): return None -# What a row carries besides its measurements. -_NOT_A_MEASUREMENT = frozenset({"tag", "op", "op_module", "ops", "params", "run_config", "result"}) - - def _prop(value) -> str: """Format one measurement for the XML. diff --git a/benchmarks/hardware/memory/hbm_bandwidth.py b/benchmarks/hardware/memory/hbm_bandwidth.py index 8088a69cc..bc64de467 100644 --- a/benchmarks/hardware/memory/hbm_bandwidth.py +++ b/benchmarks/hardware/memory/hbm_bandwidth.py @@ -30,6 +30,9 @@ _CU_SRC = Path(__file__).parent / "hbm_saturation.cu" +_MIXES = ("copy", "triad", "read", "write") + + def _compile(cu_path, binary_path, arch="sm_90"): """Compile the CUDA source. Raises on failure.""" cmd = [ @@ -57,9 +60,6 @@ def _run(binary_path, size_mb, theo_peak_gbs): return result.stdout.strip().splitlines() -_MIXES = ("copy", "triad", "read", "write") - - def _parse_peaks(lines): """Best bandwidth (GB/s) per access mix from the CSV output. diff --git a/benchmarks/ops/bench_fused_gated.py b/benchmarks/ops/bench_fused_gated.py index b144e2a0f..e05c66654 100644 --- a/benchmarks/ops/bench_fused_gated.py +++ b/benchmarks/ops/bench_fused_gated.py @@ -35,6 +35,22 @@ ) from workloads.workload_base import FixtureBase +# Scenario -> (tokens, width). +_STRATEGY_SHAPES = { + "llama-hidden-1k-tokens": (1024, 4096), + "llama-7b-ffn-1k-tokens": (1024, 11008), + "llama-hidden-4k-tokens": (4096, 4096), +} +_STRATEGY_DTYPES = (torch.float16, torch.bfloat16, torch.float32) +_STRATEGY_KERNELS = [ + ("silu_and_mul", SiluAndMulFwdKernel), + ("gelu_and_mul", GeluAndMulFwdKernel), + ("gelu_tanh_and_mul", GeluTanhAndMulFwdKernel), +] +# How far behind the fastest strategy the default may sit before the choice is +# stale. Wide enough to clear run-to-run spread, narrow enough to flag a flip. +_STRATEGY_MARGIN = 1.25 + class FusedGatedBenchmark(BenchmarkBase[FusedGatedBenchCase]): """Times the strategy decision; it records no row, so both metrics are ``None``.""" @@ -82,20 +98,6 @@ def test_gelu_tanh_and_mul_bench(call) -> None: _profile_fused_gated(GeluTanhAndMulFwdOp, call, "gelu_tanh_and_mul") -# Scenario -> (tokens, width). -_STRATEGY_SHAPES = { - "llama-hidden-1k-tokens": (1024, 4096), - "llama-7b-ffn-1k-tokens": (1024, 11008), - "llama-hidden-4k-tokens": (4096, 4096), -} -_STRATEGY_DTYPES = (torch.float16, torch.bfloat16, torch.float32) -_STRATEGY_KERNELS = [ - ("silu_and_mul", SiluAndMulFwdKernel), - ("gelu_and_mul", GeluAndMulFwdKernel), - ("gelu_tanh_and_mul", GeluTanhAndMulFwdKernel), -] - - def _strategy_params(): """Default-strategy sentinel: shape and dtype axes on the first kernel, plus one reference-point direct-vs-explicit sentinel per remaining kernel. @@ -127,11 +129,6 @@ class FusedGatedStrategyBenchFixture(FixtureBase): PARAMS = [("op_name, M, N, dtype, kernel_cls", _strategy_params())] -# How far behind the fastest strategy the default may sit before the choice is -# stale. Wide enough to clear run-to-run spread, narrow enough to flag a flip. -_STRATEGY_MARGIN = 1.25 - - @FusedGatedStrategyBenchFixture def test_fused_gated_default_strategy_is_the_fast_one( op_name: str, diff --git a/benchmarks/ops/bench_pool.py b/benchmarks/ops/bench_pool.py index f41677cfd..827adbbda 100644 --- a/benchmarks/ops/bench_pool.py +++ b/benchmarks/ops/bench_pool.py @@ -38,6 +38,19 @@ from workloads.pool import MeanPoolingCallWorkload, MeanPoolingWorkload from workloads.workload_base import CallWorkload +# Which library serves an op, and the pooling kind and rank its adapter needs. An op absent +# here has none: no library covers 1D, adaptive pooling, or 3D max-pool indices. Every row +# is also timed against torch, eager and compiled, so this table is not the whole baseline. +_BASELINE: dict[str, tuple[str, str, int]] = { + "AvgPool2dFwdOp": (FLAGGEMS_TAG, "avg", 2), + "AvgPool3dFwdOp": ("cudnn", "avg", 3), + "MaxPool2dFwdOp": (FLAGGEMS_TAG, "max", 2), + "MaxPool2dIndicesFwdOp": (FLAGGEMS_TAG, "max", 2), + "MaxPool3dFwdOp": ("cudnn", "max", 3), +} +# Autotuning is a bench-run policy; manifest workloads do not carry it. +_TUNE = True + def flaggems_pool_fn( kind: str, @@ -176,18 +189,6 @@ def run_max(x: torch.Tensor): return None -# Which library serves an op, and the pooling kind and rank its adapter needs. An op absent -# here has none: no library covers 1D, adaptive pooling, or 3D max-pool indices. Every row -# is also timed against torch, eager and compiled, so this table is not the whole baseline. -_BASELINE: dict[str, tuple[str, str, int]] = { - "AvgPool2dFwdOp": (FLAGGEMS_TAG, "avg", 2), - "AvgPool3dFwdOp": ("cudnn", "avg", 3), - "MaxPool2dFwdOp": (FLAGGEMS_TAG, "max", 2), - "MaxPool2dIndicesFwdOp": (FLAGGEMS_TAG, "max", 2), - "MaxPool3dFwdOp": ("cudnn", "max", 3), -} - - def _as_tuple(value, ndim: int) -> tuple: if isinstance(value, (tuple, list)): return tuple(value) @@ -342,10 +343,6 @@ def test_adaptive_max_pool2d_indices_bench(call) -> None: # MeanPoolingFwdOp, the chunked sequence mean. -# Autotuning is a bench-run policy; manifest workloads do not carry it. -_TUNE = True - - def _torch_view_mean(workload: MeanPoolingWorkload): """The same mean over a reshaped view, or None where the chunks are ragged. diff --git a/benchmarks/tests/test_benchmark_boundaries.py b/benchmarks/tests/test_benchmark_boundaries.py index 3e1b0e7ff..5a54f21fb 100644 --- a/benchmarks/tests/test_benchmark_boundaries.py +++ b/benchmarks/tests/test_benchmark_boundaries.py @@ -18,6 +18,16 @@ BENCHMARK_DIRS = ("benchmarks/ops",) +# A benchmark takes (flops, bytes) from its op — docs/design/roofline.md §4.2. An entry +# here declares the two methods for a reason the name below states. An entry whose +# subject is an op goes as soon as that op gains a manifest entry; an entry whose +# subject is not an op stays, because a manifest entry is something only an op can have. +_ROOFLINE_OF_ITS_OWN = { + "FusedGatedBenchmark": "times a forced kernel strategy, which no op can request and " + "no report has a row for; both metrics return None", +} + + def _benchmark_files() -> list[Path]: return [ path @@ -72,16 +82,6 @@ def test_benchmarks_do_not_author_gen_inputs() -> None: assert _scan(_defines_gen_inputs) == {} -# A benchmark takes (flops, bytes) from its op — docs/design/roofline.md §4.2. An entry -# here declares the two methods for a reason the name below states. An entry whose -# subject is an op goes as soon as that op gains a manifest entry; an entry whose -# subject is not an op stays, because a manifest entry is something only an op can have. -_ROOFLINE_OF_ITS_OWN = { - "FusedGatedBenchmark": "times a forced kernel strategy, which no op can request and " - "no report has a row for; both metrics return None", -} - - def _writes_its_own_roofline(tree: ast.AST) -> list[str]: return [ f"{node.name}.{fn.name} (line {fn.lineno})" diff --git a/scripts/lint/tilelang_idioms_lint.py b/scripts/lint/tilelang_idioms_lint.py index 907057bae..99155f50d 100755 --- a/scripts/lint/tilelang_idioms_lint.py +++ b/scripts/lint/tilelang_idioms_lint.py @@ -51,6 +51,22 @@ _DTYPE_NAME = re.compile(r"^(u?int[0-9]+|b?float[0-9]+|float8[a-z0-9_]*|bool|handle)$") +_NONSCALAR_KINDS = { + ast.List: "list", + ast.ListComp: "list", + ast.Dict: "dict", + ast.DictComp: "dict", + ast.Set: "set", + ast.SetComp: "set", + ast.Tuple: "tuple", + ast.GeneratorExp: "generator", + ast.Lambda: "function", +} +_SCOPE = (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef) +_Func = ast.FunctionDef | ast.AsyncFunctionDef +_SCALAR_ANNOTATIONS = frozenset({"int", "float", "str", "bool", "None", "NoneType"}) + + def _attr_path(node: ast.AST) -> str | None: """Dotted name of an attribute chain, e.g. ``T.reinterpret``; None otherwise.""" parts = [] @@ -150,23 +166,6 @@ def _arg(call: ast.Call, pos: int, name: str) -> ast.AST | None: return next((k.value for k in call.keywords if k.arg == name), None) -_NONSCALAR_KINDS = { - ast.List: "list", - ast.ListComp: "list", - ast.Dict: "dict", - ast.DictComp: "dict", - ast.Set: "set", - ast.SetComp: "set", - ast.Tuple: "tuple", - ast.GeneratorExp: "generator", - ast.Lambda: "function", -} - -_SCOPE = (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef) - -_Func = ast.FunctionDef | ast.AsyncFunctionDef - - def _jit_names(tree: ast.Module) -> tuple[set[str], set[str]]: """What this file binds to ``tilelang.jit``: module aliases, then bare names. @@ -239,9 +238,6 @@ def _function_tables(top: symtable.SymbolTable) -> dict[tuple[str, int], symtabl return tables -_SCALAR_ANNOTATIONS = frozenset({"int", "float", "str", "bool", "None", "NoneType"}) - - def _string_annotation_kind(annotation: str) -> str | None: """Classify a quoted annotation.""" try: diff --git a/scripts/nightly_report.py b/scripts/nightly_report.py index 5f21e2a26..77da964ad 100755 --- a/scripts/nightly_report.py +++ b/scripts/nightly_report.py @@ -85,6 +85,22 @@ # --------------------------------------------------------------------------- +# A dtype name as a case id spells it: the id ends in its dtype cases and dtype parameters. +_DTYPE_TOKEN = re.compile(r"bfloat16|bool|u?int\d+|float\d+(?:_[a-z0-9]+)?|complex\d+") +# Verdicts are drawn on device execution time, not on the span that also covers the +# gaps between a call's kernels. +_CONCLUSION_KEY = "device_busy_ms" +#: Absorbs the rounding a count recovered from ``tflops`` carries. +_DERIVED_COUNT_RTOL = 1e-3 +# A kernel file below this share of executed lines was never constructed by any +# test. The lowest genuinely-built kernel file measures 36.9%, so the threshold +# has room before it starts catching built kernels. +_KERNEL_BUILT_PCT = 25 +# Below this, one statement swings the percentage too far to read. +_COVERAGE_MIN_STMTS = 20 +_COVERAGE_WORST_N = 15 # rows in the least-covered file list + + def _get_properties(testcase: ET.Element) -> dict[str, str]: """Extract user properties from a JUnit testcase element.""" props = {} @@ -301,10 +317,6 @@ def _case_id(name: str) -> str: return name -# A dtype name as a case id spells it: the id ends in its dtype cases and dtype parameters. -_DTYPE_TOKEN = re.compile(r"bfloat16|bool|u?int\d+|float\d+(?:_[a-z0-9]+)?|complex\d+") - - def _dtypes_of(name: str) -> tuple[str, ...]: """The dtype names in a row's case id, in order.""" return tuple(t for t in _case_id(name).split("-") if _DTYPE_TOKEN.fullmatch(t)) @@ -353,11 +365,6 @@ def history_window(runs: list[dict], retention_days: int = HISTORY_RETENTION_DAY return sorted(kept.values(), key=lambda r: r["date"]) -# Verdicts are drawn on device execution time, not on the span that also covers the -# gaps between a call's kernels. -_CONCLUSION_KEY = "device_busy_ms" - - def _conclusion(cfg: dict) -> tuple[float | None, str]: """The reading a verdict is drawn on, and which key it came from. @@ -376,10 +383,6 @@ def _conclusion_ms(cfg: dict) -> float | None: return _conclusion(cfg)[0] -#: Absorbs the rounding a count recovered from ``tflops`` carries. -_DERIVED_COUNT_RTOL = 1e-3 - - class _WorkCounts(NamedTuple): flops: float | None nbytes: float | None @@ -1131,15 +1134,6 @@ def _pct(hit: int, total: int) -> str: return f"{100 * hit / total:.1f}%" if total else "-" -# A kernel file below this share of executed lines was never constructed by any -# test. The lowest genuinely-built kernel file measures 36.9%, so the threshold -# has room before it starts catching built kernels. -_KERNEL_BUILT_PCT = 25 -# Below this, one statement swings the percentage too far to read. -_COVERAGE_MIN_STMTS = 20 -_COVERAGE_WORST_N = 15 # rows in the least-covered file list - - def _coverage_signals(files: list[dict]) -> dict: """Reduce per-file coverage to the three numbers worth acting on. diff --git a/scripts/validate_manifest.py b/scripts/validate_manifest.py index a2922fe9a..c31adaac1 100755 --- a/scripts/validate_manifest.py +++ b/scripts/validate_manifest.py @@ -56,6 +56,23 @@ _STAGE_KEYS = {"name", "op", "kernel", "optional"} +_ENTRY_KEYS = { + "family": str, + "status": str, + "signature": dict, + "workloads": list, + "roofline": dict, + "ref_api": str, + "composition": dict, +} +_REQUIRED = ("family", "status", "signature", "workloads", "roofline") +_OP_KEY = re.compile(r"[A-Z][A-Za-z0-9]*(Fwd|Bwd)Op") +# Execution-policy parameters every op takes, in order with their defaults, and the reserved one +# it may take (docs/design/manifest.md § Signature). +_POLICY_PARAMETERS = {"target": None, "kernel_map": None, "tune": False} +_RESERVED_POLICY = "config" + + def _key_format_errors( op_name: str, all_op_names: Collection[str], @@ -232,18 +249,6 @@ def _check_bench_files(repo_root: Path) -> list[str]: return errors -_ENTRY_KEYS = { - "family": str, - "status": str, - "signature": dict, - "workloads": list, - "roofline": dict, - "ref_api": str, - "composition": dict, -} -_REQUIRED = ("family", "status", "signature", "workloads", "roofline") - - def _schema_errors(op_name: str, entry: dict, all_op_names) -> list[str]: """Top-level fields of an entry (docs/design/manifest.md § Top-Level Fields).""" if not isinstance(op_name, str): @@ -273,9 +278,6 @@ def _schema_errors(op_name: str, entry: dict, all_op_names) -> list[str]: return errors -_OP_KEY = re.compile(r"[A-Z][A-Za-z0-9]*(Fwd|Bwd)Op") - - def _family_errors(op_name: str, entry: dict) -> list[str]: """`family` names a public module; an implemented op is exported from it by its key.""" family = entry.get("family") @@ -317,12 +319,6 @@ def _ref_api_errors(op_name: str, ref: str) -> list[str]: return [f"[schema] {op_name}: ref_api {ref!r}: no prefix of it is an importable module"] -# Execution-policy parameters every op takes, in order with their defaults, and the reserved one -# it may take (docs/design/manifest.md § Signature). -_POLICY_PARAMETERS = {"target": None, "kernel_map": None, "tune": False} -_RESERVED_POLICY = "config" - - def _normal_default(value): """A default as the manifest writes it: a dtype by name, a tuple as a list.""" if isinstance(value, tuple): diff --git a/scripts/validate_roofline_bytes.py b/scripts/validate_roofline_bytes.py index 8fac541a0..84fa84b0a 100755 --- a/scripts/validate_roofline_bytes.py +++ b/scripts/validate_roofline_bytes.py @@ -48,12 +48,6 @@ COLD_CACHE_PREMISE = "cold-cache replay (ncu --cache-control all)" -def _op_class(op_name: str, entry: dict): - from tileops.manifest.registry import op_class - - return op_class(op_name, entry) - - # The calls whose read half is not a lower bound: op name -> (condition over the call's ``ix``, # reason). READ_BOUND_EXCEPTIONS: dict = { @@ -68,6 +62,12 @@ def _op_class(op_name: str, entry: dict): } +def _op_class(op_name: str, entry: dict): + from tileops.manifest.registry import op_class + + return op_class(op_name, entry) + + def _call(op_name: str, entry: dict, row: dict, case: dict): """One row and dtype case of an entry, instantiated.""" from tileops.manifest import load_adts diff --git a/src/tileops/backend/protocol.py b/src/tileops/backend/protocol.py index f3dbb71c1..cb63c0fa8 100644 --- a/src/tileops/backend/protocol.py +++ b/src/tileops/backend/protocol.py @@ -6,36 +6,33 @@ import torch - -class TensorSpec(NamedTuple): - """What one tensor is, without the tensor. Handed to ``build_kernel``.""" - - device: torch.device - dtype: torch.dtype - shape: tuple[int, ...] - - @staticmethod - def of(tensor: torch.Tensor) -> "TensorSpec": - """Describe *tensor*.""" - return TensorSpec(tensor.device, tensor.dtype, tuple(tensor.shape)) - - # One call's result. A purely mutating op returns ``None``: ``torch.library.custom_op`` # cannot express a return value aliasing an input. KernelResult = Union[torch.Tensor, tuple[torch.Tensor, ...], None] - # Called ``build_kernel(*inputs, **params)``: a `TensorSpec` per input in # ``signature.inputs`` order — ``None`` for an ``optional: true`` input the call did not # pass, so presence is read off the slot rather than off how many slots there are — then # ``signature.params`` by keyword. Both lists are per-op, which the type system cannot # express, hence ``...``. BuildKernel = Callable[..., Callable[..., KernelResult]] - # "Is this the kind of device my kernels are written for" — ``False``, not an exception, # for the rest. Per-call support is ``build_kernel``'s answer; it sees the dtypes too. DetectFn = Callable[[torch.device], bool] +class TensorSpec(NamedTuple): + """What one tensor is, without the tensor. Handed to ``build_kernel``.""" + + device: torch.device + dtype: torch.dtype + shape: tuple[int, ...] + + @staticmethod + def of(tensor: torch.Tensor) -> "TensorSpec": + """Describe *tensor*.""" + return TensorSpec(tensor.device, tensor.dtype, tuple(tensor.shape)) + + class _Builtin: """The type of :data:`BUILTIN`. One instance, compared by identity.""" diff --git a/src/tileops/kernels/attention/call_spec.py b/src/tileops/kernels/attention/call_spec.py index d49eec102..ec0d61e1e 100644 --- a/src/tileops/kernels/attention/call_spec.py +++ b/src/tileops/kernels/attention/call_spec.py @@ -35,6 +35,15 @@ WS_ARCH = 90 +# Tile heights the warp-specialized paged decode kernel can pick from. A tile +# divides the page size, so one tile never straddles two pages, and it splits +# evenly across the four consumer warps. +_WS_DECODE_TILES = (16, 32, 64, 128) +# Head dims that map onto one warp: the score reduction is a shuffle chain over +# 32 lanes, so a lane owns ``dim / 32`` elements of the head vector. +_WS_DECODE_LANES = 32 + + def fp8_dtype() -> Optional[torch.dtype]: """Return ``torch.float8_e4m3fn`` when the torch build carries it.""" return getattr(torch, "float8_e4m3fn", None) @@ -81,15 +90,6 @@ def uses_sliding_window(call: AttentionCall) -> bool: return call.window_size_left != -1 or call.window_size_right != -1 -# Tile heights the warp-specialized paged decode kernel can pick from. A tile -# divides the page size, so one tile never straddles two pages, and it splits -# evenly across the four consumer warps. -_WS_DECODE_TILES = (16, 32, 64, 128) -# Head dims that map onto one warp: the score reduction is a shuffle chain over -# 32 lanes, so a lane owns ``dim / 32`` elements of the head vector. -_WS_DECODE_LANES = 32 - - def paged_decode_ws_region(call: AttentionCall) -> bool: """The paged-decode region the warp-specialized MHA kernel serves. diff --git a/src/tileops/kernels/attention/gqa_dense.py b/src/tileops/kernels/attention/gqa_dense.py index 3ac31a980..7c1359ade 100644 --- a/src/tileops/kernels/attention/gqa_dense.py +++ b/src/tileops/kernels/attention/gqa_dense.py @@ -41,6 +41,31 @@ ] +# Causal warp-specialized Dense attention. +BLOCK_M = 128 +BLOCK_N = 128 +NSK = 2 +NSV = 2 +THREADS = 384 +NMMA = 256 +_pc = { + tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, + tilelang.PassConfigKey.TL_DISABLE_THREAD_STORAGE_SYNC: True, +} +_cf = [ + "-O3", + "--use_fast_math", + "-Wno-deprecated-declarations", + "-U__CUDA_NO_HALF_OPERATORS__", + "-U__CUDA_NO_HALF_CONVERSIONS__", + "-U__CUDA_NO_HALF2_OPERATORS__", + "-U__CUDA_NO_BFLOAT16_CONVERSIONS__", + "--expt-relaxed-constexpr", + "--expt-extended-lambda", + "-DNDEBUG", +] + + @functools.lru_cache(maxsize=32) @tilelang.jit(out_idx=[4, 5], pass_configs=_PASS_CONFIGS, compile_flags=_COMPILE_FLAGS) def _gqa_dense_rope_qk_kernel( @@ -201,32 +226,6 @@ def make_dense_qk_rope_preprocessor( ) -# Causal warp-specialized Dense attention. -BLOCK_M = 128 -BLOCK_N = 128 -NSK = 2 -NSV = 2 -THREADS = 384 -NMMA = 256 - -_pc = { - tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True, - tilelang.PassConfigKey.TL_DISABLE_THREAD_STORAGE_SYNC: True, -} -_cf = [ - "-O3", - "--use_fast_math", - "-Wno-deprecated-declarations", - "-U__CUDA_NO_HALF_OPERATORS__", - "-U__CUDA_NO_HALF_CONVERSIONS__", - "-U__CUDA_NO_HALF2_OPERATORS__", - "-U__CUDA_NO_BFLOAT16_CONVERSIONS__", - "--expt-relaxed-constexpr", - "--expt-extended-lambda", - "-DNDEBUG", -] - - @functools.lru_cache(maxsize=32) @tilelang.jit(out_idx=[3], pass_configs=_pc, compile_flags=_cf) def _gqa_dense_ws_kernel( diff --git a/src/tileops/kernels/attention/gqa_fwd.py b/src/tileops/kernels/attention/gqa_fwd.py index 42895b501..3be46c23a 100644 --- a/src/tileops/kernels/attention/gqa_fwd.py +++ b/src/tileops/kernels/attention/gqa_fwd.py @@ -25,14 +25,6 @@ ] -def _tile_stage_thread_configs() -> list[dict]: - """The default GQA search space: block_m x block_n x num_stages x threads.""" - return [ - {"block_m": bm, "block_n": bn, "num_stages": ns, "threads": th} - for bm, bn, ns, th in itertools.product((32, 64, 128), (32, 64, 128), (1, 2, 3), (128, 256)) - ] - - _FAST_COMPILE_FLAGS = [ "-O3", "--use_fast_math", @@ -47,6 +39,14 @@ def _tile_stage_thread_configs() -> list[dict]: ] +def _tile_stage_thread_configs() -> list[dict]: + """The default GQA search space: block_m x block_n x num_stages x threads.""" + return [ + {"block_m": bm, "block_n": bn, "num_stages": ns, "threads": th} + for bm, bn, ns, th in itertools.product((32, 64, 128), (32, 64, 128), (1, 2, 3), (128, 256)) + ] + + def _make_apply_softcap_no_mask_guard(score_scale, softcap, accum_dtype, block_rows, block_cols): @T.macro def apply_softcap(acc_s): diff --git a/src/tileops/kernels/elementwise/_broadcast.py b/src/tileops/kernels/elementwise/_broadcast.py index 23025adf9..91684a4a1 100644 --- a/src/tileops/kernels/elementwise/_broadcast.py +++ b/src/tileops/kernels/elementwise/_broadcast.py @@ -4,6 +4,9 @@ import torch +# CUDA caps grid.y at 65535, and one grid axis carries the rows. +_CUDA_MAX_GRID_Y = 65535 + def _flat(t): """The flat view every PrimFunc here takes.""" @@ -102,10 +105,6 @@ def _is_contiguous_same_shape(coalesced_shape, a_strides, b_strides): ) -# CUDA caps grid.y at 65535, and one grid axis carries the rows. -_CUDA_MAX_GRID_Y = 65535 - - def row_broadcast_split(coalesced_shape, a_strides, b_strides): """``(rows, inner)`` when the innermost coalesced dim reads at stride 0 or 1. diff --git a/src/tileops/kernels/elementwise/_dtype.py b/src/tileops/kernels/elementwise/_dtype.py index ca83efb32..b4370a385 100644 --- a/src/tileops/kernels/elementwise/_dtype.py +++ b/src/tileops/kernels/elementwise/_dtype.py @@ -9,11 +9,6 @@ BOOL_STORAGE_DTYPE = "int8" -def log_for_output_precision(value, wide): - """Return ``log(wide)`` computed to the precision *value*'s dtype can keep.""" - return T.log(wide) if value.dtype == "float32" else T.__log(wide) - - _BITWISE_DTYPES = ( torch.bool, torch.uint8, @@ -22,33 +17,28 @@ def log_for_output_precision(value, wide): torch.int32, torch.int64, ) - - # The dtypes every elementwise kernel refuses. _FP8_DTYPES = ( torch.float8_e4m3fn, torch.float8_e5m2, ) - - _FLOAT_DTYPES = ( torch.float16, torch.bfloat16, torch.float32, ) - - _LOGICAL_DTYPES = _BITWISE_DTYPES + _FLOAT_DTYPES - - _BINARY_FULL_DTYPES = _BITWISE_DTYPES + ( torch.float16, torch.bfloat16, torch.float32, ) +_BINARY_NO_BOOL_DTYPES = tuple(dt for dt in _BINARY_FULL_DTYPES if dt is not torch.bool) -_BINARY_NO_BOOL_DTYPES = tuple(dt for dt in _BINARY_FULL_DTYPES if dt is not torch.bool) +def log_for_output_precision(value, wide): + """Return ``log(wide)`` computed to the precision *value*'s dtype can keep.""" + return T.log(wide) if value.dtype == "float32" else T.__log(wide) def _torch_dtype_nbytes(dtype: torch.dtype) -> int: diff --git a/src/tileops/kernels/fft.py b/src/tileops/kernels/fft.py index 59321753b..03ba6da5e 100644 --- a/src/tileops/kernels/fft.py +++ b/src/tileops/kernels/fft.py @@ -20,6 +20,89 @@ __all__ = ["FFTC2CCall", "FFTC2CDecomposedKernel", "FFTC2COneCTAKernel"] +# n -> the radices of the one-CTA plan's passes. +_RADIX_PLAN = { + 2: (2,), + 4: (4,), + 8: (8,), + 16: (16,), + 32: (32,), + 64: (8, 8), + 128: (8, 8, 2), + 256: (16, 16), + 512: (8, 8, 8), + 1024: (16, 16, 4), + 2048: (16, 16, 8), + 4096: (16, 16, 16), + 8192: (16, 16, 16, 2), + 16384: (16, 16, 16, 4), +} +# (n, dtype) -> four-step factors, outermost first, one kernel each. +_FOUR_STEP_PLAN = { + (1 << 14, "complex128"): (256, 64), + (1 << 15, "complex64"): (256, 128), + (1 << 15, "complex128"): (256, 128), + (1 << 16, "complex64"): (256, 256), + (1 << 16, "complex128"): (256, 256), + (1 << 17, "complex64"): (128, 1024), + (1 << 17, "complex128"): (128, 1024), + (1 << 18, "complex64"): (256, 1024), + (1 << 18, "complex128"): (256, 1024), + (1 << 19, "complex64"): (256, 2048), + (1 << 19, "complex128"): (256, 2048), + (1 << 20, "complex64"): (1024, 1024), + (1 << 20, "complex128"): (1024, 1024), + (1 << 21, "complex64"): (1024, 2048), + (1 << 21, "complex128"): (1024, 2048), + (1 << 22, "complex64"): (2048, 2048), + (1 << 22, "complex128"): (2048, 2048), + (1 << 23, "complex64"): (2048, 4096), + (1 << 23, "complex128"): (2048, 4096), + (1 << 24, "complex64"): (4096, 4096), + (1 << 24, "complex128"): (4096, 4096), + (1 << 25, "complex64"): (256, 512, 256), + (1 << 25, "complex128"): (256, 512, 256), + (1 << 26, "complex64"): (512, 512, 256), + (1 << 26, "complex128"): (512, 512, 256), + (1 << 27, "complex64"): (256, 512, 1024), + (1 << 27, "complex128"): (256, 512, 1024), + (1 << 28, "complex64"): (512, 512, 1024), + (1 << 28, "complex128"): (512, 512, 1024), +} +# (n, dtype) -> transforms per CTA for each kernel of the plan. +_FOUR_STEP_TILE = { + (1 << 14, "complex128"): (4, 4), + (1 << 15, "complex64"): (16, 4), + (1 << 15, "complex128"): (8, 4), + (1 << 16, "complex64"): (16, 8), + (1 << 16, "complex128"): (8, 4), + (1 << 17, "complex64"): (16, 4), + (1 << 17, "complex128"): (16, 4), + (1 << 18, "complex64"): (16, 4), + (1 << 18, "complex128"): (8, 4), + (1 << 19, "complex64"): (32, 4), + (1 << 19, "complex128"): (8, 2), + (1 << 20, "complex64"): (8, 4), + (1 << 20, "complex128"): (4, 4), + (1 << 21, "complex64"): (8, 4), + (1 << 21, "complex128"): (4, 2), + (1 << 22, "complex64"): (8, 4), + (1 << 22, "complex128"): (4, 4), + (1 << 23, "complex64"): (8, 4), + (1 << 23, "complex128"): (4, 2), + (1 << 24, "complex64"): (4, 4), + (1 << 24, "complex128"): (2, 2), + (1 << 25, "complex64"): (32, 16, 8), + (1 << 25, "complex128"): (8, 4, 8), + (1 << 26, "complex64"): (16, 16, 8), + (1 << 26, "complex128"): (4, 4, 8), + (1 << 27, "complex64"): (32, 16, 8), + (1 << 27, "complex128"): (32, 4, 4), + (1 << 28, "complex64"): (16, 16, 8), + (1 << 28, "complex128"): (4, 4, 4), +} + + @dataclasses.dataclass(frozen=True) class FFTC2CCall(CallSpec): """What selects the C2C kernel; the batch is symbolic in every kernel, so it is absent.""" @@ -1317,91 +1400,6 @@ def main( return _func -# n -> the radices of the one-CTA plan's passes. -_RADIX_PLAN = { - 2: (2,), - 4: (4,), - 8: (8,), - 16: (16,), - 32: (32,), - 64: (8, 8), - 128: (8, 8, 2), - 256: (16, 16), - 512: (8, 8, 8), - 1024: (16, 16, 4), - 2048: (16, 16, 8), - 4096: (16, 16, 16), - 8192: (16, 16, 16, 2), - 16384: (16, 16, 16, 4), -} - -# (n, dtype) -> four-step factors, outermost first, one kernel each. -_FOUR_STEP_PLAN = { - (1 << 14, "complex128"): (256, 64), - (1 << 15, "complex64"): (256, 128), - (1 << 15, "complex128"): (256, 128), - (1 << 16, "complex64"): (256, 256), - (1 << 16, "complex128"): (256, 256), - (1 << 17, "complex64"): (128, 1024), - (1 << 17, "complex128"): (128, 1024), - (1 << 18, "complex64"): (256, 1024), - (1 << 18, "complex128"): (256, 1024), - (1 << 19, "complex64"): (256, 2048), - (1 << 19, "complex128"): (256, 2048), - (1 << 20, "complex64"): (1024, 1024), - (1 << 20, "complex128"): (1024, 1024), - (1 << 21, "complex64"): (1024, 2048), - (1 << 21, "complex128"): (1024, 2048), - (1 << 22, "complex64"): (2048, 2048), - (1 << 22, "complex128"): (2048, 2048), - (1 << 23, "complex64"): (2048, 4096), - (1 << 23, "complex128"): (2048, 4096), - (1 << 24, "complex64"): (4096, 4096), - (1 << 24, "complex128"): (4096, 4096), - (1 << 25, "complex64"): (256, 512, 256), - (1 << 25, "complex128"): (256, 512, 256), - (1 << 26, "complex64"): (512, 512, 256), - (1 << 26, "complex128"): (512, 512, 256), - (1 << 27, "complex64"): (256, 512, 1024), - (1 << 27, "complex128"): (256, 512, 1024), - (1 << 28, "complex64"): (512, 512, 1024), - (1 << 28, "complex128"): (512, 512, 1024), -} - -# (n, dtype) -> transforms per CTA for each kernel of the plan. -_FOUR_STEP_TILE = { - (1 << 14, "complex128"): (4, 4), - (1 << 15, "complex64"): (16, 4), - (1 << 15, "complex128"): (8, 4), - (1 << 16, "complex64"): (16, 8), - (1 << 16, "complex128"): (8, 4), - (1 << 17, "complex64"): (16, 4), - (1 << 17, "complex128"): (16, 4), - (1 << 18, "complex64"): (16, 4), - (1 << 18, "complex128"): (8, 4), - (1 << 19, "complex64"): (32, 4), - (1 << 19, "complex128"): (8, 2), - (1 << 20, "complex64"): (8, 4), - (1 << 20, "complex128"): (4, 4), - (1 << 21, "complex64"): (8, 4), - (1 << 21, "complex128"): (4, 2), - (1 << 22, "complex64"): (8, 4), - (1 << 22, "complex128"): (4, 4), - (1 << 23, "complex64"): (8, 4), - (1 << 23, "complex128"): (4, 2), - (1 << 24, "complex64"): (4, 4), - (1 << 24, "complex128"): (2, 2), - (1 << 25, "complex64"): (32, 16, 8), - (1 << 25, "complex128"): (8, 4, 8), - (1 << 26, "complex64"): (16, 16, 8), - (1 << 26, "complex128"): (4, 4, 8), - (1 << 27, "complex64"): (32, 16, 8), - (1 << 27, "complex128"): (32, 4, 4), - (1 << 28, "complex64"): (16, 16, 8), - (1 << 28, "complex128"): (4, 4, 4), -} - - def _plan_table() -> Dict[tuple, FFTPlan]: """One record per served (length, dtype); every one-CTA builder takes (row, grp).""" records = {} diff --git a/src/tileops/kernels/gemm/dense.py b/src/tileops/kernels/gemm/dense.py index d138b53cd..9e53f6411 100644 --- a/src/tileops/kernels/gemm/dense.py +++ b/src/tileops/kernels/gemm/dense.py @@ -42,6 +42,10 @@ _FP8_WS_BLOCK_K = 128 +_TILE_K = 8 +_SMEM_CAP = 224 * 1024 + + def _tma_misalignment( m: int, n: int, k: int, dtype: torch.dtype, trans_a: bool, trans_b: bool ) -> Optional[str]: @@ -2904,10 +2908,6 @@ def _gemm_small_batch_main( return _gemm_small_batch_func -_TILE_K = 8 -_SMEM_CAP = 224 * 1024 - - def _bandwidth_autotune_grid(rts: tuple, bns: tuple, nss: tuple) -> list[dict]: """Config grid for the bandwidth-mode kernels, guarded by thread and SMEM caps.""" return [ diff --git a/src/tileops/kernels/gemm/heuristics.py b/src/tileops/kernels/gemm/heuristics.py index c9c92cbe3..56e0e112a 100644 --- a/src/tileops/kernels/gemm/heuristics.py +++ b/src/tileops/kernels/gemm/heuristics.py @@ -67,6 +67,11 @@ _NS_CAP = {"basic": 4, "splitk": 4, "coop2": 4, "coop2_splitk": 4} +#: The row count the small-M split-K band was fitted at, by calibration key. A board +#: without an entry has no band. +_SMALL_M_SPLITK_M = {"h200": 32} + + @dataclass(frozen=True) class _Calibration: """The scorer's ranking constants for one board. @@ -483,11 +488,6 @@ def small_batch_config(n: int, k: int, sm_count: int) -> dict: return cfg -#: The row count the small-M split-K band was fitted at, by calibration key. A board -#: without an entry has no band. -_SMALL_M_SPLITK_M = {"h200": 32} - - def small_m_splitk_config( m: int, n: int, k: int, sm_count: int, device_name: str ) -> Optional[dict]: diff --git a/src/tileops/kernels/gemm/w4a16.py b/src/tileops/kernels/gemm/w4a16.py index d54f476bf..bc95f7a00 100644 --- a/src/tileops/kernels/gemm/w4a16.py +++ b/src/tileops/kernels/gemm/w4a16.py @@ -24,6 +24,10 @@ ) +# What identifies a tile shape, as opposed to how its K loop is sliced. +_TILE_KEYS = ("block_m", "block_n", "block_k", "num_stages", "threads") + + @dataclass(frozen=True) class _Layout: """Constants fixed by the packed-weight ABI, not tuning parameters.""" @@ -275,10 +279,6 @@ class _TileBuffers(NamedTuple): out_shared: Any -# What identifies a tile shape, as opposed to how its K loop is sliced. -_TILE_KEYS = ("block_m", "block_n", "block_k", "num_stages", "threads") - - def _select_config(m: int, n: int, k: int, group_size: int, sms: int) -> dict: """Choose a tile shape, then its lowest-cost whole-K, split-K, or stream-K variant.""" legal = list(_legal_configs(m, n, k, group_size, sms)) diff --git a/src/tileops/kernels/grouped_gemm/heuristics.py b/src/tileops/kernels/grouped_gemm/heuristics.py index f4a11d678..cc44ae277 100644 --- a/src/tileops/kernels/grouped_gemm/heuristics.py +++ b/src/tileops/kernels/grouped_gemm/heuristics.py @@ -25,6 +25,12 @@ ] +# Gated activations the epilogue can fuse: B stacks gate and up along N; a tile's B +# half-loads block_n / 2 gate columns and the matching up columns, one accumulator +# holds both, and the epilogue stores act(gate) * up, so C has N / 2 columns. +ACTIVATIONS = ("none", "silu_and_mul", "gelu_and_mul") + + @dataclasses.dataclass(frozen=True) class _HeuristicPolicy: """The constants the selector reads, in three kinds a reader must tell apart. @@ -148,12 +154,6 @@ class GemmType(str, enum.Enum): _FLAT_LIKE_TYPES = (GemmType.DENSE, GemmType.BATCHED, GemmType.K_GROUPED_CONTIGUOUS) -# Gated activations the epilogue can fuse: B stacks gate and up along N; a tile's B -# half-loads block_n / 2 gate columns and the matching up columns, one accumulator -# holds both, and the epilogue stores act(gate) * up, so C has N / 2 columns. -ACTIVATIONS = ("none", "silu_and_mul", "gelu_and_mul") - - class Major(str, enum.Enum): """Which logical dim is contiguous in memory for an operand.""" diff --git a/src/tileops/kernels/norm/fused_add_norm.py b/src/tileops/kernels/norm/fused_add_norm.py index 8c226b9a6..a61ede64e 100644 --- a/src/tileops/kernels/norm/fused_add_norm.py +++ b/src/tileops/kernels/norm/fused_add_norm.py @@ -37,6 +37,10 @@ # Fused Add + LayerNorm kernel +# This kernel serves 16-bit dtypes only, so one 16-byte access moves eight elements. +_VEC = VECTOR_ACCESS_BYTES // 2 + + @functools.lru_cache(maxsize=32) def _fused_add_layer_norm_kernel(M, N, eps, dtype): N_padded = align_up(N, ALIGNMENT) @@ -229,9 +233,6 @@ def forward( # Fused Add + RMSNorm kernel -# This kernel serves 16-bit dtypes only, so one 16-byte access moves eight elements. -_VEC = VECTOR_ACCESS_BYTES // 2 - @functools.lru_cache(maxsize=32) def _fused_add_rms_norm_kernel(M, N, eps, dtype, splits): diff --git a/src/tileops/kernels/pool/common.py b/src/tileops/kernels/pool/common.py index 0aff54713..5d665004f 100644 --- a/src/tileops/kernels/pool/common.py +++ b/src/tileops/kernels/pool/common.py @@ -7,6 +7,10 @@ from tileops.kernels.constants import STATIC_SHARED_BYTES, VECTOR_ACCESS_BYTES from tileops.kernels.kernel_base import Kernel +# Window sums promote to fp32 and cast back at the store: a narrow accumulator loses the +# low bits of a window this wide. +ACCUM_DTYPE = "float" + def dtype_itemsize(dtype: str) -> int: """Bytes one element of *dtype* takes, over the dtypes these kernels accept.""" @@ -103,11 +107,6 @@ def pool_output_dim( return max(out, 0) -# Window sums promote to fp32 and cast back at the store: a narrow accumulator loses the -# low bits of a window this wide. -ACCUM_DTYPE = "float" - - class AvgPoolWindow(NamedTuple): """One average-pooling problem, and the extents and facts that follow from it. diff --git a/src/tileops/kernels/reduction/_primitives.py b/src/tileops/kernels/reduction/_primitives.py index b00980d79..e5d887464 100644 --- a/src/tileops/kernels/reduction/_primitives.py +++ b/src/tileops/kernels/reduction/_primitives.py @@ -75,6 +75,11 @@ FRAGMENT_ELEMS_PER_THREAD: int = 64 +# Largest integer count fp32 carries exactly; a statistic folded through +# fp32 counts or weights is trusted only below it. +FP32_EXACT_INT_LIMIT = 1 << 24 + + def ceildiv_int(x: int, y: int) -> int: """Return ``ceil(x / y)`` for positive integer dimensions.""" return -(-x // y) @@ -856,11 +861,6 @@ class _LeadingAxisReducePolicy: _LEADING_POLICY = _LeadingAxisReducePolicy() -# Largest integer count fp32 carries exactly; a statistic folded through -# fp32 counts or weights is trusted only below it. -FP32_EXACT_INT_LIMIT = 1 << 24 - - def edge_axis_plan( shape: "tuple[int, ...]", k: int, diff --git a/src/tileops/kernels/reduction/call_spec.py b/src/tileops/kernels/reduction/call_spec.py index 809ce9736..836687910 100644 --- a/src/tileops/kernels/reduction/call_spec.py +++ b/src/tileops/kernels/reduction/call_spec.py @@ -15,6 +15,12 @@ ] +# The fused pass runs one block per kept column and has no other parallelism, so +# it takes over only where that alone is enough: the fewest kept columns that fill the +# device, per calibrated board. A board without an entry uses the general implementation. +_EDGE_FUSED_MIN_KEPT = {"h200": 32} + + @dataclasses.dataclass(frozen=True) class LogicalReduceCall(CallSpec): """Semantic and shape facts used to select a logical reduction implementation.""" @@ -37,12 +43,6 @@ def logical_reduce_region(call: LogicalReduceCall) -> bool: return call.op_kind in {"any", "all", "count_nonzero"} -# The fused pass runs one block per kept column and has no other parallelism, so -# it takes over only where that alone is enough: the fewest kept columns that fill the -# device, per calibrated board. A board without an entry uses the general implementation. -_EDGE_FUSED_MIN_KEPT = {"h200": 32} - - def logical_edge_fused_region(call: LogicalReduceCall) -> bool: """The edge-axis logical reduction region a calibrated board serves with the fused pass.""" diff --git a/src/tileops/kernels/reduction/logical_reduce.py b/src/tileops/kernels/reduction/logical_reduce.py index 307307a19..2cee9fa3e 100644 --- a/src/tileops/kernels/reduction/logical_reduce.py +++ b/src/tileops/kernels/reduction/logical_reduce.py @@ -70,6 +70,15 @@ _UNSUPPORTED_STORAGE_DTYPES = _BYTE_REINTERPRETED_DTYPES | _WIDENED_STORAGE_DTYPES +# Elements a lane folds in the edge-fused pass. One block holds a `trail`-wide +# fp32 fragment, so this is what fixes its register footprint per lane rather +# than letting it grow with the row. Eight is the flat optimum at every width +# the manifest asks for. +_FUSED_EDGE_ELEMS_PER_LANE = 8 +_FUSED_EDGE_MIN_THREADS = 64 +_FUSED_EDGE_MAX_THREADS = 1024 + + def storage_dtype_for(dtype: torch.dtype) -> torch.dtype: """The dtype the prim_func declares for an input of *dtype*.""" if dtype in _BYTE_REINTERPRETED_DTYPES: @@ -108,15 +117,6 @@ def _logical_out_dtype(op_kind: str, partial: bool) -> str: return "float32" if partial else "int64" -# Elements a lane folds in the edge-fused pass. One block holds a `trail`-wide -# fp32 fragment, so this is what fixes its register footprint per lane rather -# than letting it grow with the row. Eight is the flat optimum at every width -# the manifest asks for. -_FUSED_EDGE_ELEMS_PER_LANE = 8 -_FUSED_EDGE_MIN_THREADS = 64 -_FUSED_EDGE_MAX_THREADS = 1024 - - def fused_edge_threads(trail: int) -> int: """The thread width the edge-fused pass runs a ``trail``-wide row at.""" lanes = ceildiv_int(trail, _FUSED_EDGE_ELEMS_PER_LANE) diff --git a/src/tileops/kernels/reduction/reduce.py b/src/tileops/kernels/reduction/reduce.py index 719cfb787..516b5200f 100644 --- a/src/tileops/kernels/reduction/reduce.py +++ b/src/tileops/kernels/reduction/reduce.py @@ -45,6 +45,9 @@ } +_LEADING_AXIS_KINDS = frozenset({"sum", "mean", "amax", "amin"}) + + @dataclass(frozen=True) class ProductReducePolicy: """Launch heuristics for product reductions.""" @@ -62,9 +65,6 @@ class ProductReducePolicy: # Simple reduce kernel -_LEADING_AXIS_KINDS = frozenset({"sum", "mean", "amax", "amin"}) - - class ReduceKernel(Kernel): """Unified reduce kernel supporting sum/mean/amin/amax/prod/std/var/var_mean. diff --git a/src/tileops/manifest/dtype_rules.py b/src/tileops/manifest/dtype_rules.py index e3a85948b..108acaa58 100644 --- a/src/tileops/manifest/dtype_rules.py +++ b/src/tileops/manifest/dtype_rules.py @@ -1,28 +1,37 @@ -"""The dtype registry: every dtype name a manifest may write, with its bits per element. +"""The dtype registry: every dtype name a manifest may write, with its bits per element and +its category. This package depends on nothing beyond the standard library and PyYAML, so dtypes are names. """ from __future__ import annotations -__all__ = ["DTYPE_BITS"] +__all__ = ["DTYPE_BITS", "DTYPE_CATEGORY", "FLOAT8_DTYPES"] -DTYPE_BITS: dict[str, int] = { - "bool": 8, - "uint8": 8, - "int8": 8, - "int16": 16, - "int32": 32, - "int64": 64, - "float16": 16, - "bfloat16": 16, - "float32": 32, - "float64": 64, - "complex64": 64, - "complex128": 128, - "float8_e4m3fn": 8, - "float8_e5m2": 8, - "float8_e4m3": 8, - "float8_e5m2fnuz": 8, - "float8_e4m3fnuz": 8, +# name: (bits per element, category) +_DTYPES: dict[str, tuple[int, str]] = { + "bool": (8, "bool"), + "uint8": (8, "int"), + "int8": (8, "int"), + "int16": (16, "int"), + "int32": (32, "int"), + "int64": (64, "int"), + "float16": (16, "float"), + "bfloat16": (16, "float"), + "float32": (32, "float"), + "float64": (64, "float"), + "complex64": (64, "complex"), + "complex128": (128, "complex"), + "float8_e4m3fn": (8, "float"), + "float8_e5m2": (8, "float"), + "float8_e4m3": (8, "float"), + "float8_e5m2fnuz": (8, "float"), + "float8_e4m3fnuz": (8, "float"), } + +DTYPE_BITS: dict[str, int] = {name: bits for name, (bits, _) in _DTYPES.items()} +# 'bool', 'int', 'float' or 'complex'. +DTYPE_CATEGORY: dict[str, str] = {name: kind for name, (_, kind) in _DTYPES.items()} +FLOAT8_DTYPES: frozenset[str] = frozenset( + name for name, (bits, kind) in _DTYPES.items() if bits == 8 and kind == "float" +) diff --git a/src/tileops/manifest/expr.py b/src/tileops/manifest/expr.py index ba7cbf545..8e588f466 100644 --- a/src/tileops/manifest/expr.py +++ b/src/tileops/manifest/expr.py @@ -54,17 +54,8 @@ ] -class SignatureError(ValueError): - """A declaration outside the schema; the message names it.""" - - -class EvaluationError(SignatureError): - """An expression that failed to evaluate; the message names its declaration.""" - - # The value of an expression that reads more than a point fixes. OPEN = object() - _COMPREHENSION_CALLEES = frozenset({"all", "sum", "max", "min"}) _NODES = ( ast.BoolOp, @@ -102,6 +93,15 @@ class EvaluationError(SignatureError): ast.comprehension, ast.keyword, ) +_UNFOLDED = (ast.Constant, ast.Name, ast.List, ast.Starred, ast.Slice, ast.GeneratorExp) + + +class SignatureError(ValueError): + """A declaration outside the schema; the message names it.""" + + +class EvaluationError(SignatureError): + """An expression that failed to evaluate; the message names its declaration.""" # ---------------------------------------------------------------- parsing and the language @@ -644,7 +644,6 @@ def visit_Call(self, node): _NAMESPACE = namespace() -_UNFOLDED = (ast.Constant, ast.Name, ast.List, ast.Starred, ast.Slice, ast.GeneratorExp) def bind(node: ast.expr, point: dict) -> ast.expr: diff --git a/src/tileops/manifest/primitives.py b/src/tileops/manifest/primitives.py index 4db576b90..1c8571026 100644 --- a/src/tileops/manifest/primitives.py +++ b/src/tileops/manifest/primitives.py @@ -10,7 +10,7 @@ import numbers from types import SimpleNamespace -from .dtype_rules import DTYPE_BITS +from .dtype_rules import DTYPE_BITS, DTYPE_CATEGORY # The seed both conftests give the global RNG; every private workload RNG derives from it. WORKLOAD_SEED = 1235 @@ -32,6 +32,23 @@ ] +# The largest finite value of each floating dtype; the lowest is its negation. +_FLOAT_MAX = { + "float16": 65504.0, + "bfloat16": 3.3895313892515355e38, + "float32": 3.4028234663852886e38, + "float64": 1.7976931348623157e308, + "float8_e4m3fn": 448.0, + "float8_e4m3": 240.0, + "float8_e5m2": 57344.0, + "float8_e4m3fnuz": 240.0, + "float8_e5m2fnuz": 57344.0, +} +_COMPLEX_PART = {"complex64": "float32", "complex128": "float64"} +# Floating formats without an infinity. +_NO_INF = frozenset({"float8_e4m3fn", "float8_e4m3fnuz", "float8_e5m2fnuz"}) + + def normalize_axis(axis: int, rank: int) -> int: """At rank 0, `0` and `-1` name the scalar axis; otherwise an axis lies in [-rank, rank).""" lo, hi = (-1, 1) if rank == 0 else (-rank, rank) @@ -163,7 +180,7 @@ def moe_capacity(layout, rows, experts): def promote_int_to_float(dtype): """float32 for an integral dtype, else the dtype itself.""" - return "float32" if dtype in ("uint8", "int8", "int16", "int32", "int64") else dtype + return "float32" if category(dtype) == "int" else dtype def coalesce_dtype(value, dtype): @@ -193,29 +210,10 @@ def repeat(value, count): return [value] * count -# The largest finite value of each floating dtype; the lowest is its negation. -_FLOAT_MAX = { - "float16": 65504.0, - "bfloat16": 3.3895313892515355e38, - "float32": 3.4028234663852886e38, - "float64": 1.7976931348623157e308, - "float8_e4m3fn": 448.0, - "float8_e4m3": 240.0, - "float8_e5m2": 57344.0, - "float8_e4m3fnuz": 240.0, - "float8_e5m2fnuz": 57344.0, -} -_COMPLEX_PART = {"complex64": "float32", "complex128": "float64"} - - def category(x): """`'bool'`, `'int'`, `'float'` or `'complex'`: the category of a number or a dtype name.""" if isinstance(x, str): - if x == "bool": - return "bool" - if x in _COMPLEX_PART: - return "complex" - return "float" if x in _FLOAT_MAX else "int" + return DTYPE_CATEGORY.get(x, "int") if isinstance(x, bool): return "bool" if isinstance(x, numbers.Integral): @@ -223,10 +221,6 @@ def category(x): return "float" if isinstance(x, numbers.Real) else "complex" -# Floating formats without an infinity. -_NO_INF = frozenset({"float8_e4m3fn", "float8_e4m3fnuz", "float8_e5m2fnuz"}) - - def _fits_float(v, dtype): if math.isinf(v): return dtype not in _NO_INF diff --git a/src/tileops/ops/_signature_codegen.py b/src/tileops/ops/_signature_codegen.py index bc762e05c..9a42b16c4 100644 --- a/src/tileops/ops/_signature_codegen.py +++ b/src/tileops/ops/_signature_codegen.py @@ -51,6 +51,9 @@ __all__ = ["CheckError", "SignatureCall", "install", "maybe_install_signature"] +_SCHEMA_TYPES = {int: "SymInt", float: "float", bool: "bool", str: "str"} + + def operator_name(family: str, class_name: str) -> str: """``("norm", "RMSNormFwdOp")`` -> ``"norm_rms_norm_fwd"``; a class whose own name already opens with the family, such as ``MoePrePermuteFwdOp``, names it once.""" @@ -1234,9 +1237,6 @@ def binder(self, cls: type): return scope["_call_boundary"] -_SCHEMA_TYPES = {int: "SymInt", float: "float", bool: "bool", str: "str"} - - def _schema_type(cls: type, parameter: inspect.Parameter) -> str: """The operator-schema type of one execution parameter, from its annotation.""" annotation = parameter.annotation diff --git a/src/tileops/ops/elementwise/_base.py b/src/tileops/ops/elementwise/_base.py index 22a7dce0d..3f6ad5db5 100644 --- a/src/tileops/ops/elementwise/_base.py +++ b/src/tileops/ops/elementwise/_base.py @@ -23,6 +23,15 @@ from ..op_base import Op +_MANIFEST_INT_DTYPES = ( + torch.uint8, + torch.int8, + torch.int16, + torch.int32, + torch.int64, +) +_PREDICATE_FALLBACK_DTYPES = _MANIFEST_INT_DTYPES + (torch.bool,) + class _PerDtypeKernels: """The family's one way to reach a kernel: ``self._kernel(inputs, dtype, *dims)``. @@ -333,23 +342,11 @@ def _build_kernel_instance(self, tune, dtype, impl, a_shape, b_shape): return impl(a_shape, b_shape, dtype, tune=tune, alpha=self.alpha) -_MANIFEST_INT_DTYPES = ( - torch.uint8, - torch.int8, - torch.int16, - torch.int32, - torch.int64, -) - - def _int_identity(input: torch.Tensor) -> torch.Tensor: """The default integer answer: the op leaves such a value unchanged.""" return input.clone() -_PREDICATE_FALLBACK_DTYPES = _MANIFEST_INT_DTYPES + (torch.bool,) - - class _IntFallbackCall: """What ``_IntIdentityUnaryOp`` builds for a dtype the shipped kernels do not serve. diff --git a/src/tileops/ops/elementwise/arithmetic.py b/src/tileops/ops/elementwise/arithmetic.py index 635d9c3ff..5e5a1f53c 100644 --- a/src/tileops/ops/elementwise/arithmetic.py +++ b/src/tileops/ops/elementwise/arithmetic.py @@ -24,6 +24,8 @@ from ..op_base import Op from ._base import BinaryOp, _AlphaScaledBinaryOp, _PerDtypeKernels +_DIV_KEY_BY_ROUNDING_MODE = {None: "div", "trunc": "div_trunc", "floor": "floor_divide"} + class AddFwdOp(_AlphaScaledBinaryOp): """Element-wise addition with broadcast: y = input + alpha * other. @@ -53,9 +55,6 @@ class MulFwdOp(BinaryOp): kernel_types = {"mul": MulFwdKernel} -_DIV_KEY_BY_ROUNDING_MODE = {None: "div", "trunc": "div_trunc", "floor": "floor_divide"} - - class DivFwdOp(BinaryOp): """Element-wise division with broadcast: y = input / other. diff --git a/src/tileops/ops/op_base.py b/src/tileops/ops/op_base.py index 44a40bb59..993d4bcfd 100644 --- a/src/tileops/ops/op_base.py +++ b/src/tileops/ops/op_base.py @@ -38,6 +38,11 @@ _Entry = TypeVar("_Entry") +# Every dispatch key a created op class declares in ``kernel_types``. Constructing an op imports +# it and every sub-op it builds, so every key that can replace something in that op is here. +_DISPATCH_KEYS: set[str] = set() + + class _Unresolved: """The type of :data:`_UNRESOLVED`, so a traceback says what it is.""" @@ -52,10 +57,6 @@ def __repr__(self) -> str: _UNRESOLVED = _Unresolved() -# Every dispatch key a created op class declares in ``kernel_types``. Constructing an op imports -# it and every sub-op it builds, so every key that can replace something in that op is here. -_DISPATCH_KEYS: set[str] = set() - # The calls in progress on this thread, innermost last, each with the checked calls completed # inside it: what a composite's call collects from its sub-ops. _OPEN_CALLS = threading.local() diff --git a/src/tileops/ops/pool.py b/src/tileops/ops/pool.py index 7cd08b7bf..f8e3c56df 100644 --- a/src/tileops/ops/pool.py +++ b/src/tileops/ops/pool.py @@ -45,6 +45,17 @@ ] +# Per-axis name suffixes, indexed by spatial dimensionality. +_POOL_DIM_NAMES: Dict[int, Tuple[str, ...]] = {1: ("l",), 2: ("h", "w"), 3: ("d", "h", "w")} +# Kernel-kwarg suffixes for kernel_size/stride/padding(/dilation). +# Why: the 1d max-pool kernels name their pooling axis `w`, not `l`. +_MAX_POOL_PARAM_SUFFIXES: Dict[int, Tuple[str, ...]] = { + 1: ("w",), + 2: ("h", "w"), + 3: ("d", "h", "w"), +} + + def _per_axis(value: "int | Sequence[int]", ndim: int) -> tuple[int, ...]: """A pooling parameter as one value per spatial axis, as ``per_axis`` reads it.""" return (value,) * ndim if isinstance(value, int) else tuple(value) @@ -292,17 +303,6 @@ def _check_ragged( raise ValueError("indices must name each chunk offsets implies exactly once") -# Per-axis name suffixes, indexed by spatial dimensionality. -_POOL_DIM_NAMES: Dict[int, Tuple[str, ...]] = {1: ("l",), 2: ("h", "w"), 3: ("d", "h", "w")} -# Kernel-kwarg suffixes for kernel_size/stride/padding(/dilation). -# Why: the 1d max-pool kernels name their pooling axis `w`, not `l`. -_MAX_POOL_PARAM_SUFFIXES: Dict[int, Tuple[str, ...]] = { - 1: ("w",), - 2: ("h", "w"), - 3: ("d", "h", "w"), -} - - class _AvgPoolFwdOpBase(Op): """Generic average-pooling forward, parametrized by class-attribute ``ndim``. diff --git a/src/tileops/perf/formulas.py b/src/tileops/perf/formulas.py index b9980d600..c807c0a91 100644 --- a/src/tileops/perf/formulas.py +++ b/src/tileops/perf/formulas.py @@ -13,6 +13,8 @@ from math import prod from typing import TYPE_CHECKING +from tileops.manifest.dtype_rules import FLOAT8_DTYPES + if TYPE_CHECKING: from tileops.manifest.workload import CallView @@ -65,6 +67,17 @@ ] +# Per gated element: the activation of the gate (silu: 5; the erf gelu: 5) and the multiply +# by the up projection. +_GATED_ACTIVATION = 6 +# Per score: the scale, the running max, the subtraction, the exp and the sum of a softmax. +_SOFTMAX_PER_SCORE = 5 +# Per score: the divide, tanh and multiply of a logit softcap. +_SOFTCAP_PER_SCORE = 3 +# Per head score: the relu, the weight multiply and the add into the sum over heads. +_INDEXER_EPILOGUE_PER_SCORE = 3 + + def _distribute_total(total: int, batch: int, max_len: int) -> list[int]: lengths = [0] * batch remaining = total @@ -97,11 +110,6 @@ def _expert_weight_bytes(call) -> int: return (call.bytes("w_gate_up") + call.bytes("w_down")) // experts -# Per gated element: the activation of the gate (silu: 5; the erf gelu: 5) and the multiply -# by the up projection. -_GATED_ACTIVATION = 6 - - def _routing_flops(call) -> int: """Routing as FusedTopKFwdOp prices it, per token: scoring (sigmoid 4 per logit; softmax 3, plus the row sum and a divide per kept weight unless renormalizing), top_k @@ -406,12 +414,6 @@ def visible_scores(q_len: int, kv_len: int, is_causal: bool, left: int, right: i return visible_score_rows(q_len, kv_len, is_causal, left, right)[0] -# Per score: the scale, the running max, the subtraction, the exp and the sum of a softmax. -_SOFTMAX_PER_SCORE = 5 -# Per score: the divide, tanh and multiply of a logit softcap. -_SOFTCAP_PER_SCORE = 3 - - def attention_flops( heads: int, scores: int, rows: int, qk_dim: int, v_dim: int, softcap: bool = False ) -> int: @@ -731,10 +733,6 @@ def lightning_indexer_scored_keys(call: "CallView") -> int: ) -# Per head score: the relu, the weight multiply and the add into the sum over heads. -_INDEXER_EPILOGUE_PER_SCORE = 3 - - def fp8_lightning_indexer_roofline(call: "CallView") -> tuple[int, int]: """Lightning indexer: per query head and windowed key, a D-long contraction and the relu, weight and head-sum epilogue; each tensor moves once, ``logits`` written whole.""" @@ -773,9 +771,6 @@ def gqa_prefill_paged_cache_rows(call: "CallView") -> int: # ---------------------------------------------------------------- paged caches and MLA -_FP8 = "float8_e4m3fn" - - def _elem_bytes(call: "CallView", name: str) -> int: """Bytes of one element of tensor *name*.""" return call.bytes(name) // max(1, prod(call.tensors[name][0])) @@ -805,7 +800,7 @@ def mla_paged_fwd_roofline(call: "CallView") -> tuple[int, int]: pairs = [visible_score_rows(ix["S_q"], c, ix["is_causal"], -1, -1) for c in lengths] scores, rows = sum(p[0] for p in pairs), sum(p[1] for p in pairs) flops = attention_flops(ix["H"], scores, rows, ix["DK"], ix["kv_lora_rank"]) - if call.tensors["kv_cache"][1] == _FP8: + if call.tensors["kv_cache"][1] in FLOAT8_DTYPES: flops += ix["H"] * rows moved = _derived_bytes(call) - call.bytes("kv_cache") - call.bytes("block_table") moved += _paged_cache_read_bytes(call, "kv_cache", "block_table", lengths) @@ -870,7 +865,7 @@ def paged_kv_cache_write_roofline(call: "CallView") -> tuple[int, int]: into both caches; an FP8 cache scales and saturates each value (2 per value).""" shape = call.tensors["k"][0] tokens, row = _written_slots(call), prod(shape[1:]) - flops = 4 * tokens * row if call.tensors["k_pages"][1] == _FP8 else 0 + flops = 4 * tokens * row if call.tensors["k_pages"][1] in FLOAT8_DTYPES else 0 moved = 2 * tokens * row * (_elem_bytes(call, "k") + _elem_bytes(call, "k_pages")) moved += call.bytes("slot_mapping") if tokens * row: @@ -886,7 +881,7 @@ def mla_kv_cache_write_roofline(call: "CallView") -> tuple[int, int]: width, pe = ix["DC"] + ix["PE"], ix["PE"] slots = call.values("slot_mapping") tokens = _written_slots(call) - flops = 2 * tokens * width if call.tensors["kv_cache"][1] == _FP8 else 0 + flops = 2 * tokens * width if call.tensors["kv_cache"][1] in FLOAT8_DTYPES else 0 moved = tokens * width * (_elem_bytes(call, "kv_c") + _elem_bytes(call, "kv_cache")) moved += call.bytes("slot_mapping") moved += call.bytes("scale") if tokens * width and call.present("scale") else 0 @@ -904,7 +899,7 @@ def paged_kv_cache_gather_roofline(call: "CallView") -> tuple[int, int]: lengths = _segments(call, "cu_seq_lens") starts = call.values("seq_starts") if call.present("seq_starts") else [0] * len(lengths) ends = [s + n for s, n in zip(starts, lengths, strict=True)] - fp8 = call.tensors["cache"][1] == _FP8 + fp8 = call.tensors["cache"][1] in FLOAT8_DTYPES flops = prod(call.tensors["dst"][0]) if fp8 else 0 moved = _derived_bytes(call) - call.bytes("cache") - call.bytes("block_table") moved += _paged_cache_read_bytes(call, "cache", "block_table", ends, starts) diff --git a/src/tileops/perf/profile.py b/src/tileops/perf/profile.py index 8c3e7766f..b8d919fd7 100644 --- a/src/tileops/perf/profile.py +++ b/src/tileops/perf/profile.py @@ -20,6 +20,18 @@ ) +# Tensor-core dtype keys, by the dtype the contraction consumes. fp32 maps to +# tf32 because that is the unit an fp32 contraction runs on when tensor cores +# serve it. Encode side of the roof-key format; ``resolve_roof`` is the decode. +_TENSOR_CORE_DTYPE_KEYS = { + "float16": "fp16", + "bfloat16": "bf16", + "float32": "tf32", + "float8_e4m3fn": "fp8", + "float8_e5m2": "fp8", +} + + def get_profile_path(gpu_name: str) -> Path: """Return the path to a GPU profile YAML. @@ -71,18 +83,6 @@ def _inject_effective(profile): section["effective"] = section["theoretical"] * section["calibration"] -# Tensor-core dtype keys, by the dtype the contraction consumes. fp32 maps to -# tf32 because that is the unit an fp32 contraction runs on when tensor cores -# serve it. Encode side of the roof-key format; ``resolve_roof`` is the decode. -_TENSOR_CORE_DTYPE_KEYS = { - "float16": "fp16", - "bfloat16": "bf16", - "float32": "tf32", - "float8_e4m3fn": "fp8", - "float8_e5m2": "fp8", -} - - def tensor_core_roof(dtype) -> str: """Tensor-core roof key for a contraction computing at *dtype*. diff --git a/src/tileops/trace/record.py b/src/tileops/trace/record.py index 77949364f..34f9ca8e2 100644 --- a/src/tileops/trace/record.py +++ b/src/tileops/trace/record.py @@ -21,40 +21,36 @@ ] -class EventKind(IntEnum): - """Event kind packed into ``w1`` bits 24..27.""" - - RANGE_BEGIN = 0 - RANGE_END = 1 - INSTANT = 2 - # Reserved: ``trace.dag`` is now a build-time declaration (no runtime record), - # so no DAG record is ever emitted. Kept to keep the enum value stable. - DAG = 3 - - # Per-slot record capacity (config default; callers may override). MAX_EVENTS_DEFAULT = 768 - # Field widths and bit offsets within w1. _EVENT_ID_BITS = 24 _KIND_BITS = 4 _LANE_BITS = 4 _PAYLOAD_BITS = 32 - _EVENT_ID_SHIFT = 0 _KIND_SHIFT = 24 _LANE_SHIFT = 28 _PAYLOAD_SHIFT = 32 - _EVENT_ID_MASK = (1 << _EVENT_ID_BITS) - 1 _KIND_MASK = (1 << _KIND_BITS) - 1 _LANE_MASK = (1 << _LANE_BITS) - 1 _PAYLOAD_MASK = (1 << _PAYLOAD_BITS) - 1 - # Max distinct lanes that fit the 4-bit lane field. MAX_LANES = 1 << _LANE_BITS +class EventKind(IntEnum): + """Event kind packed into ``w1`` bits 24..27.""" + + RANGE_BEGIN = 0 + RANGE_END = 1 + INSTANT = 2 + # Reserved: ``trace.dag`` is now a build-time declaration (no runtime record), + # so no DAG record is ever emitted. Kept to keep the enum value stable. + DAG = 3 + + def pack_w1(event_id: int, kind: int, lane: int, payload: int) -> int: """Pack the four ``w1`` fields into a single unsigned 64-bit word. diff --git a/src/tileops/trace/ui.py b/src/tileops/trace/ui.py index d04cf3fd5..6a2c6b688 100644 --- a/src/tileops/trace/ui.py +++ b/src/tileops/trace/ui.py @@ -74,6 +74,59 @@ _INK = "#191a16" +# Plotly config: horizontal-only zoom + pan. With yaxis.fixedrange set, scrollZoom +# stretches only x and pan moves only x; the vertical / box / autoscale buttons are +# stripped so the y axis can never be rescaled. +_CONFIG = { + "scrollZoom": True, + "displaylogo": False, + "responsive": True, + "modeBarButtonsToRemove": [ + "zoom2d", + "select2d", + "lasso2d", + "zoomIn2d", + "zoomOut2d", + "autoScale2d", + ], +} +_PLOTLY_CDN = "https://cdn.plot.ly/plotly-2.35.2.min.js" +_HTML_TEMPLATE = """ + +{title} + + +
{tab_buttons}
+
+""" + + def _lane_label(gid: int, lane: int, group_id_to_name: dict, lane_id_to_name: dict) -> str: """Build a lane's y-axis label ``" / "``. @@ -348,61 +401,6 @@ def _figure_for_cta( return {"data": data, "layout": layout} -# Plotly config: horizontal-only zoom + pan. With yaxis.fixedrange set, scrollZoom -# stretches only x and pan moves only x; the vertical / box / autoscale buttons are -# stripped so the y axis can never be rescaled. -_CONFIG = { - "scrollZoom": True, - "displaylogo": False, - "responsive": True, - "modeBarButtonsToRemove": [ - "zoom2d", - "select2d", - "lasso2d", - "zoomIn2d", - "zoomOut2d", - "autoScale2d", - ], -} - -_PLOTLY_CDN = "https://cdn.plot.ly/plotly-2.35.2.min.js" - -_HTML_TEMPLATE = """ - -{title} - - -
{tab_buttons}
-
-""" - - def export_timeline_html( events: list, path: str, diff --git a/src/tileops/utils/utils.py b/src/tileops/utils/utils.py index f4d585ef9..39d25c1f6 100644 --- a/src/tileops/utils/utils.py +++ b/src/tileops/utils/utils.py @@ -21,6 +21,17 @@ # `get_device_name` string scan would run on every forward. +# Spin cycles queued before a device_busy_of measurement: tens of milliseconds +# on any supported clock, ample to enqueue every timed call first. +_BUSY_TIMING_SPIN_CYCLES = 50_000_000 + + +# Calibrated boards: the key selection tables use -> the name fragment CUDA reports. +# All SKUs of a board share its key. GPU profiles match the full name instead +# (:func:`tileops.perf.find_profile`): a speed-of-light reading is not shared. +_CALIBRATION_BOARDS = {"h200": "H200"} + + @functools.lru_cache(maxsize=16) def _device_name(index: int) -> str: return torch.cuda.get_device_name(index) @@ -32,12 +43,6 @@ def _sm_version(index: int) -> int: return major * 10 + minor -# Calibrated boards: the key selection tables use -> the name fragment CUDA reports. -# All SKUs of a board share its key. GPU profiles match the full name instead -# (:func:`tileops.perf.find_profile`): a speed-of-light reading is not shared. -_CALIBRATION_BOARDS = {"h200": "H200"} - - def calibration_key(device_name: str) -> "str | None": """The key of the calibrated board *device_name* belongs to, or ``None``. @@ -102,11 +107,6 @@ def forget_device_properties() -> None: _device_facts.cache_clear() -# Spin cycles queued before a device_busy_of measurement: tens of milliseconds -# on any supported clock, ample to enqueue every timed call first. -_BUSY_TIMING_SPIN_CYCLES = 50_000_000 - - def device_busy_of(call, device: "torch.device", warmup: int = 5, rep: int = 20) -> float: """Mean device time of *call* in milliseconds with host gaps excluded. diff --git a/workloads/elementwise.py b/workloads/elementwise.py index 1582efe43..b284987ab 100644 --- a/workloads/elementwise.py +++ b/workloads/elementwise.py @@ -141,7 +141,7 @@ def gen_inputs(self) -> tuple[torch.Tensor]: if self.dtype == torch.uint8: x = torch.randint(0, 8, (self.n_total,), device=run_device(), dtype=self.dtype) - elif self.dtype in (torch.int8, torch.int16, torch.int32, torch.int64): + elif not (self.dtype.is_floating_point or self.dtype.is_complex) and self.dtype.is_signed: x = torch.randint(-4, 4, (self.n_total,), device=run_device(), dtype=self.dtype) else: x = torch.randn(self.n_total, device=run_device(), dtype=self.dtype) diff --git a/workloads/gemm.py b/workloads/gemm.py index b774d7220..adff6e46d 100644 --- a/workloads/gemm.py +++ b/workloads/gemm.py @@ -11,6 +11,9 @@ W4A16_GROUP_SIZE = 128 +_FP8_INIT_SCALE: float = 0.25 + + class GemmWorkload(WorkloadBase): def __init__( self, @@ -334,9 +337,6 @@ def ref_program(self, a: torch.Tensor, b: torch.Tensor) -> torch.Tensor: return torch.bmm(a, b) -_FP8_INIT_SCALE: float = 0.25 - - class BmmFp8Workload(WorkloadBase): """Workload for batched FP8 GEMM. From 565c1bfe0877ebe26822002d8c20b680687210cb Mon Sep 17 00:00:00 2001 From: lcy-seso Date: Sun, 27 Sep 2026 23:46:49 +0800 Subject: [PATCH 4/5] [Test][Manifest] Check spec-only entries against their references and recount their rooflines test_spec_reference.py runs each new entry's reference, the torch expression its issue names or FLA's recurrence for Kimi Delta Attention, on every workload row with data on meta and metadata on the CPU. The reference's outputs must have the signature's names, shapes and dtypes, it must write exactly the inputs the call's effects mark written, and one call the signature rejects must be rejected by the reference too. The roofline oracle now ranges over spec-only entries as well: the binder builds a spec-only op from its signature alone, since a recount needs no implementation. The entries whose traffic follows their metadata get hand-written recounts of bytes and flops, and the three whose flops alone follow it get flops recounts. --- tests/roofline_binder.py | 27 +- tests/test_roofline_oracle.py | 290 +++++++++++++++++- tests/test_spec_reference.py | 551 ++++++++++++++++++++++++++++++++++ 3 files changed, 855 insertions(+), 13 deletions(-) create mode 100644 tests/test_spec_reference.py diff --git a/tests/roofline_binder.py b/tests/roofline_binder.py index 5a3574473..adfe37919 100644 --- a/tests/roofline_binder.py +++ b/tests/roofline_binder.py @@ -1,6 +1,8 @@ """Build a bytes-oracle case for each workload row of an op from its manifest entry alone. Each row is instantiated, the op constructed from it and its call checked on meta tensors. +An implemented entry's op is its class; a spec-only entry's is a class carrying only what its +signature generates, since the recount needs no implementation. In parallel the oracle counts the traffic the checked call implies -- one read per input it binds, one write per output, both for a written input -- and the caller requires the two to be equal. The `roofline` block is never read. @@ -19,15 +21,36 @@ from tileops.manifest.plan import entry_plan from tileops.manifest.registry import op_class from tileops.manifest.workload import instantiate +from tileops.ops._signature_codegen import install +from tileops.ops.op_base import Op -__all__ = ["manifest_cases"] +__all__ = ["manifest_cases", "signature_class"] + + +def signature_class(op_name: str, entry: dict) -> type: + """An `Op` subclass with *entry*'s generated methods and no kernel.""" + + def construct(self, **params): + vars(self).update(params) + self.dispatch_kernel(None) + + body = {"__init__": construct, "default_kernel_map": property(lambda self: {})} + body["forward"] = body["_eager_forward"] = lambda self, *args: None + cls = type(f"Signature{op_name}", (Op,), body) + if not install(cls, entry): + raise ValueError(f"{op_name}: the signature does not generate") + return cls def manifest_cases(op_name: str): """Yield ``(label, dtype case, op, oracle bytes, oracle read bytes)`` per row and dtype case.""" entry = load_manifest()[op_name] plan = entry_plan(op_name, entry, load_adts()) - cls = op_class(op_name, entry) + cls = ( + op_class(op_name, entry) + if entry["status"] == "implemented" + else signature_class(op_name, entry) + ) for row in entry["workloads"]: for case in row.get("dtype_cases") or [{}]: call = instantiate(plan, row, case) diff --git a/tests/test_roofline_oracle.py b/tests/test_roofline_oracle.py index 3a718ffd9..72aaed731 100644 --- a/tests/test_roofline_oracle.py +++ b/tests/test_roofline_oracle.py @@ -521,7 +521,268 @@ def test_grouped_gemm_does_not_charge_the_padding_offsets_it_ignores(self): assert self._priced(GroupedGemmFwdOp(), tensors)[1] == oracle -# Coverage levels. Every implemented op sits at +def _evaluated(op_name: str, row: dict, case: dict, **values): + """``(flops, bytes)`` the generated evaluator prices for a row of a spec-only entry, and the + call; *values* replace the named metadata tensors' generated contents.""" + import dataclasses + + from tests.roofline_binder import signature_class + from tileops.manifest import load_adts, load_manifest + from tileops.manifest.plan import entry_plan + from tileops.manifest.workload import instantiate + + entry = load_manifest()[op_name] + plan = entry_plan(op_name, entry, load_adts()) + call = instantiate(plan, {**row, "label": "recount"}, case) + specs = { + **call.specs, + **{n: dataclasses.replace(call.specs[n], values=v) for n, v in values.items()}, + } + call = dataclasses.replace(call, specs=specs) + tensors = call.materialize("meta") + cls = signature_class(op_name, entry) + op = cls(**call.arguments(tensors)) + checked = cls._signature.check(op, {t: tensors[t] for t in plan.sig.inputs}) + metadata = {n: torch.tensor(call.values(n)) for n in checked.metadata} + op._signature_call = dataclasses.replace(checked, metadata=metadata) + return op.eval_roofline(), call + + +def _attention_flops(heads, qk, v, scores, rows): + # Per score two contractions and the softmax (5); per output element its divide. + return heads * (scores * (2 * qk + 2 * v + 5) + rows * v) + + +_BF16, _F32, _FP8, _I32, _I64 = ( + torch.bfloat16, + torch.float32, + torch.float8_e4m3fn, + torch.int32, + torch.int64, +) + + +class TestSpecOnlyRecounts: + """Spec-only entries whose traffic or arithmetic follows their metadata values, recounted + by walking the call. Each case picks metadata that reaches the branches the values decide.""" + + @pytest.mark.parametrize( + "row,case", + [ + # Two requests with disjoint pages; two causal query rows each. + ( + {"S_q": 2, "NP": 8, "W": 4, "cache_lens": [5, 9]}, + {"T": "bfloat16", "KV": "bfloat16"}, + ), + # A pool smaller than the requests' pages, so requests share rows; an FP8 cache. + ( + {"S_q": 1, "NP": 3, "W": 3, "cache_lens": [5, 9, 12], "some": ["kv_scale"]}, + {"T": "bfloat16", "KV": "float8_e4m3fn"}, + ), + ], + ) + def test_mla_paged_reads_the_rows_its_block_table_reaches(self, row, case): + name = "MultiHeadLatentAttentionPagedFwdOp" + row = {"H": 3, "DK": 12, "PS": 4, "kv_lora_rank": 8, **row} + (flops, moved), call = _evaluated(name, row, case) + heads, dk, rank, page, s_q = row["H"], row["DK"], row["kv_lora_rank"], row["PS"], row["S_q"] + table, lengths = call.values("block_table"), call.values("cache_seqlens") + fp8 = case["KV"] == "float8_e4m3fn" + scores = rows = 0 + for c in lengths: + for i in range(s_q): + seen = c - s_q + i + 1 + scores, rows = scores + seen, rows + 1 + cache_rows = { + (table[b][j // page], j % page) for b, c in enumerate(lengths) for j in range(c) + } + pages = {(b, j // page) for b, c in enumerate(lengths) for j in range(c)} + batch = len(lengths) + assert moved == _ledger( + name, + q=((batch, s_q, heads, dk), _BF16), + kv_cache=((len(cache_rows), dk), _FP8 if fp8 else _BF16), + block_table=((len(pages),), _I32), + cache_seqlens=((batch,), _I32), + kv_scale=((1,), _F32) if fp8 else None, + o=((batch, s_q, heads, rank), _BF16), + lse=((batch, s_q, heads), _F32), + ) + assert flops == _attention_flops(heads, dk, rank, scores, rows) + ( + heads * rows if fp8 else 0 + ) + + def test_dsa_paged_scores_each_valid_slot_and_reads_the_rows_they_name(self): + name = "DeepSeekSparseAttentionPagedFwdOp" + row = {"S_q": 2, "H": 2, "K": 4, "NP": 6, "PS": 4, "W": 3, "cache_lens": [5, 10]} + # Request 0: query 0 selects nothing, query 1 three positions on two pages. + # Request 1: a repeated slot, then nothing. + indices = [[[-1, -1, -1, -1], [0, 4, 3, -1]], [[3, 3, -1, -1], [-1, -1, -1, -1]]] + (flops, moved), call = _evaluated(name, row, {}, indices=indices) + table, lengths, page = call.values("block_table"), call.values("cache_seqlens"), row["PS"] + valid = [ + [[j for j in slots if 0 <= j < lengths[b]] for slots in per_q] + for b, per_q in enumerate(indices) + ] + scores = sum(len(v) for per_q in valid for v in per_q) + rows = sum(1 for per_q in valid for v in per_q if v) + cache_rows = { + (table[b][j // page], j % page) + for b, per_q in enumerate(valid) + for v in per_q + for j in v + } + pages = {(b, j // page) for b, per_q in enumerate(valid) for v in per_q for j in v} + assert moved == _ledger( + name, + q=((rows, row["H"], 576), _BF16), # the query rows that score something + kv_cache=((len(cache_rows), 656), torch.uint8), + block_table=((len(pages),), _I32), + cache_seqlens=((2,), _I32), + indices=((2, 2, 4), _I32), + o=((2, 2, row["H"], 512), _BF16), + lse=((2, 2, row["H"]), _F32), + ) + # Each distinct row's 512 latent values are dequantized once. + assert flops == _attention_flops(row["H"], 576, 512, scores, rows) + 512 * len(cache_rows) + + def test_paged_kv_cache_write_moves_only_the_tokens_with_a_slot(self): + name = "PagedKVCacheWriteFwdOp" + row = {"N": 5, "H_kv": 2, "D": 4, "NP": 4, "PS": 3, "some": ["k_scale"]} + slots = [7, -1, 2, -1, 11] + (flops, moved), _call = _evaluated( + name, row, {"T": "bfloat16", "KV": "float8_e4m3fn"}, slot_mapping=slots + ) + written = ((3, 2, 4), _FP8) + assert moved == _ledger( + name, + k=((3, 2, 4), _BF16), + v=((3, 2, 4), _BF16), + k_pages_unread=True, + k_pages_write=written, + v_pages_unread=True, + v_pages_write=written, + slot_mapping=((5,), _I64), + k_scale=((1,), _F32), + v_scale=((1,), _F32), + ) + # Per stored FP8 value, the scale and the saturating cast. + assert flops == 2 * 2 * 3 * 2 * 4 + + def test_mla_kv_cache_write_rotates_and_moves_only_the_tokens_with_a_slot(self): + name = "MultiHeadLatentAttentionKVCacheWriteFwdOp" + row = { + "DC": 6, + "PE": 4, + "NP": 4, + "PS": 3, + "P": 16, + "seq_lens": [3, 2], + "fuse_rope": True, + "some": ["scale"], + } + slots = [5, -1, 0, 8, -1] # tokens 0, 2 and 3, at positions 0, 2 and 0 + case = {"T": "bfloat16", "KV": "float8_e4m3fn", "C": "float32"} + (flops, moved), _call = _evaluated(name, row, case, slot_mapping=slots) + assert moved == _ledger( + name, + kv_c=((3, 6), _BF16), + k_pe=((3, 4), _BF16), + kv_cache_unread=True, + kv_cache_write=((3, 10), _FP8), + slot_mapping=((5,), _I64), + scale=((1,), _F32), + positions=((3,), _I64), + cos_sin_cache=((2, 4), _F32), # positions 0 and 2 + ) + assert flops == 3 * (2 * 10 + 3 * 4) + + @pytest.mark.parametrize( + "row,case", + [ + ({"seq_lens": [3, 4], "starts": [2, 5], "some": ["seq_starts"]}, {"KV": "bfloat16"}), + ( + {"seq_lens": [3, 4], "out_dtype": "float16", "some": ["scale"]}, + {"KV": "float8_e4m3fn"}, + ), + ], + ) + def test_paged_kv_cache_gather_reads_the_rows_each_range_reaches(self, row, case): + name = "PagedKVCacheGatherFwdOp" + row = {"T_q": 7, "NP": 6, "PS": 4, "W": 3, "E": [2, 3], **row} + (flops, moved), call = _evaluated(name, row, case) + table, page = call.values("block_table"), row["PS"] + starts = row.get("starts", [0, 0]) + ranges = [range(s, s + n) for s, n in zip(starts, row["seq_lens"], strict=True)] + cache_rows = {(table[b][j // page], j % page) for b, r in enumerate(ranges) for j in r} + pages = {(b, j // page) for b, r in enumerate(ranges) for j in r} + fp8 = case["KV"] == "float8_e4m3fn" + assert moved == _ledger( + name, + dst_unread=True, + dst_write=((7, 2, 3), torch.float16 if fp8 else _BF16), + cache=((len(cache_rows), 2, 3), _FP8 if fp8 else _BF16), + block_table=((len(pages),), _I32), + cu_seq_lens=((3,), _I32), + seq_starts=None if fp8 else ((2,), _I32), + scale=((1,), _F32) if fp8 else None, + ) + assert flops == (7 * 2 * 3 if fp8 else 0) + + def test_fused_qk_norm_rope_touches_the_q_and_k_columns_and_the_named_rows(self): + name = "FusedQKNormRopeFwdOp" + row = {"D": 8, "P": 16, "R": 4, "num_heads": 3, "num_kv_heads": 1, "seq_lens": [3, 2]} + (flops, moved), _call = _evaluated(name, row, {"T": "bfloat16", "C": "float32"}) + qk = ((5, 4 * 8), _BF16) # 5 tokens, 3 q heads and 1 k head of width 8 + assert moved == _ledger( + name, + qkv=qk, + qkv_write=qk, + q_weight=((8,), _BF16), + k_weight=((8,), _BF16), + cos_sin_cache=((3, 4), _F32), # positions 0, 1, 2 + positions=((5,), _I64), + ) + assert flops == 5 * 4 * (4 * 8 + 3 * 4) + + def test_chain_speculative_sampling_prices_the_cheaper_outcome(self): + name = "ChainSpeculativeSamplingFwdOp" + batch, n, vocab = 2, 3, 10 + (flops, moved), _call = _evaluated(name, {"B": batch, "N": n, "V": vocab}, {}) + # With N < V, every draft accepted is cheaper than a rejection at the first: the N + # ratio tests and a draw from target row N; the ids, 2 N probabilities, one row. + assert moved == _ledger( + name, + draft_probs=((batch * n,), _F32), + draft_token_ids=((batch, n), _I32), + target_probs=((batch * (n + vocab),), _F32), + seed=((1,), _I64), + offset=((1,), _I64), + output_token_ids=((batch, n + 1), _I32), + num_accepted=((batch,), _I32), + ) + assert flops == batch * (2 * n + 3 * vocab + 1) + + def test_top_k_masks_pay_nothing_on_a_row_k_leaves_whole(self): + vocab, ks = 10, [3, 10, 12, 1] + row = {"V": vocab, "k_list": ks} + (flops, _moved), _call = _evaluated("TopKMaskFwdOp", row, {"T": "float32"}) + assert flops == sum(2 * vocab for k in ks if k < vocab) + (flops, _moved), _call = _evaluated("TopKTopPMaskFwdOp", row, {"T": "float32"}) + # Top-k where k restricts; over the survivors max, subtract, exp, sum, compare and + # accumulate; p times the sum; the final mask. + assert flops == sum((vocab if k < vocab else 0) + 6 * min(k, vocab) + 1 + vocab for k in ks) + + def test_mla_varlen_scores_each_request_under_its_causal_mask(self): + row = {"T_q": 7, "H": 2, "DN": 4, "PE": 2, "DV": 3, "seq_lens": [3, 4]} + (flops, _moved), _call = _evaluated( + "MultiHeadLatentAttentionVarlenFwdOp", row, {"T": "bfloat16"} + ) + scores = sum(i + 1 for n in row["seq_lens"] for i in range(n)) + assert flops == _attention_flops(2, 6, 3, scores, 7) + + +# Coverage levels. Every op, implemented or spec-only, sits at # exactly one, and the level says what an independent recount rests on. # # one The binder builds the case from the manifest: signature, one workload @@ -554,6 +815,14 @@ def test_grouped_gemm_does_not_charge_the_padding_offsets_it_ignores(self): "NSAVarlenFwdOp": "how much it reads follows the values in `block_counts`", "NSATopkVarlenFwdOp": "`lse_in` is passed and the kernel recomputes the lse instead of reading it", "IndexedExpertMLPFwdOp": "the routed weight reads follow the values in `topk_ids`", + "GroupedQueryAttentionPagedFwdOp": "it reads the rows its page table names, not the pool", + "MultiHeadLatentAttentionPagedFwdOp": "it reads the cache rows its block table reaches, not the pool", + "DeepSeekSparseAttentionPagedFwdOp": "it reads the cache rows its valid index slots name", + "PagedKVCacheWriteFwdOp": "only the tokens `slot_mapping` gives a slot are read and written", + "MultiHeadLatentAttentionKVCacheWriteFwdOp": "only the tokens `slot_mapping` gives a slot are read and written", + "PagedKVCacheGatherFwdOp": "it reads the cache rows each request's range reaches, not the pool", + "FusedQKNormRopeFwdOp": "it leaves the v columns untouched and reads only the named cos/sin rows", + "ChainSpeculativeSamplingFwdOp": "where the chain stops is drawn at run time, so it prices the cheaper outcome", } # Level three: no independent recount is available. Empty, and an entry here has @@ -561,12 +830,11 @@ def test_grouped_gemm_does_not_charge_the_padding_offsets_it_ignores(self): NOT_RECOUNTABLE: dict[str, str] = {} -def _implemented_ops() -> list[str]: +def _entries() -> list[str]: + """Every entry: a spec-only one is recounted from its signature, needing no implementation.""" from tileops.manifest import load_manifest - return sorted( - name for name, entry in load_manifest().items() if entry.get("status") == "implemented" - ) + return sorted(load_manifest()) def _draws_metadata(op_name: str) -> bool: @@ -606,13 +874,13 @@ def _binder_agrees(op_name: str) -> bool: class TestCoverageLevels: - """Every implemented op sits at exactly one level, and the level is the truth.""" + """Every op sits at exactly one level, and the level is the truth.""" def test_a_generated_case_equals_its_op(self): from tests.roofline_binder import manifest_cases checked = 0 - for op_name in _implemented_ops(): + for op_name in _entries(): if op_name in HAND_WRITTEN or op_name in NOT_RECOUNTABLE: continue for label, dtype, op, oracle, _reads in manifest_cases(op_name): @@ -627,7 +895,7 @@ def test_a_generated_case_agrees_on_the_read_half(self): from tests.roofline_binder import manifest_cases checked = 0 - for op_name in _implemented_ops(): + for op_name in _entries(): if op_name in HAND_WRITTEN or op_name in NOT_RECOUNTABLE: continue for label, dtype, op, _oracle, reads in manifest_cases(op_name): @@ -637,11 +905,11 @@ def test_a_generated_case_agrees_on_the_read_half(self): checked += 1 assert checked > 0 - def test_every_implemented_op_sits_at_one_level(self): + def test_every_op_sits_at_one_level(self): both = sorted(set(HAND_WRITTEN) & set(NOT_RECOUNTABLE)) assert not both, f"declared at two levels: {both}" - unknown = sorted((set(HAND_WRITTEN) | set(NOT_RECOUNTABLE)) - set(_implemented_ops())) - assert not unknown, f"declared but not implemented: {unknown}" + unknown = sorted((set(HAND_WRITTEN) | set(NOT_RECOUNTABLE)) - set(_entries())) + assert not unknown, f"declared but not in the manifest: {unknown}" def test_a_declared_op_is_one_the_manifest_does_not_already_check(self): """Level two and three are for ops the manifest cannot recount, not a queue. diff --git a/tests/test_spec_reference.py b/tests/test_spec_reference.py new file mode 100644 index 000000000..1029d101d --- /dev/null +++ b/tests/test_spec_reference.py @@ -0,0 +1,551 @@ +"""Reference conformance of spec-only entries that have no implementation yet. + +Each entry's reference is the torch expression its issue names, or the library reference +where one exists. Every workload row is instantiated; data tensors live on ``meta`` and +metadata tensors on the CPU with the row's generated values, so the reference runs as +written at the row's full size and costs nothing. Checked against the signature: the +reference's outputs have the inferred names, shapes and dtypes, and it writes exactly the +inputs the call's effects mark written. One call the signature rejects is rejected by the +reference too. +""" + +from __future__ import annotations + +import math + +import pytest +import torch +import torch.nn.functional as F +from torch.utils._python_dispatch import TorchDispatchMode + +from tests.roofline_binder import signature_class +from tileops.manifest import load_adts, load_manifest +from tileops.manifest.plan import entry_plan +from tileops.manifest.workload import instantiate + +pytestmark = pytest.mark.smoke + +_INF = float("inf") + + +# ---------------------------------------------------------------- quantization + + +def _int8_scale(amax): + """``amax / 127``, and 1.0 for an all-zero group.""" + return torch.where(amax > 0, amax / 127, torch.ones_like(amax)) + + +def _int8_per_tensor(p, t): + x = t["x"] + _m, _k = x.shape + xf = x.float() + scale = _int8_scale(xf.abs().amax()).reshape(1) + q = torch.round(xf / scale).clamp(-127, 127).to(torch.int8) + return {"q": q, "scale": scale} + + +def _int8_per_channel(p, t): + w = t["w"] + _n, _k = w.shape + wf = w.float() + scale = _int8_scale(wf.abs().amax(dim=1)) + return {"q": torch.round(wf / scale[:, None]).clamp(-127, 127).to(torch.int8), "scale": scale} + + +def _blocks(xf, block=128): + m, k = xf.shape + nb = -(-k // block) + return F.pad(xf, (0, nb * block - k)).view(m, nb, block), nb + + +def _int8_per_block(p, t): + xf = t["x"].float() + m, k = xf.shape + blocks, _nb = _blocks(xf) + scale = _int8_scale(blocks.abs().amax(-1)) + q = torch.round(xf / scale.repeat_interleave(128, 1)[:, :k]).clamp(-127, 127) + return {"q": q.to(torch.int8), "scale": scale} + + +def _int4_per_group(p, t): + w, g = t["w"], p["group_size"] + n, k = w.shape + wg = w.float().view(n, k // g, g) + lo, hi = wg.amin(-1), wg.amax(-1) + scale = torch.where(hi > lo, (hi - lo) / 15, torch.ones_like(hi)) + zero = torch.round(-lo / scale).clamp(0, 15) + q = torch.round(wg / scale[..., None] + zero[..., None]).clamp(0, 15).to(torch.uint8) + # Two values per byte in row order; the byte order GemmW4A16FwdOp consumes is a permutation + # of these bytes that only its repack kernel states, accepted by the GEMM round trip. + q = q.view(n, k // 2, 2) + packed = q[..., 0] | (q[..., 1] << 4) + return { + "packed_weight": packed, + "weight_scale": scale.to(w.dtype), + "weight_zero": zero.to(torch.uint8), + } + + +def _smooth_quant(p, t): + x, smooth = t["x"], t["smooth"] + _m, _k = x.shape + xs = x.float() / smooth + scale = _int8_scale(xs.abs().amax(dim=1)) + return {"q": torch.round(xs / scale[:, None]).clamp(-127, 127).to(torch.int8), "scale": scale} + + +def _dequant(expand): + def reference(p, t): + q, scale = t["q"], t["scale"] + m, k = q.shape + return {"x": (q.float() * expand(scale, m, k)).to(p["out_dtype"])} + + return reference + + +def _fp8_per_block(p, t): + w = t["w"] + n, k = w.shape + nn, nk = -(-n // 128), -(-k // 128) + wf = F.pad(w.float(), (0, nk * 128 - k, 0, nn * 128 - n)).view(nn, 128, nk, 128) + amax = wf.abs().amax(dim=(1, 3)) + scale = torch.where(amax > 0, amax / 448, torch.ones_like(amax)) + full = scale.repeat_interleave(128, 0).repeat_interleave(128, 1)[:n, :k] + return {"q": (w.float() / full).clamp(-448, 448).to(torch.float8_e4m3fn), "scale": scale} + + +# ---------------------------------------------------------------- sampling + + +def _top_k(logits, k): + batch, vocab = logits.shape + k = k.view(batch) + kth = ( + logits.float().sort(-1, descending=True).values.gather(1, (k.clamp(max=vocab) - 1)[:, None]) + ) + kept = (logits.float() >= kth) | (k >= vocab).to(logits.device)[:, None] + return logits.masked_fill(~kept, -_INF) + + +def _top_p(logits, p): + probs = logits.float().softmax(-1) + p = p.view(logits.shape[0]) + sorted_probs, order = probs.sort(-1, descending=True) + exclusive = sorted_probs.cumsum(-1) - sorted_probs + removed = torch.zeros_like(probs, dtype=torch.bool).scatter( + 1, order, exclusive >= p.view(-1, 1) + ) + return logits.masked_fill(removed, -_INF) + + +def _top_k_mask(p, t): + return {"masked_logits": _top_k(t["logits"], t["k"])} + + +def _min_p_mask(p, t): + logits, min_p = t["logits"], t["min_p"] + probs = logits.float().softmax(-1) + removed = probs < min_p.view(logits.shape[0], 1) * probs.amax(-1, keepdim=True) + return {"masked_logits": logits.masked_fill(removed, -_INF)} + + +def _top_p_mask(p, t): + return {"masked_logits": _top_p(t["logits"], t["p"])} + + +def _top_k_top_p_mask(p, t): + return {"masked_logits": _top_p(_top_k(t["logits"], t["k"]), t["p"])} + + +def _sampling_from_probs(p, t): + return {"samples": torch.multinomial(t["probs"], 1).squeeze(1).to(torch.int32)} + + +def _chain_speculative_sampling(p, t): + draft, target, ids = t["draft_probs"], t["target_probs"], t["draft_token_ids"].long() + batch, n, vocab = draft.shape + d = draft.gather(2, ids[..., None]).squeeze(-1) + q = target[:, :n].gather(2, ids[..., None]).squeeze(-1) + accepted = torch.rand(batch, n, device=draft.device) * d < q + num = accepted.int().cumprod(1).sum(1) + padded_draft = torch.cat([draft, torch.zeros_like(draft[:, :1])], 1) + pick = num[:, None, None].expand(batch, 1, vocab) + residual = (target.gather(1, pick) - padded_draft.gather(1, pick)).squeeze(1).clamp_min(0) + resampled = torch.multinomial(residual, 1) + position = torch.arange(n + 1, device=draft.device)[None] + drafts = torch.cat([ids.to(draft.device), torch.full_like(resampled, -1)], 1) + tokens = torch.where( + position < num[:, None], drafts, torch.where(position == num[:, None], resampled, -1) + ) + return {"output_token_ids": tokens.to(torch.int32), "num_accepted": num.to(torch.int32)} + + +# ---------------------------------------------------------------- attention and caches + + +def _attend(q, k, v, scale, visible): + """Float32 softmax attention of ``q [S, H, Dk]`` over ``k [N, Dk]``/``v [N, Dv]`` shared by + every head, or per head when ``k`` is ``[N, H, Dk]``; ``visible [S, N]`` on the CPU.""" + kh = k if k.dim() == 3 else k[:, None].expand(-1, q.shape[1], -1) + vh = v if v.dim() == 3 else v[:, None].expand(-1, q.shape[1], -1) + scores = torch.einsum("shd,nhd->hsn", q.float(), kh.float()) * scale + scores = scores.masked_fill(~visible.to(q.device)[None], -_INF) + return torch.einsum("hsn,nhd->shd", scores.softmax(-1), vh.float()), scores.logsumexp(-1).T + + +def _causal(queries, keys, is_causal): + """Bottom-right aligned visibility of ``queries`` rows over ``keys``.""" + rows = torch.arange(queries)[:, None] + keys - queries + return ( + torch.arange(keys)[None] <= rows + if is_causal + else torch.ones(queries, keys, dtype=torch.bool) + ) + + +def _paged_rows(cache, table, start, end): + """Rows ``[start, end)`` of one request, read through its block table.""" + page_size = cache.shape[1] + pages = cache[table[: -(-end // page_size)].long()] + return pages.flatten(0, 1)[start:end] + + +def _mla_paged(p, t): + q, cache, table, lens = t["q"], t["kv_cache"], t["block_table"], t["cache_seqlens"] + batch, s_q, _h, dk = q.shape + rank = p["kv_lora_rank"] + scale = p["sm_scale"] if p["sm_scale"] is not None else dk**-0.5 + outs, lses = [], [] + for b in range(batch): + kv = _paged_rows(cache, table[b], 0, int(lens[b])).float() + if t.get("kv_scale") is not None: + kv = kv * t["kv_scale"] + o, lse = _attend(q[b], kv, kv[:, :rank], scale, _causal(s_q, kv.shape[0], p["is_causal"])) + outs.append(o) + lses.append(lse) + return {"o": torch.stack(outs).to(q.dtype), "lse": torch.stack(lses)} + + +def _mla_varlen(p, t): + q, k_nope, k_pe, v, cu = t["q"], t["k_nope"], t["k_pe"], t["v"], t["cu_seqlens"].tolist() + heads = q.shape[1] + scale = p["sm_scale"] if p["sm_scale"] is not None else q.shape[-1] ** -0.5 + outs, lses = [], [] + for a, e in zip(cu, cu[1:], strict=False): + k = torch.cat([k_nope[a:e], k_pe[a:e, None].expand(-1, heads, -1)], -1) + o, lse = _attend(q[a:e], k, v[a:e], scale, _causal(e - a, e - a, p["is_causal"])) + outs.append(o) + lses.append(lse) + return {"o": torch.cat(outs).to(q.dtype), "lse": torch.cat(lses)} + + +def _stored(rows, cache, scale): + """Rows as ``cache`` stores them, one cache row each: divided by the scale when it has one.""" + rows = rows.reshape(rows.shape[0], *cache.shape[2:]) + return (rows.float() / scale if scale is not None else rows).to(cache.dtype) + + +def _paged_kv_cache_write(p, t): + slots = t["slot_mapping"] + tokens = (slots >= 0).nonzero().squeeze(1) + for name, pages, scale in (("k", "k_pages", "k_scale"), ("v", "v_pages", "v_scale")): + cache = t[pages] + cache.view(-1, *cache.shape[2:])[slots[tokens]] = _stored( + t[name][tokens], cache, t.get(scale) + ) + return {} + + +def _rope(x, cos_sin, positions, layout): + """Rotate ``x [N, ..., R']`` on its first ``R`` columns by ``cos_sin [P, R]`` at ``positions``.""" + half = cos_sin.shape[1] // 2 + cos, sin = cos_sin[positions].float().chunk(2, -1) + shape = (x.shape[0],) + (1,) * (x.dim() - 2) + (half,) + cos, sin = cos.view(shape), sin.view(shape) + rot, rest = x[..., : 2 * half].float(), x[..., 2 * half :] + if layout == "neox": + a, b = rot[..., :half], rot[..., half:] + rotated = torch.cat([a * cos - b * sin, b * cos + a * sin], -1) + else: + a, b = rot[..., 0::2], rot[..., 1::2] + rotated = torch.stack([a * cos - b * sin, b * cos + a * sin], -1).flatten(-2) + return torch.cat([rotated.to(x.dtype), rest], -1) + + +def _mla_kv_cache_write(p, t): + slots, cache = t["slot_mapping"], t["kv_cache"] + tokens = (slots >= 0).nonzero().squeeze(1) + k_pe = t["k_pe"][tokens] + if p["fuse_rope"]: + k_pe = _rope(k_pe, t["cos_sin_cache"], t["positions"][tokens], p["rope_layout"]) + rows = torch.cat([t["kv_c"][tokens], k_pe], -1) + cache.view(-1, cache.shape[-1])[slots[tokens]] = _stored(rows, cache, t.get("scale")) + return {} + + +def _fused_qk_norm_rope(p, t): + qkv, eps = t["qkv"], p["eps"] + heads, kv_heads = p["num_heads"], p["num_kv_heads"] + dim = t["q_weight"].shape[0] + tokens = qkv.shape[0] + for first, count, weight in ((0, heads, t["q_weight"]), (heads, kv_heads, t["k_weight"])): + view = qkv[:, first * dim : (first + count) * dim].view(tokens, count, dim) + normed = F.rms_norm(view.float(), (dim,), weight.float(), eps).to(qkv.dtype) + view.copy_(_rope(normed, t["cos_sin_cache"], t["positions"], p["rope_layout"])) + return {} + + +def _merge_attention_states(p, t): + s_a, s_b = t["s_a"], t["s_b"] + top = torch.maximum(s_a, s_b) + empty = top == -_INF + shift = torch.where(empty, torch.zeros_like(top), top) + w_a, w_b = (s_a - shift).exp(), (s_b - shift).exp() + total = torch.where(empty, torch.ones_like(top), w_a + w_b) + v = t["v_a"].float() * (w_a / total)[..., None] + t["v_b"].float() * (w_b / total)[..., None] + return {"v": v.to(t["v_a"].dtype), "s": torch.where(empty, top, top + total.log())} + + +def _dsa_paged(p, t): + q, cache, table, lens, indices = ( + t["q"], + t["kv_cache"], + t["block_table"], + t["cache_seqlens"], + t["indices"], + ) + batch, s_q, _h, dk = q.shape + page_size = cache.shape[1] + scale = p["sm_scale"] if p["sm_scale"] is not None else dk**-0.5 + flat = cache.view(-1, cache.shape[-1]) + outs, lses = [], [] + for b in range(batch): + for s in range(s_q): + pos = indices[b, s].long() + pos = pos[(pos >= 0) & (pos < lens[b])] + rows = flat[table[b][pos // page_size].long() * page_size + pos % page_size] + latent = rows[:, :512].view(torch.float8_e4m3fn).float() + scales = rows[:, 512:528].view(torch.float32).repeat_interleave(128, 1) + value = latent * scales + key = torch.cat([value, rows[:, 528:].view(torch.bfloat16).float()], -1) + visible = torch.ones(1, key.shape[0], dtype=torch.bool) + o, lse = _attend(q[b, s : s + 1], key, value, scale, visible) + outs.append(o[0]) + lses.append(lse[0]) + shape = (batch, s_q) + return { + "o": torch.stack(outs).view(*shape, *outs[0].shape).to(q.dtype), + "lse": torch.stack(lses).view(*shape, -1), + } + + +def _paged_kv_cache_gather(p, t): + dst, cache, table = t["dst"], t["cache"], t["block_table"] + cu = t["cu_seq_lens"].tolist() + starts = t["seq_starts"].tolist() if t.get("seq_starts") is not None else [0] * (len(cu) - 1) + for b, (a, e) in enumerate(zip(cu, cu[1:], strict=False)): + rows = _paged_rows(cache, table[b], starts[b], starts[b] + e - a) + if t.get("scale") is not None: + rows = rows.float() * t["scale"] + dst[a:e] = rows.to(dst.dtype) + return {} + + +# ---------------------------------------------------------------- linear attention + + +def _kda(p, t): + naive = pytest.importorskip( + "fla.ops.kda.naive", + reason="KimiDeltaAttentionFwdOp unverified: its reference, FLA (package `fla`), is not installed", + ).naive_recurrent_kda + q, k, v, g, beta = t["q"], t["k"], t["v"], t["g"], t["beta"] + if p["use_qk_l2norm_in_kernel"]: + q, k = ( + F.normalize(q.float(), dim=-1).to(q.dtype), + F.normalize(k.float(), dim=-1).to(k.dtype), + ) + if p["use_gate_in_kernel"]: + heads = v.shape[2] + a = t["A_log"].exp()[:, None] + x = g.float() + (t["dt_bias"].view(heads, -1) if t.get("dt_bias") is not None else 0) + lower = p["lower_bound"] + g = lower * torch.sigmoid(a * x) if lower is not None else -a * F.softplus(x) + if p["use_beta_sigmoid_in_kernel"]: + beta = beta.float().sigmoid() * (2 if p["allow_neg_eigval"] else 1) + cu = t["cu_seqlens"].tolist() if t.get("cu_seqlens") is not None else None + spans = list(zip(cu, cu[1:], strict=False)) if cu else [(0, q.shape[1])] + init = t.get("initial_state") + if init is not None and p["state_v_first"]: + init = init.transpose(-1, -2) + outs, states = [], [] + for i, (a, e) in enumerate(spans): + sel = slice(a, e) + h0 = init[i : i + 1] if cu and init is not None else init + o, s = naive(q[:, sel], k[:, sel], v[:, sel], g[:, sel], beta[:, sel], p["scale"], h0, True) + outs.append(o) + states.append(s) + state = torch.cat(states) if cu else states[0] + if p["state_v_first"]: + state = state.transpose(-1, -2) + return {"o": torch.cat(outs, 1).to(v.dtype), "final_state": state.float()} + + +# ---------------------------------------------------------------- entries + + +def _narrow(name, axis=-1): + """A rejected call: tensor *name* one element shorter along *axis*.""" + + def edit(tensors): + x = tensors[name] + tensors[name] = x.narrow(axis, 0, x.shape[axis] - 1) + + return edit + + +def _unsqueeze(name): + def edit(tensors): + tensors[name] = tensors[name][None] + + return edit + + +def _empty_rows(name): + def edit(tensors): + tensors[name] = tensors[name][:0] + + return edit + + +REFERENCES = { + # op: (reference, the edit of the first row's tensors that the signature rejects) + "INT8QuantPerTensorFwdOp": (_int8_per_tensor, _empty_rows("x")), + "INT8QuantPerChannelFwdOp": (_int8_per_channel, _unsqueeze("w")), + "INT8QuantPerBlockFwdOp": (_int8_per_block, _unsqueeze("x")), + "INT4QuantPerGroupFwdOp": (_int4_per_group, _narrow("w")), + "SmoothQuantFwdOp": (_smooth_quant, _narrow("smooth")), + "INT8DequantPerTensorFwdOp": (_dequant(lambda s, m, k: s), _unsqueeze("q")), + "INT8DequantPerChannelFwdOp": (_dequant(lambda s, m, k: s[:, None]), _narrow("scale")), + "INT8DequantPerBlockFwdOp": ( + _dequant(lambda s, m, k: s.repeat_interleave(128, 1)[:, :k]), + _narrow("scale", 0), + ), + "FP8QuantPerBlockFwdOp": (_fp8_per_block, _unsqueeze("w")), + "TopKMaskFwdOp": (_top_k_mask, _narrow("k")), + "MinPMaskFwdOp": (_min_p_mask, _narrow("min_p", 0)), + "TopPMaskFwdOp": (_top_p_mask, _narrow("p", 0)), + "TopKTopPMaskFwdOp": (_top_k_top_p_mask, _narrow("p", 0)), + "SamplingFromProbsFwdOp": (_sampling_from_probs, _unsqueeze("probs")), + "ChainSpeculativeSamplingFwdOp": (_chain_speculative_sampling, _narrow("draft_token_ids")), + "MultiHeadLatentAttentionPagedFwdOp": (_mla_paged, _narrow("kv_cache")), + "MultiHeadLatentAttentionVarlenFwdOp": (_mla_varlen, _narrow("k_nope")), + "PagedKVCacheWriteFwdOp": (_paged_kv_cache_write, _narrow("v")), + "FusedQKNormRopeFwdOp": (_fused_qk_norm_rope, _narrow("k_weight")), + "MultiHeadLatentAttentionKVCacheWriteFwdOp": (_mla_kv_cache_write, _narrow("kv_cache")), + "MergeAttentionStatesFwdOp": (_merge_attention_states, _narrow("v_b")), + "DeepSeekSparseAttentionPagedFwdOp": (_dsa_paged, _narrow("kv_cache")), + "PagedKVCacheGatherFwdOp": (_paged_kv_cache_gather, _narrow("dst")), + "KimiDeltaAttentionFwdOp": (_kda, _narrow("g")), +} + + +class _Writes(TorchDispatchMode): + """The storages the dispatched ops write, by their schema's mutable arguments.""" + + def __init__(self): + super().__init__() + self.storages = set() + + def __torch_dispatch__(self, func, types, args=(), kwargs=None): + kwargs = kwargs or {} + for i, arg in enumerate(func._schema.arguments): + value = args[i] if i < len(args) else kwargs.get(arg.name) + if ( + arg.alias_info is not None + and arg.alias_info.is_write + and isinstance(value, torch.Tensor) + ): + self.storages.add(value.untyped_storage()._cdata) + return func(*args, **kwargs) + + +def _calls(name): + """Per row and dtype case: the instantiated call, its meta tensors and the reference's + tensors (metadata on the CPU with the row's values).""" + entry = load_manifest()[name] + plan = entry_plan(name, entry, load_adts()) + for row in entry["workloads"]: + for case in row.get("dtype_cases") or [{}]: + call = instantiate(plan, row, case) + meta = call.materialize("meta") + host = dict(meta) + for n, spec in call.specs.items(): + if spec is not None and spec.values is not None and n in meta: + host[n] = torch.tensor(spec.values, dtype=getattr(torch, spec.dtype)).reshape( + spec.shape + ) + yield row["label"], case, plan, call, meta, host + + +def _spec_only_without_class(): + return sorted(name for name in REFERENCES if load_manifest()[name]["status"] == "spec-only") + + +def test_every_entry_named_here_is_spec_only(): + """A reference here stands in for a missing implementation; one with a class is tested by + its own op tests.""" + assert _spec_only_without_class() == sorted(REFERENCES) + + +@pytest.mark.parametrize("name", sorted(REFERENCES)) +def test_reference_agrees_with_the_signature(name): + entry = load_manifest()[name] + reference, _reject = REFERENCES[name] + cls = signature_class(name, entry) + for label, case, plan, call, meta, host in _calls(name): + where = f"{name} {label} {case}" + op = cls(**call.arguments(meta)) + checked = cls._signature.check(op, {n: meta[n] for n in plan.sig.inputs}) + inputs = {n: host[n] for n in plan.sig.inputs} + with _Writes() as writes: + outputs = reference(call.arguments(meta), inputs) + declared = {n: call.tensors[n] for n in plan.sig.outputs} + got = {n: (tuple(v.shape), str(v.dtype).removeprefix("torch.")) for n, v in outputs.items()} + assert got == {n: (tuple(s), d) for n, (s, d) in declared.items()}, where + written = { + n + for n, v in inputs.items() + if v is not None and v.untyped_storage()._cdata in writes.storages + } + assert written == set(checked.written), where + + +@pytest.mark.parametrize("name", sorted(REFERENCES)) +def test_a_call_the_signature_rejects_the_reference_rejects(name): + reference, reject = REFERENCES[name] + label, case, plan, call, meta, host = next(_calls(name)) + cls = signature_class(name, load_manifest()[name]) + op = cls(**call.arguments(meta)) + bad_meta = {n: meta[n] for n in plan.sig.inputs} + bad_host = {n: host[n] for n in plan.sig.inputs} + reject(bad_meta) + reject(bad_host) + with pytest.raises((ValueError, TypeError)): + cls._signature.check(op, bad_meta) + with pytest.raises((RuntimeError, ValueError, IndexError, TypeError)): + reference(call.arguments(meta), bad_host) + + +def test_a_write_only_buffer_is_written_whole(): + """`PagedKVCacheGatherFwdOp`'s `dst` is write-only: on a small call with real values the + reference leaves no row of it unwritten.""" + name = "PagedKVCacheGatherFwdOp" + plan = entry_plan(name, load_manifest()[name], load_adts()) + row = {"T_q": 7, "NP": 6, "PS": 4, "W": 3, "E": [2, 3], "seq_lens": [3, 4], "starts": [2, 5]} + call = instantiate(plan, {**row, "some": ["seq_starts"], "label": "s"}, {"KV": "bfloat16"}) + tensors = call.materialize("cpu") + tensors["dst"].fill_(math.nan) + _paged_kv_cache_gather(call.arguments(tensors), {n: tensors[n] for n in plan.sig.inputs}) + assert not tensors["dst"].isnan().any() From deace79805ebb403c6037d1b8a362cb6f0547871 Mon Sep 17 00:00:00 2001 From: lcy-seso Date: Mon, 28 Sep 2026 06:40:17 +0800 Subject: [PATCH 5/5] [Test][Manifest] Prune scaffolding from the spec-only checks The spec-only-membership check restated a process step, and the two-case MLA paged and gather recounts each split one call's branches across two rows; one row per entry now reaches every branch. --- tests/test_spec_reference.py | 10 ---------- 1 file changed, 10 deletions(-) diff --git a/tests/test_spec_reference.py b/tests/test_spec_reference.py index 1029d101d..c53a26fc3 100644 --- a/tests/test_spec_reference.py +++ b/tests/test_spec_reference.py @@ -489,16 +489,6 @@ def _calls(name): yield row["label"], case, plan, call, meta, host -def _spec_only_without_class(): - return sorted(name for name in REFERENCES if load_manifest()[name]["status"] == "spec-only") - - -def test_every_entry_named_here_is_spec_only(): - """A reference here stands in for a missing implementation; one with a class is tested by - its own op tests.""" - assert _spec_only_without_class() == sorted(REFERENCES) - - @pytest.mark.parametrize("name", sorted(REFERENCES)) def test_reference_agrees_with_the_signature(name): entry = load_manifest()[name]