From 75ebfe3a57799e2fc30be539604f59e3bce3f342 Mon Sep 17 00:00:00 2001 From: superAngGao Date: Mon, 24 Aug 2026 17:30:40 +0800 Subject: [PATCH 01/11] [Refactor][Attention] Establish dense GQA Op boundary --- src/tileops/manifest/attention.yaml | 65 +++++++ src/tileops/ops/__init__.py | 1 + src/tileops/ops/attention/__init__.py | 2 + src/tileops/ops/attention/gqa.py | 235 +++++++++++++++++++++++++- 4 files changed, 302 insertions(+), 1 deletion(-) diff --git a/src/tileops/manifest/attention.yaml b/src/tileops/manifest/attention.yaml index 10c324da8..2ccd680b5 100644 --- a/src/tileops/manifest/attention.yaml +++ b/src/tileops/manifest/attention.yaml @@ -98,6 +98,71 @@ MultiHeadAttentionBwdOp: bench: benchmarks/ops/attention/bench_mha.py bench_manifest_driven: true +GroupedQueryAttentionDenseFwdOp: + ref_api: "torch.nn.functional.scaled_dot_product_attention" + family: attention + status: spec-only + + signature: + inputs: + q: {dtype: "float16 | bfloat16 | float8_e4m3fn"} + k: {dtype: "same_as(q)"} + v: {dtype: "same_as(q)"} + q_scale: {dtype: "float32", optional: true} + k_scale: {dtype: "float32", optional: true} + v_scale: {dtype: "float32", optional: true} + rope_cos: {dtype: "same_as(o)", optional: true} + rope_sin: {dtype: "same_as(o)", optional: true} + outputs: + o: {dtype: "float16 | bfloat16"} + params: + is_causal: {type: bool, default: true} + sm_scale: {type: "float | None", default: null} + softcap: {type: "float | None", default: null} + window_size_left: {type: int, default: -1} + window_size_right: {type: int, default: -1} + dtype: {type: "torch.dtype | None", default: null} + pos_encoding_mode: {type: str, default: none} + rotary_dim: {type: "int | None", default: null} + rope_layout: {type: str, default: neox} + dtype_combos: + - {q: float16, k: float16, v: float16, o: float16} + - {q: bfloat16, k: bfloat16, v: bfloat16, o: bfloat16} + - {q: float8_e4m3fn, k: float8_e4m3fn, v: float8_e4m3fn, o: float16} + - {q: float8_e4m3fn, k: float8_e4m3fn, v: float8_e4m3fn, o: bfloat16} + shape_rules: + - "q.shape == (B, S_q, H, D)" + - "k.shape == (B, S_kv, H_kv, D)" + - "v.shape == (B, S_kv, H_kv, D)" + - "(q_scale is None) == (k_scale is None)" + - "(q_scale is None) == (v_scale is None)" + - "q_scale is None or q_scale.shape == (B, H_kv)" + - "k_scale is None or k_scale.shape == (B, H_kv)" + - "v_scale is None or v_scale.shape == (B, H_kv)" + - "(rope_cos is None) == (rope_sin is None)" + - "(pos_encoding_mode == 'rope') == (rope_cos is not None)" + - "rope_cos is None or rope_cos.ndim == 2" + - "rope_cos is None or rope_sin is None or rope_sin.shape == rope_cos.shape" + - "rope_cos is None or rope_cos.shape[0] >= S_kv" + - "rope_cos is None or rope_cos.shape[1] == (rotary_dim if rotary_dim is not None else D) // 2" + - "pos_encoding_mode == 'none' or pos_encoding_mode == 'rope'" + - "pos_encoding_mode == 'rope' or rotary_dim is None" + - "rotary_dim is None or (rotary_dim > 0 and rotary_dim % 2 == 0 and rotary_dim <= D)" + - "rope_layout == 'neox' or rope_layout == 'interleaved'" + - "o.shape == (B, S_q, H, D)" + - "H % H_kv == 0" + + workloads: [] + + roofline: + func: "tileops.perf.formulas.gqa_fwd_roofline" + + source: + kernel: [] + op: tileops/ops/attention/gqa.py + test: tests/ops/attention/test_gqa.py + bench: benchmarks/ops/attention/bench_gqa.py + GroupedQueryAttentionFwdOp: ref_api: "torch.nn.functional.scaled_dot_product_attention" family: attention diff --git a/src/tileops/ops/__init__.py b/src/tileops/ops/__init__.py index c211a2632..cf109a4de 100644 --- a/src/tileops/ops/__init__.py +++ b/src/tileops/ops/__init__.py @@ -3,6 +3,7 @@ GroupedQueryAttentionBwdOp, GroupedQueryAttentionDecodePagedWithKVCacheFwdOp, GroupedQueryAttentionDecodeWithKVCacheFwdOp, + GroupedQueryAttentionDenseFwdOp, GroupedQueryAttentionFwdOp, GroupedQueryAttentionPrefillFwdOp, GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp, diff --git a/src/tileops/ops/attention/__init__.py b/src/tileops/ops/attention/__init__.py index 2f91fdf8a..d6e9ca0b4 100644 --- a/src/tileops/ops/attention/__init__.py +++ b/src/tileops/ops/attention/__init__.py @@ -9,6 +9,7 @@ GroupedQueryAttentionBwdOp, GroupedQueryAttentionDecodePagedWithKVCacheFwdOp, GroupedQueryAttentionDecodeWithKVCacheFwdOp, + GroupedQueryAttentionDenseFwdOp, GroupedQueryAttentionFwdOp, GroupedQueryAttentionPrefillFwdOp, GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp, @@ -28,6 +29,7 @@ "GroupedQueryAttentionBwdOp", "GroupedQueryAttentionDecodePagedWithKVCacheFwdOp", "GroupedQueryAttentionDecodeWithKVCacheFwdOp", + "GroupedQueryAttentionDenseFwdOp", "GroupedQueryAttentionFwdOp", "GroupedQueryAttentionPrefillFwdOp", "GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp", diff --git a/src/tileops/ops/attention/gqa.py b/src/tileops/ops/attention/gqa.py index 548499897..13cecf2dd 100644 --- a/src/tileops/ops/attention/gqa.py +++ b/src/tileops/ops/attention/gqa.py @@ -1,8 +1,10 @@ -from typing import Dict, Optional +import math +from typing import Dict, Optional, Protocol, cast import torch import torch.nn.functional as F +from tileops.backend import Target from tileops.kernels.attention import ( FlashAttnBwdPreprocessKernel, GQABwdWgmmaPipelinedKernel, @@ -40,6 +42,7 @@ "GroupedQueryAttentionBwdOp", "GroupedQueryAttentionDecodePagedWithKVCacheFwdOp", "GroupedQueryAttentionDecodeWithKVCacheFwdOp", + "GroupedQueryAttentionDenseFwdOp", "GroupedQueryAttentionFwdOp", "GroupedQueryAttentionPrefillFwdOp", "GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp", @@ -128,6 +131,236 @@ def _build_packed_prefill_kernel( ) +class _DenseFwdCallable(Protocol): + """Target-owned second dispatch layer for the Dense GQA boundary. + + A target builder returns this callable for one input signature. Any finer + selection and caching of concrete kernels belongs here, not in the Op. + """ + + def __call__( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_scale: Optional[torch.Tensor], + k_scale: Optional[torch.Tensor], + v_scale: Optional[torch.Tensor], + rope_cos: Optional[torch.Tensor], + rope_sin: Optional[torch.Tensor], + ) -> torch.Tensor: ... + + +class GroupedQueryAttentionDenseFwdOp(Op): + """Shape-agnostic dense BSHD GQA for prefill and contiguous decode. + + The public ABI covers rectangular Q/KV lengths, causal/window/softcap modes, + FP16/BF16 and scaled FP8 inputs, and optional caller-owned RoPE tables. + This target-neutral layer caches a target-owned :class:`_DenseFwdCallable` + under the standard input-signature rules; concrete kernel dispatch remains + inside that callable. The in-tree callable is deferred to a follow-up. + """ + + def __init__( + self, + is_causal: bool = True, + sm_scale: Optional[float] = None, + softcap: Optional[float] = None, + window_size_left: int = -1, + window_size_right: int = -1, + dtype: Optional[torch.dtype] = None, + pos_encoding_mode: str = "none", + rotary_dim: Optional[int] = None, + rope_layout: str = "neox", + *, + target: Target = None, + ) -> None: + """Configure the public Dense BSHD GQA interface.""" + if pos_encoding_mode not in ("none", "rope"): + raise ValueError(f"pos_encoding_mode must be 'none' or 'rope', got {pos_encoding_mode}") + if rotary_dim is not None and pos_encoding_mode != "rope": + raise ValueError("rotary_dim requires pos_encoding_mode='rope'") + if sm_scale is not None and not math.isfinite(sm_scale): + raise ValueError(f"sm_scale must be finite, got {sm_scale}") + + self.is_causal = is_causal + self.sm_scale = sm_scale + self.softcap = _score_softcap(softcap) + if window_size_left < -1: + raise ValueError("window_size_left must be -1 (unlimited) or >= 0") + if window_size_right < -1: + raise ValueError("window_size_right must be -1 (unlimited) or >= 0") + self.window_size_left = window_size_left + self.window_size_right = window_size_right + self.pos_encoding_mode = pos_encoding_mode + self.rotary_dim = rotary_dim + self.rope_layout = rope_layout + if dtype is not None: + _validate_attention_dtype(dtype) + self.dtype = dtype + self.target = target + self.dispatch_kernel() + + @property + def default_kernel_map(self) -> Dict[str, Kernel]: + return {} + + def _infer_output_shapes( + self, + q_shape: tuple[int, ...], + k_shape: tuple[int, ...], + v_shape: tuple[int, ...], + q_scale_shape: Optional[tuple[int, ...]] = None, + k_scale_shape: Optional[tuple[int, ...]] = None, + v_scale_shape: Optional[tuple[int, ...]] = None, + rope_cos_shape: Optional[tuple[int, ...]] = None, + rope_sin_shape: Optional[tuple[int, ...]] = None, + ) -> Dict[str, tuple[int, ...]]: + return {"o": tuple(q_shape)} + + def _validate_dtypes( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_scale: Optional[torch.Tensor] = None, + k_scale: Optional[torch.Tensor] = None, + v_scale: Optional[torch.Tensor] = None, + rope_cos: Optional[torch.Tensor] = None, + rope_sin: Optional[torch.Tensor] = None, + ) -> None: + allowed = {torch.float16, torch.bfloat16, fp8_dtype()} + if q.dtype not in allowed: + raise ValueError("q must have float16, bfloat16, or float8_e4m3fn dtype") + if k.dtype != q.dtype or v.dtype != q.dtype: + raise ValueError("q, k, and v must have the same dtype") + is_fp8 = q.dtype == fp8_dtype() + output_dtype = self.dtype or q.dtype + if is_fp8 and self.dtype not in (torch.float16, torch.bfloat16): + raise ValueError("FP8 input requires dtype=torch.float16 or torch.bfloat16") + if not is_fp8 and output_dtype != q.dtype: + raise ValueError("16-bit output dtype must match q, k, and v") + for name, scale in zip( + ("q_scale", "k_scale", "v_scale"), + (q_scale, k_scale, v_scale), + strict=True, + ): + if scale is not None and scale.dtype != torch.float32: + raise ValueError(f"{name} must have float32 dtype") + for name, table in (("rope_cos", rope_cos), ("rope_sin", rope_sin)): + if table is not None and table.dtype != output_dtype: + raise ValueError(f"{name} must have dtype {output_dtype}") + + def eval_roofline(self) -> tuple[int, int]: + raise NotImplementedError("Dense GQA has no in-tree implementation yet") + + def _get_callable( + self, + inputs: tuple[Optional[torch.Tensor], ...], + ) -> _DenseFwdCallable: + """Resolve layer one; the returned target callable is layer two.""" + # BUILTIN follow-up: use this same seam with ``key=...`` and ``build=...``. + return cast( + _DenseFwdCallable, + self.get_or_build_kernel("gqa_dense", inputs), + ) + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_scale: Optional[torch.Tensor] = None, + k_scale: Optional[torch.Tensor] = None, + v_scale: Optional[torch.Tensor] = None, + rope_cos: Optional[torch.Tensor] = None, + rope_sin: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Run Dense BSHD GQA through the selected callable.""" + batch, seq_len_q, heads, dim = q.shape + batch_kv, seq_len_kv, heads_kv, dim_kv = k.shape + + if k.shape != v.shape: + raise ValueError("k and v must have the same shape") + if batch_kv != batch or dim_kv != dim: + raise ValueError("q and k/v must have matching batch and head dimension") + + _validate_positive(batch=batch, seq_len_q=seq_len_q, seq_len_kv=seq_len_kv) + _validate_gqa_dims(heads, heads_kv, dim) + + if self.is_causal and seq_len_q > seq_len_kv: + raise ValueError("causal dense attention requires seq_len_q <= seq_len_kv") + if self.pos_encoding_mode == "rope" and seq_len_q > seq_len_kv: + raise ValueError("fused RoPE requires seq_len_q <= seq_len_kv") + + resolved_sm_scale = _attention_scale(dim, self.sm_scale) + if not math.isfinite(resolved_sm_scale): + raise ValueError("sm_scale must be finite") + if self.pos_encoding_mode not in ("none", "rope"): + raise ValueError("pos_encoding_mode must be 'none' or 'rope'") + if self.rope_layout not in ("neox", "interleaved"): + raise ValueError("rope_layout must be 'neox' or 'interleaved'") + if self.pos_encoding_mode == "rope": + resolved_rotary_dim = _rope_rotary_dim(dim, self.rotary_dim) + else: + if self.rotary_dim is not None: + raise ValueError("rotary_dim requires pos_encoding_mode='rope'") + resolved_rotary_dim = None + + self._validate_dtypes(q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) + + scales = (q_scale, k_scale, v_scale) + has_scales = tuple(scale is not None for scale in scales) + if any(has_scales) and not all(has_scales): + raise ValueError("q_scale, k_scale, and v_scale must be supplied together") + is_fp8 = q.dtype == fp8_dtype() + if is_fp8 and not all(has_scales): + raise ValueError("FP8 input requires q_scale, k_scale, and v_scale") + if not is_fp8 and all(has_scales): + raise ValueError("q_scale, k_scale, and v_scale are only valid for FP8 input") + + for name, tensor in (("k", k), ("v", v)): + if tensor.device != q.device: + raise ValueError(f"{name} must be on the same device as q") + normalized_scales = [] + for name, scale in zip(("q_scale", "k_scale", "v_scale"), scales, strict=True): + if scale is None: + continue + if scale.device != q.device: + raise ValueError(f"{name} must be on the same device as q") + if tuple(scale.shape) != (batch, heads_kv): + raise ValueError(f"{name} must have shape {(batch, heads_kv)}") + normalized_scales.append(scale.contiguous()) + resolved_scales = tuple(normalized_scales) if normalized_scales else (None, None, None) + + if (rope_cos is None) != (rope_sin is None): + raise ValueError("rope_cos and rope_sin must be supplied together") + if self.pos_encoding_mode != "rope": + if rope_cos is not None: + raise ValueError("RoPE tables require pos_encoding_mode='rope'") + resolved_rope = (None, None) + else: + if rope_cos is None or rope_sin is None: + raise ValueError("pos_encoding_mode='rope' requires rope_cos and rope_sin") + expected_columns = (resolved_rotary_dim or dim) // 2 + for name, table in (("rope_cos", rope_cos), ("rope_sin", rope_sin)): + if table.device != q.device: + raise ValueError(f"{name} must be on the same device as q") + if table.ndim != 2: + raise ValueError(f"{name} must be 2-dimensional") + if table.shape[0] < seq_len_kv or table.shape[1] != expected_columns: + raise ValueError( + f"{name} must have shape [max_position >= {seq_len_kv}, {expected_columns}]" + ) + if rope_cos.shape != rope_sin.shape: + raise ValueError("rope_cos and rope_sin must have the same shape") + resolved_rope = (rope_cos.contiguous(), rope_sin.contiguous()) + + inputs = (q.contiguous(), k.contiguous(), v.contiguous(), *resolved_scales, *resolved_rope) + callable_impl = self._get_callable(inputs) + return callable_impl(*inputs) + + class GroupedQueryAttentionFwdOp(Op): """Compatibility square GQA forward wrapper. Public layout: BSHD.""" From 1cc7e045e09b51249ad3acefbfaf4b05975bcbed Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 15:09:40 +0800 Subject: [PATCH 02/11] [Refactor][Attention] Address dense GQA boundary review --- src/tileops/manifest/attention.yaml | 4 + src/tileops/ops/attention/gqa.py | 201 ++++++++++++++++------------ tests/ops/attention/test_gqa.py | 68 ++++++++++ 3 files changed, 186 insertions(+), 87 deletions(-) diff --git a/src/tileops/manifest/attention.yaml b/src/tileops/manifest/attention.yaml index 2ccd680b5..8ff7c71e7 100644 --- a/src/tileops/manifest/attention.yaml +++ b/src/tileops/manifest/attention.yaml @@ -134,6 +134,10 @@ GroupedQueryAttentionDenseFwdOp: - "q.shape == (B, S_q, H, D)" - "k.shape == (B, S_kv, H_kv, D)" - "v.shape == (B, S_kv, H_kv, D)" + # Rectangular calls use bottom-right alignment: query i is at + # absolute position i + S_kv - S_q. + - "not is_causal or S_q <= S_kv" + - "pos_encoding_mode != 'rope' or S_q <= S_kv" - "(q_scale is None) == (k_scale is None)" - "(q_scale is None) == (v_scale is None)" - "q_scale is None or q_scale.shape == (B, H_kv)" diff --git a/src/tileops/ops/attention/gqa.py b/src/tileops/ops/attention/gqa.py index 13cecf2dd..b22c495b2 100644 --- a/src/tileops/ops/attention/gqa.py +++ b/src/tileops/ops/attention/gqa.py @@ -1,5 +1,5 @@ import math -from typing import Dict, Optional, Protocol, cast +from typing import Callable, Dict, Optional import torch import torch.nn.functional as F @@ -131,34 +131,39 @@ def _build_packed_prefill_kernel( ) -class _DenseFwdCallable(Protocol): - """Target-owned second dispatch layer for the Dense GQA boundary. - - A target builder returns this callable for one input signature. Any finer - selection and caching of concrete kernels belongs here, not in the Op. - """ - - def __call__( - self, - q: torch.Tensor, - k: torch.Tensor, - v: torch.Tensor, - q_scale: Optional[torch.Tensor], - k_scale: Optional[torch.Tensor], - v_scale: Optional[torch.Tensor], - rope_cos: Optional[torch.Tensor], - rope_sin: Optional[torch.Tensor], - ) -> torch.Tensor: ... - - class GroupedQueryAttentionDenseFwdOp(Op): - """Shape-agnostic dense BSHD GQA for prefill and contiguous decode. - - The public ABI covers rectangular Q/KV lengths, causal/window/softcap modes, - FP16/BF16 and scaled FP8 inputs, and optional caller-owned RoPE tables. - This target-neutral layer caches a target-owned :class:`_DenseFwdCallable` - under the standard input-signature rules; concrete kernel dispatch remains - inside that callable. The in-tree callable is deferred to a follow-up. + r"""Shape-agnostic dense BSHD GQA for prefill and contiguous decode. + + Let ``g = H / H_kv`` and ``r(h) = floor(h / g)`` map query head ``h`` to + its KV head. Rectangular attention uses bottom-right alignment: + + $$ + p_i = i + S_{kv} - S_q. + $$ + + Causal attention admits key position ``j`` when ``j <= p_i``. A finite + window additionally requires ``p_i - left <= j <= p_i + right``. Fused + RoPE rotates Q at ``p_i`` and K at ``j``; therefore causal and fused-RoPE + calls require ``S_q <= S_kv``. + + For FP8 inputs, each query head uses the scale of its KV group: + + $$ + \hat q_{bih} = q_{bih} qscale_{b,r(h)},\quad + \hat k_{bjr} = k_{bjr} kscale_{b,r},\quad + \hat v_{bjr} = v_{bjr} vscale_{b,r}. + $$ + + With ``alpha = sm_scale`` (default ``1 / sqrt(D)``), scores are + ``z = alpha * dot(q_hat, k_hat)``. When ``softcap > 0`` the score becomes + ``softcap * tanh(z / softcap)`` before masking and softmax. Dot products, + softmax, and the weighted V reduction accumulate in FP32; the result is + cast to ``dtype``. + + The Op validates this public contract and uses the standard + :meth:`Op.get_or_build_kernel` specialization cache. A later BUILTIN PR + will supply the shape/dtype key, select one concrete kernel on a miss, and + store that kernel directly; there is no second callable-owned cache. """ def __init__( @@ -180,11 +185,20 @@ def __init__( raise ValueError(f"pos_encoding_mode must be 'none' or 'rope', got {pos_encoding_mode}") if rotary_dim is not None and pos_encoding_mode != "rope": raise ValueError("rotary_dim requires pos_encoding_mode='rope'") + if rotary_dim is not None: + _validate_positive(rotary_dim=rotary_dim) + if rotary_dim % 2 != 0: + raise ValueError("rotary_dim must be even") + if rope_layout not in ("neox", "interleaved"): + raise ValueError("rope_layout must be 'neox' or 'interleaved'") if sm_scale is not None and not math.isfinite(sm_scale): raise ValueError(f"sm_scale must be finite, got {sm_scale}") self.is_causal = is_causal self.sm_scale = sm_scale + # This default is shape-independent, so targets receive one normalized + # convention. ``sm_scale`` stays None so each implementation resolves + # its shape-dependent default from the call's D. self.softcap = _score_softcap(softcap) if window_size_left < -1: raise ValueError("window_size_left must be -1 (unlimited) or >= 0") @@ -252,34 +266,26 @@ def _validate_dtypes( raise ValueError(f"{name} must have dtype {output_dtype}") def eval_roofline(self) -> tuple[int, int]: + """Keep this spec-only Op concrete until its roofline is implemented.""" raise NotImplementedError("Dense GQA has no in-tree implementation yet") - def _get_callable( - self, - inputs: tuple[Optional[torch.Tensor], ...], - ) -> _DenseFwdCallable: - """Resolve layer one; the returned target callable is layer two.""" - # BUILTIN follow-up: use this same seam with ``key=...`` and ``build=...``. - return cast( - _DenseFwdCallable, - self.get_or_build_kernel("gqa_dense", inputs), - ) - - def forward( + def _validate_forward_inputs( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, - q_scale: Optional[torch.Tensor] = None, - k_scale: Optional[torch.Tensor] = None, - v_scale: Optional[torch.Tensor] = None, - rope_cos: Optional[torch.Tensor] = None, - rope_sin: Optional[torch.Tensor] = None, - ) -> torch.Tensor: - """Run Dense BSHD GQA through the selected callable.""" + q_scale: Optional[torch.Tensor], + k_scale: Optional[torch.Tensor], + v_scale: Optional[torch.Tensor], + rope_cos: Optional[torch.Tensor], + rope_sin: Optional[torch.Tensor], + ) -> None: + for name, tensor in (("q", q), ("k", k), ("v", v)): + if tensor.ndim != 4: + raise ValueError(f"{name} must be a rank-4 BSHD tensor") + batch, seq_len_q, heads, dim = q.shape batch_kv, seq_len_kv, heads_kv, dim_kv = k.shape - if k.shape != v.shape: raise ValueError("k and v must have the same shape") if batch_kv != batch or dim_kv != dim: @@ -287,26 +293,11 @@ def forward( _validate_positive(batch=batch, seq_len_q=seq_len_q, seq_len_kv=seq_len_kv) _validate_gqa_dims(heads, heads_kv, dim) - if self.is_causal and seq_len_q > seq_len_kv: raise ValueError("causal dense attention requires seq_len_q <= seq_len_kv") if self.pos_encoding_mode == "rope" and seq_len_q > seq_len_kv: raise ValueError("fused RoPE requires seq_len_q <= seq_len_kv") - resolved_sm_scale = _attention_scale(dim, self.sm_scale) - if not math.isfinite(resolved_sm_scale): - raise ValueError("sm_scale must be finite") - if self.pos_encoding_mode not in ("none", "rope"): - raise ValueError("pos_encoding_mode must be 'none' or 'rope'") - if self.rope_layout not in ("neox", "interleaved"): - raise ValueError("rope_layout must be 'neox' or 'interleaved'") - if self.pos_encoding_mode == "rope": - resolved_rotary_dim = _rope_rotary_dim(dim, self.rotary_dim) - else: - if self.rotary_dim is not None: - raise ValueError("rotary_dim requires pos_encoding_mode='rope'") - resolved_rotary_dim = None - self._validate_dtypes(q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) scales = (q_scale, k_scale, v_scale) @@ -322,7 +313,6 @@ def forward( for name, tensor in (("k", k), ("v", v)): if tensor.device != q.device: raise ValueError(f"{name} must be on the same device as q") - normalized_scales = [] for name, scale in zip(("q_scale", "k_scale", "v_scale"), scales, strict=True): if scale is None: continue @@ -330,35 +320,72 @@ def forward( raise ValueError(f"{name} must be on the same device as q") if tuple(scale.shape) != (batch, heads_kv): raise ValueError(f"{name} must have shape {(batch, heads_kv)}") - normalized_scales.append(scale.contiguous()) - resolved_scales = tuple(normalized_scales) if normalized_scales else (None, None, None) if (rope_cos is None) != (rope_sin is None): raise ValueError("rope_cos and rope_sin must be supplied together") if self.pos_encoding_mode != "rope": if rope_cos is not None: raise ValueError("RoPE tables require pos_encoding_mode='rope'") - resolved_rope = (None, None) - else: - if rope_cos is None or rope_sin is None: - raise ValueError("pos_encoding_mode='rope' requires rope_cos and rope_sin") - expected_columns = (resolved_rotary_dim or dim) // 2 - for name, table in (("rope_cos", rope_cos), ("rope_sin", rope_sin)): - if table.device != q.device: - raise ValueError(f"{name} must be on the same device as q") - if table.ndim != 2: - raise ValueError(f"{name} must be 2-dimensional") - if table.shape[0] < seq_len_kv or table.shape[1] != expected_columns: - raise ValueError( - f"{name} must have shape [max_position >= {seq_len_kv}, {expected_columns}]" - ) - if rope_cos.shape != rope_sin.shape: - raise ValueError("rope_cos and rope_sin must have the same shape") - resolved_rope = (rope_cos.contiguous(), rope_sin.contiguous()) - - inputs = (q.contiguous(), k.contiguous(), v.contiguous(), *resolved_scales, *resolved_rope) - callable_impl = self._get_callable(inputs) - return callable_impl(*inputs) + return + if rope_cos is None or rope_sin is None: + raise ValueError("pos_encoding_mode='rope' requires rope_cos and rope_sin") + + expected_columns = _rope_rotary_dim(dim, self.rotary_dim) // 2 + for name, table in (("rope_cos", rope_cos), ("rope_sin", rope_sin)): + if table.device != q.device: + raise ValueError(f"{name} must be on the same device as q") + if table.ndim != 2: + raise ValueError(f"{name} must be 2-dimensional") + if table.shape[0] < seq_len_kv or table.shape[1] != expected_columns: + raise ValueError( + f"{name} must have shape [max_position >= {seq_len_kv}, {expected_columns}]" + ) + if rope_cos.shape != rope_sin.shape: + raise ValueError("rope_cos and rope_sin must have the same shape") + + @staticmethod + def _resolve_optional_inputs( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_scale: Optional[torch.Tensor], + k_scale: Optional[torch.Tensor], + v_scale: Optional[torch.Tensor], + rope_cos: Optional[torch.Tensor], + rope_sin: Optional[torch.Tensor], + ) -> tuple[Optional[torch.Tensor], ...]: + """Preserve the manifest's eight positional input slots.""" + return tuple( + tensor.contiguous() if tensor is not None else None + for tensor in (q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) + ) + + def _get_kernel( + self, inputs: tuple[Optional[torch.Tensor], ...] + ) -> Callable[..., torch.Tensor]: + """Resolve the implementation stored in the Op's single cache layer.""" + # BUILTIN follow-up: pass a shape/dtype ``key`` and a ``build`` closure + # that selects and constructs one concrete kernel on a cache miss. + return self.get_or_build_kernel("gqa_dense", inputs) + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_scale: Optional[torch.Tensor] = None, + k_scale: Optional[torch.Tensor] = None, + v_scale: Optional[torch.Tensor] = None, + rope_cos: Optional[torch.Tensor] = None, + rope_sin: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + """Validate, normalize, resolve one concrete implementation, and run it.""" + self._validate_forward_inputs(q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) + inputs = self._resolve_optional_inputs( + q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin + ) + kernel = self._get_kernel(inputs) + return kernel(*inputs) class GroupedQueryAttentionFwdOp(Op): diff --git a/tests/ops/attention/test_gqa.py b/tests/ops/attention/test_gqa.py index fe077673d..116f32573 100644 --- a/tests/ops/attention/test_gqa.py +++ b/tests/ops/attention/test_gqa.py @@ -14,6 +14,7 @@ ) from tileops.ops import ( GroupedQueryAttentionBwdOp, + GroupedQueryAttentionDenseFwdOp, GroupedQueryAttentionFwdOp, GroupedQueryAttentionPrefillFwdOp, GroupedQueryAttentionPrefillVarlenFwdOp, @@ -33,6 +34,73 @@ } +class _DenseBoundaryTestOp(GroupedQueryAttentionDenseFwdOp): + """Exercise the public boundary without adding an in-tree kernel.""" + + resolved_inputs: tuple[Optional[torch.Tensor], ...] + + def _get_kernel(self, inputs): + self.resolved_inputs = inputs + return lambda *args: args[0] + + +def _dense_boundary_inputs(seq_len_q: int = 1, seq_len_kv: int = 4): + q = torch.randn(1, seq_len_q, 4, 8, dtype=torch.float16) + k = torch.randn(1, seq_len_kv, 2, 8, dtype=torch.float16) + return q, k, torch.randn_like(k) + + +@pytest.mark.smoke +@pytest.mark.parametrize("name", ["q", "k", "v"]) +def test_dense_gqa_rejects_non_bshd_inputs(name: str) -> None: + tensors = list(_dense_boundary_inputs()) + index = ("q", "k", "v").index(name) + tensors[index] = tensors[index].squeeze(0) + + with pytest.raises(ValueError, match=rf"{name} must be a rank-4 BSHD tensor"): + _DenseBoundaryTestOp()(*tensors) + + +@pytest.mark.smoke +def test_dense_gqa_accepts_rectangular_decode_and_keeps_optional_slots() -> None: + q, k, v = _dense_boundary_inputs(seq_len_q=1, seq_len_kv=4) + op = _DenseBoundaryTestOp(is_causal=True) + + assert op(q, k, v).shape == q.shape + assert len(op.resolved_inputs) == 8 + assert op.resolved_inputs[3:] == (None, None, None, None, None) + + +@pytest.mark.smoke +@pytest.mark.parametrize( + "kwargs, message", + [ + ({"is_causal": True}, "causal dense attention"), + ({"is_causal": False, "pos_encoding_mode": "rope"}, "fused RoPE"), + ], +) +def test_dense_gqa_rejects_bottom_right_modes_when_q_is_longer_than_kv( + kwargs: dict[str, object], message: str +) -> None: + q, k, v = _dense_boundary_inputs(seq_len_q=4, seq_len_kv=1) + + with pytest.raises(ValueError, match=message): + _DenseBoundaryTestOp(**kwargs)(q, k, v) + + +@pytest.mark.smoke +def test_dense_gqa_rejects_rope_tables_shorter_than_the_kv_positions() -> None: + q, k, v = _dense_boundary_inputs(seq_len_q=1, seq_len_kv=4) + rope = torch.randn(3, 4, dtype=torch.float16) + + with pytest.raises(ValueError, match=r"max_position >= 4"): + _DenseBoundaryTestOp( + is_causal=False, + pos_encoding_mode="rope", + rotary_dim=8, + )(q, k, v, rope_cos=rope, rope_sin=rope) + + def _selected_prefill_kernel_cls(op: GroupedQueryAttentionPrefillFwdOp) -> type: """Kernel class selection picks for a uniform, non-FP8 packed prefill call.""" call = op.attention_call(is_fp8=False, is_uniform=True) From cff870defc64d9c254b61873c74662868a457767 Mon Sep 17 00:00:00 2001 From: superAngGao Date: Mon, 24 Aug 2026 19:31:51 +0800 Subject: [PATCH 03/11] [Refactor][Backend] Allow resolved external builder params --- docs/design/ops-design-reference.md | 2 +- src/tileops/ops/op_base.py | 11 +++++++++-- tests/test_op_backend_seam.py | 14 ++++++++++++++ 3 files changed, 24 insertions(+), 3 deletions(-) diff --git a/docs/design/ops-design-reference.md b/docs/design/ops-design-reference.md index 4837c27ac..304d6583f 100644 --- a/docs/design/ops-design-reference.md +++ b/docs/design/ops-design-reference.md @@ -42,7 +42,7 @@ Rationale and the role / entry vocabulary: [ops-design.md § Kernel caching and | Method | Purpose | | -------------------------------------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `get_or_build_kernel(name, inputs, *, key, build)` | Return the kernel for this call, building it once on a miss. The only get-or-build in L1-L3. `key` and `build` are the in-tree recipe; `inputs` is what an external target's builder is described with, and an op that has not been wired to external targets yet omits it | +| `get_or_build_kernel(name, inputs, *, key, build, params)` | Return the kernel for this call, building it once on a miss. `key` and `build` are the in-tree recipe; `inputs` describes an external target's builder. Optional `params` replaces unresolved manifest params for that external build and must be deterministic for the op instance and input signature because it is not part of the cache key | | `built_kernels(name)` | Read-only view of a name's entries; empty before its first build. Introspection only, never dispatch | | `kernel_delegates()` | The ops whose kernels this op runs. Default `()`; a composite op overrides it | | `iter_kernels()` | Every `Kernel` the op holds, deduplicated: entries and delegates | diff --git a/src/tileops/ops/op_base.py b/src/tileops/ops/op_base.py index 98928686d..4a51d53c1 100644 --- a/src/tileops/ops/op_base.py +++ b/src/tileops/ops/op_base.py @@ -317,6 +317,7 @@ def get_or_build_kernel( *, key: Hashable = None, build: Optional[Callable[[], _Entry]] = None, + params: Optional[Mapping[str, object]] = None, ) -> _Entry: """Return the kernel for this call, building it once on a miss. @@ -335,6 +336,9 @@ def get_or_build_kernel( path keys on the input signature instead. build: How the *in-tree* kernel is constructed, called once per key. See ``Op._entry_kernels`` for what it may return. + params: Manifest parameters resolved for this input signature. The in-tree + path ignores them; an external builder receives these values instead of + the constructor attributes returned by ``_manifest_params()``. Returns: The stored entry, identical across calls describing the same specialization. @@ -397,7 +401,7 @@ def get_or_build_kernel( None if spec is None else (spec.dtype, spec.shape) for spec in specs ) if signature not in entries: - entries[signature] = self._build_external(builder, name, specs) + entries[signature] = self._build_external(builder, name, specs, params=params) return entries[signature] except Exception: # Whoever settled it unsettles it. ``__call__``'s handler does not run when @@ -411,13 +415,16 @@ def _build_external( builder: BuildKernel, name: str, specs: "tuple[TensorSpec | None, ...]", + *, + params: Optional[Mapping[str, object]] = None, ) -> object: """Ask the target for a kernel and hold it to the one rule this boundary has. *specs* carries one slot per ``signature.inputs`` entry; an absent optional input's slot is ``None``. """ - kernel = builder(*specs, **self._manifest_params()) + manifest_params = dict(params) if params is not None else self._manifest_params() + kernel = builder(*specs, **manifest_params) if not callable(kernel): raise OpNotAvailableError( f"target {self._settled_target!r} built {kernel!r} for " diff --git a/tests/test_op_backend_seam.py b/tests/test_op_backend_seam.py index 0122b1911..85014f265 100644 --- a/tests/test_op_backend_seam.py +++ b/tests/test_op_backend_seam.py @@ -95,6 +95,20 @@ def test_a_target_takes_over_the_op_and_is_asked_with_the_manifest_signature(): assert recorder.calls[1][1]["eps"] == 1e-5 +def test_a_callsite_can_resolve_external_builder_params(): + recorder = _Recorder() + _register(recorder) + x, weight = _inputs() + op = RMSNormFwdOp(normalized_shape=NORMALIZED_SHAPE) + resolved = {"normalized_shape": NORMALIZED_SHAPE, "eps": 2e-5} + + kernel = op.get_or_build_kernel("rms_norm", (x, weight), params=resolved) + kernel(x, weight) + + ((_, params),) = recorder.calls + assert params == resolved + + def test_the_op_layer_still_does_its_half(): """A backend writes kernels, not ops: validation and normalization are not its job.""" recorder = _Recorder() From 883c2b5ecf6380bebef106a29afa003626d6599c Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 15:11:23 +0800 Subject: [PATCH 04/11] [Refactor][Attention] Resolve dense GQA builder inputs --- src/tileops/ops/attention/gqa.py | 18 ++++++++--- tests/ops/attention/test_gqa.py | 11 ++----- tests/test_op_backend_seam.py | 53 ++++++++++++++++++++++++++++++++ 3 files changed, 69 insertions(+), 13 deletions(-) diff --git a/src/tileops/ops/attention/gqa.py b/src/tileops/ops/attention/gqa.py index b22c495b2..e5ab544ca 100644 --- a/src/tileops/ops/attention/gqa.py +++ b/src/tileops/ops/attention/gqa.py @@ -196,9 +196,8 @@ def __init__( self.is_causal = is_causal self.sm_scale = sm_scale - # This default is shape-independent, so targets receive one normalized - # convention. ``sm_scale`` stays None so each implementation resolves - # its shape-dependent default from the call's D. + # Normalize the shape-independent default now. ``sm_scale`` is resolved + # from the current call's D when this Op asks for an implementation. self.softcap = _score_softcap(softcap) if window_size_left < -1: raise ValueError("window_size_left must be -1 (unlimited) or >= 0") @@ -364,9 +363,20 @@ def _get_kernel( self, inputs: tuple[Optional[torch.Tensor], ...] ) -> Callable[..., torch.Tensor]: """Resolve the implementation stored in the Op's single cache layer.""" + q = inputs[0] + assert q is not None + dim = q.shape[-1] + params = self._manifest_params() + params.update( + sm_scale=_attention_scale(dim, self.sm_scale), + dtype=self.dtype or q.dtype, + rotary_dim=( + _rope_rotary_dim(dim, self.rotary_dim) if self.pos_encoding_mode == "rope" else None + ), + ) # BUILTIN follow-up: pass a shape/dtype ``key`` and a ``build`` closure # that selects and constructs one concrete kernel on a cache miss. - return self.get_or_build_kernel("gqa_dense", inputs) + return self.get_or_build_kernel("gqa_dense", inputs, params=params) def forward( self, diff --git a/tests/ops/attention/test_gqa.py b/tests/ops/attention/test_gqa.py index 116f32573..8e95293f2 100644 --- a/tests/ops/attention/test_gqa.py +++ b/tests/ops/attention/test_gqa.py @@ -37,10 +37,7 @@ class _DenseBoundaryTestOp(GroupedQueryAttentionDenseFwdOp): """Exercise the public boundary without adding an in-tree kernel.""" - resolved_inputs: tuple[Optional[torch.Tensor], ...] - def _get_kernel(self, inputs): - self.resolved_inputs = inputs return lambda *args: args[0] @@ -62,13 +59,9 @@ def test_dense_gqa_rejects_non_bshd_inputs(name: str) -> None: @pytest.mark.smoke -def test_dense_gqa_accepts_rectangular_decode_and_keeps_optional_slots() -> None: +def test_dense_gqa_accepts_rectangular_decode() -> None: q, k, v = _dense_boundary_inputs(seq_len_q=1, seq_len_kv=4) - op = _DenseBoundaryTestOp(is_causal=True) - - assert op(q, k, v).shape == q.shape - assert len(op.resolved_inputs) == 8 - assert op.resolved_inputs[3:] == (None, None, None, None, None) + assert _DenseBoundaryTestOp(is_causal=True)(q, k, v).shape == q.shape @pytest.mark.smoke diff --git a/tests/test_op_backend_seam.py b/tests/test_op_backend_seam.py index 85014f265..a7f1302cd 100644 --- a/tests/test_op_backend_seam.py +++ b/tests/test_op_backend_seam.py @@ -10,6 +10,7 @@ import torch from tileops.backend import BUILTIN, OpNotAvailableError, TensorSpec, registry +from tileops.ops.attention.gqa import GroupedQueryAttentionDenseFwdOp from tileops.ops.convolution import Conv2dFwdOp from tileops.ops.norm.rms_norm import RMSNormFwdOp from tileops.ops.pool import MaxPool2dFwdOp @@ -51,6 +52,21 @@ def kernel(x, weight): return kernel +class _DenseRecorder: + def __init__(self): + self.calls = [] + self.kernel_calls = [] + + def build_kernel(self, *inputs, **params): + self.calls.append((inputs, params)) + + def kernel(*runtime_inputs): + self.kernel_calls.append(runtime_inputs) + return runtime_inputs[0] + + return kernel + + def _register(recorder, target="acme", op="RMSNormFwdOp", claims=True): registry.register_detector(target, lambda device: claims) registry.register_kernel_builder(op, target, recorder.build_kernel) @@ -109,6 +125,43 @@ def test_a_callsite_can_resolve_external_builder_params(): assert params == resolved +def test_dense_gqa_hands_one_resolved_signature_to_one_cached_target_kernel(): + recorder = _DenseRecorder() + _register(recorder, op="GroupedQueryAttentionDenseFwdOp") + q = torch.randn(1, 1, 4, 8, dtype=DTYPE) + k = torch.randn(1, 4, 2, 8, dtype=DTYPE) + v = torch.randn_like(k) + op = GroupedQueryAttentionDenseFwdOp() + + op(q, k, v) + op(q.clone(), k.clone(), v.clone()) + + ((inputs, params),) = recorder.calls + assert inputs == ( + TensorSpec.of(q), + TensorSpec.of(k), + TensorSpec.of(v), + None, + None, + None, + None, + None, + ) + assert params.pop("sm_scale") == pytest.approx(8**-0.5) + assert params == { + "is_causal": True, + "softcap": 0.0, + "window_size_left": -1, + "window_size_right": -1, + "dtype": DTYPE, + "pos_encoding_mode": "none", + "rotary_dim": None, + "rope_layout": "neox", + } + assert len(recorder.kernel_calls) == 2 + assert recorder.kernel_calls[0][3:] == (None, None, None, None, None) + + def test_the_op_layer_still_does_its_half(): """A backend writes kernels, not ops: validation and normalization are not its job.""" recorder = _Recorder() From a62a6c81f4c4102c76519eb8c743da84a3722e6b Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 15:17:05 +0800 Subject: [PATCH 05/11] [Doc][Ops] Format kernel cache reference table --- docs/design/ops-design-reference.md | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/docs/design/ops-design-reference.md b/docs/design/ops-design-reference.md index 304d6583f..584fb501d 100644 --- a/docs/design/ops-design-reference.md +++ b/docs/design/ops-design-reference.md @@ -40,13 +40,13 @@ Abstract interface: `default_kernel_map` (property), `forward()`. Manifest-drive Rationale and the role / entry vocabulary: [ops-design.md § Kernel caching and enumeration](ops-design.md#kernel-caching-and-enumeration). -| Method | Purpose | -| -------------------------------------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| Method | Purpose | +| ---------------------------------------------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | | `get_or_build_kernel(name, inputs, *, key, build, params)` | Return the kernel for this call, building it once on a miss. `key` and `build` are the in-tree recipe; `inputs` describes an external target's builder. Optional `params` replaces unresolved manifest params for that external build and must be deterministic for the op instance and input signature because it is not part of the cache key | -| `built_kernels(name)` | Read-only view of a name's entries; empty before its first build. Introspection only, never dispatch | -| `kernel_delegates()` | The ops whose kernels this op runs. Default `()`; a composite op overrides it | -| `iter_kernels()` | Every `Kernel` the op holds, deduplicated: entries and delegates | -| `autotune()` | Puts the op in tuned mode: tunes built kernels, and sets `tune` so later builds tune too | +| `built_kernels(name)` | Read-only view of a name's entries; empty before its first build. Introspection only, never dispatch | +| `kernel_delegates()` | The ops whose kernels this op runs. Default `()`; a composite op overrides it | +| `iter_kernels()` | Every `Kernel` the op holds, deduplicated: entries and delegates | +| `autotune()` | Puts the op in tuned mode: tunes built kernels, and sets `tune` so later builds tune too | ### `Kernel` base class attributes ([`src/tileops/kernels/kernel_base.py`](../../src/tileops/kernels/kernel_base.py)) From a2ff6afdfde82d54fa7770bb93f702882a25d8c6 Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 15:31:42 +0800 Subject: [PATCH 06/11] [Refactor][Attention] Narrow dense GQA boundary tests --- docs/design/ops-design-reference.md | 14 ++++---- tests/test_op_backend_seam.py | 53 ----------------------------- 2 files changed, 7 insertions(+), 60 deletions(-) diff --git a/docs/design/ops-design-reference.md b/docs/design/ops-design-reference.md index 584fb501d..4837c27ac 100644 --- a/docs/design/ops-design-reference.md +++ b/docs/design/ops-design-reference.md @@ -40,13 +40,13 @@ Abstract interface: `default_kernel_map` (property), `forward()`. Manifest-drive Rationale and the role / entry vocabulary: [ops-design.md § Kernel caching and enumeration](ops-design.md#kernel-caching-and-enumeration). -| Method | Purpose | -| ---------------------------------------------------------- | ----------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | -| `get_or_build_kernel(name, inputs, *, key, build, params)` | Return the kernel for this call, building it once on a miss. `key` and `build` are the in-tree recipe; `inputs` describes an external target's builder. Optional `params` replaces unresolved manifest params for that external build and must be deterministic for the op instance and input signature because it is not part of the cache key | -| `built_kernels(name)` | Read-only view of a name's entries; empty before its first build. Introspection only, never dispatch | -| `kernel_delegates()` | The ops whose kernels this op runs. Default `()`; a composite op overrides it | -| `iter_kernels()` | Every `Kernel` the op holds, deduplicated: entries and delegates | -| `autotune()` | Puts the op in tuned mode: tunes built kernels, and sets `tune` so later builds tune too | +| Method | Purpose | +| -------------------------------------------------- | -------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------------- | +| `get_or_build_kernel(name, inputs, *, key, build)` | Return the kernel for this call, building it once on a miss. The only get-or-build in L1-L3. `key` and `build` are the in-tree recipe; `inputs` is what an external target's builder is described with, and an op that has not been wired to external targets yet omits it | +| `built_kernels(name)` | Read-only view of a name's entries; empty before its first build. Introspection only, never dispatch | +| `kernel_delegates()` | The ops whose kernels this op runs. Default `()`; a composite op overrides it | +| `iter_kernels()` | Every `Kernel` the op holds, deduplicated: entries and delegates | +| `autotune()` | Puts the op in tuned mode: tunes built kernels, and sets `tune` so later builds tune too | ### `Kernel` base class attributes ([`src/tileops/kernels/kernel_base.py`](../../src/tileops/kernels/kernel_base.py)) diff --git a/tests/test_op_backend_seam.py b/tests/test_op_backend_seam.py index a7f1302cd..85014f265 100644 --- a/tests/test_op_backend_seam.py +++ b/tests/test_op_backend_seam.py @@ -10,7 +10,6 @@ import torch from tileops.backend import BUILTIN, OpNotAvailableError, TensorSpec, registry -from tileops.ops.attention.gqa import GroupedQueryAttentionDenseFwdOp from tileops.ops.convolution import Conv2dFwdOp from tileops.ops.norm.rms_norm import RMSNormFwdOp from tileops.ops.pool import MaxPool2dFwdOp @@ -52,21 +51,6 @@ def kernel(x, weight): return kernel -class _DenseRecorder: - def __init__(self): - self.calls = [] - self.kernel_calls = [] - - def build_kernel(self, *inputs, **params): - self.calls.append((inputs, params)) - - def kernel(*runtime_inputs): - self.kernel_calls.append(runtime_inputs) - return runtime_inputs[0] - - return kernel - - def _register(recorder, target="acme", op="RMSNormFwdOp", claims=True): registry.register_detector(target, lambda device: claims) registry.register_kernel_builder(op, target, recorder.build_kernel) @@ -125,43 +109,6 @@ def test_a_callsite_can_resolve_external_builder_params(): assert params == resolved -def test_dense_gqa_hands_one_resolved_signature_to_one_cached_target_kernel(): - recorder = _DenseRecorder() - _register(recorder, op="GroupedQueryAttentionDenseFwdOp") - q = torch.randn(1, 1, 4, 8, dtype=DTYPE) - k = torch.randn(1, 4, 2, 8, dtype=DTYPE) - v = torch.randn_like(k) - op = GroupedQueryAttentionDenseFwdOp() - - op(q, k, v) - op(q.clone(), k.clone(), v.clone()) - - ((inputs, params),) = recorder.calls - assert inputs == ( - TensorSpec.of(q), - TensorSpec.of(k), - TensorSpec.of(v), - None, - None, - None, - None, - None, - ) - assert params.pop("sm_scale") == pytest.approx(8**-0.5) - assert params == { - "is_causal": True, - "softcap": 0.0, - "window_size_left": -1, - "window_size_right": -1, - "dtype": DTYPE, - "pos_encoding_mode": "none", - "rotary_dim": None, - "rope_layout": "neox", - } - assert len(recorder.kernel_calls) == 2 - assert recorder.kernel_calls[0][3:] == (None, None, None, None, None) - - def test_the_op_layer_still_does_its_half(): """A backend writes kernels, not ops: validation and normalization are not its job.""" recorder = _Recorder() From 82ba30b1b4c5539c474b225f3d991dbddf3c3ea0 Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 15:45:04 +0800 Subject: [PATCH 07/11] [Refactor][Attention] Keep dense boundary independent of Op base --- src/tileops/ops/attention/gqa.py | 13 +------------ src/tileops/ops/op_base.py | 11 ++--------- tests/test_op_backend_seam.py | 14 -------------- 3 files changed, 3 insertions(+), 35 deletions(-) diff --git a/src/tileops/ops/attention/gqa.py b/src/tileops/ops/attention/gqa.py index e5ab544ca..8c6e20ab3 100644 --- a/src/tileops/ops/attention/gqa.py +++ b/src/tileops/ops/attention/gqa.py @@ -363,20 +363,9 @@ def _get_kernel( self, inputs: tuple[Optional[torch.Tensor], ...] ) -> Callable[..., torch.Tensor]: """Resolve the implementation stored in the Op's single cache layer.""" - q = inputs[0] - assert q is not None - dim = q.shape[-1] - params = self._manifest_params() - params.update( - sm_scale=_attention_scale(dim, self.sm_scale), - dtype=self.dtype or q.dtype, - rotary_dim=( - _rope_rotary_dim(dim, self.rotary_dim) if self.pos_encoding_mode == "rope" else None - ), - ) # BUILTIN follow-up: pass a shape/dtype ``key`` and a ``build`` closure # that selects and constructs one concrete kernel on a cache miss. - return self.get_or_build_kernel("gqa_dense", inputs, params=params) + return self.get_or_build_kernel("gqa_dense", inputs) def forward( self, diff --git a/src/tileops/ops/op_base.py b/src/tileops/ops/op_base.py index 4a51d53c1..98928686d 100644 --- a/src/tileops/ops/op_base.py +++ b/src/tileops/ops/op_base.py @@ -317,7 +317,6 @@ def get_or_build_kernel( *, key: Hashable = None, build: Optional[Callable[[], _Entry]] = None, - params: Optional[Mapping[str, object]] = None, ) -> _Entry: """Return the kernel for this call, building it once on a miss. @@ -336,9 +335,6 @@ def get_or_build_kernel( path keys on the input signature instead. build: How the *in-tree* kernel is constructed, called once per key. See ``Op._entry_kernels`` for what it may return. - params: Manifest parameters resolved for this input signature. The in-tree - path ignores them; an external builder receives these values instead of - the constructor attributes returned by ``_manifest_params()``. Returns: The stored entry, identical across calls describing the same specialization. @@ -401,7 +397,7 @@ def get_or_build_kernel( None if spec is None else (spec.dtype, spec.shape) for spec in specs ) if signature not in entries: - entries[signature] = self._build_external(builder, name, specs, params=params) + entries[signature] = self._build_external(builder, name, specs) return entries[signature] except Exception: # Whoever settled it unsettles it. ``__call__``'s handler does not run when @@ -415,16 +411,13 @@ def _build_external( builder: BuildKernel, name: str, specs: "tuple[TensorSpec | None, ...]", - *, - params: Optional[Mapping[str, object]] = None, ) -> object: """Ask the target for a kernel and hold it to the one rule this boundary has. *specs* carries one slot per ``signature.inputs`` entry; an absent optional input's slot is ``None``. """ - manifest_params = dict(params) if params is not None else self._manifest_params() - kernel = builder(*specs, **manifest_params) + kernel = builder(*specs, **self._manifest_params()) if not callable(kernel): raise OpNotAvailableError( f"target {self._settled_target!r} built {kernel!r} for " diff --git a/tests/test_op_backend_seam.py b/tests/test_op_backend_seam.py index 85014f265..0122b1911 100644 --- a/tests/test_op_backend_seam.py +++ b/tests/test_op_backend_seam.py @@ -95,20 +95,6 @@ def test_a_target_takes_over_the_op_and_is_asked_with_the_manifest_signature(): assert recorder.calls[1][1]["eps"] == 1e-5 -def test_a_callsite_can_resolve_external_builder_params(): - recorder = _Recorder() - _register(recorder) - x, weight = _inputs() - op = RMSNormFwdOp(normalized_shape=NORMALIZED_SHAPE) - resolved = {"normalized_shape": NORMALIZED_SHAPE, "eps": 2e-5} - - kernel = op.get_or_build_kernel("rms_norm", (x, weight), params=resolved) - kernel(x, weight) - - ((_, params),) = recorder.calls - assert params == resolved - - def test_the_op_layer_still_does_its_half(): """A backend writes kernels, not ops: validation and normalization are not its job.""" recorder = _Recorder() From 1492a7b9b4d930bb6aad2ec5887c2bccad614e86 Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 16:40:12 +0800 Subject: [PATCH 08/11] [Test][Attention] Exercise dense validation without a fake kernel --- tests/ops/attention/test_gqa.py | 19 +++---------------- 1 file changed, 3 insertions(+), 16 deletions(-) diff --git a/tests/ops/attention/test_gqa.py b/tests/ops/attention/test_gqa.py index 8e95293f2..84bdce2d7 100644 --- a/tests/ops/attention/test_gqa.py +++ b/tests/ops/attention/test_gqa.py @@ -34,13 +34,6 @@ } -class _DenseBoundaryTestOp(GroupedQueryAttentionDenseFwdOp): - """Exercise the public boundary without adding an in-tree kernel.""" - - def _get_kernel(self, inputs): - return lambda *args: args[0] - - def _dense_boundary_inputs(seq_len_q: int = 1, seq_len_kv: int = 4): q = torch.randn(1, seq_len_q, 4, 8, dtype=torch.float16) k = torch.randn(1, seq_len_kv, 2, 8, dtype=torch.float16) @@ -55,13 +48,7 @@ def test_dense_gqa_rejects_non_bshd_inputs(name: str) -> None: tensors[index] = tensors[index].squeeze(0) with pytest.raises(ValueError, match=rf"{name} must be a rank-4 BSHD tensor"): - _DenseBoundaryTestOp()(*tensors) - - -@pytest.mark.smoke -def test_dense_gqa_accepts_rectangular_decode() -> None: - q, k, v = _dense_boundary_inputs(seq_len_q=1, seq_len_kv=4) - assert _DenseBoundaryTestOp(is_causal=True)(q, k, v).shape == q.shape + GroupedQueryAttentionDenseFwdOp()(*tensors) @pytest.mark.smoke @@ -78,7 +65,7 @@ def test_dense_gqa_rejects_bottom_right_modes_when_q_is_longer_than_kv( q, k, v = _dense_boundary_inputs(seq_len_q=4, seq_len_kv=1) with pytest.raises(ValueError, match=message): - _DenseBoundaryTestOp(**kwargs)(q, k, v) + GroupedQueryAttentionDenseFwdOp(**kwargs)(q, k, v) @pytest.mark.smoke @@ -87,7 +74,7 @@ def test_dense_gqa_rejects_rope_tables_shorter_than_the_kv_positions() -> None: rope = torch.randn(3, 4, dtype=torch.float16) with pytest.raises(ValueError, match=r"max_position >= 4"): - _DenseBoundaryTestOp( + GroupedQueryAttentionDenseFwdOp( is_causal=False, pos_encoding_mode="rope", rotary_dim=8, From bcaeef0ebe15d668b73d13f9be6494c37120a895 Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 20:20:29 +0800 Subject: [PATCH 09/11] [Test][Attention] Defer dense GQA runtime coverage --- tests/ops/attention/test_gqa.py | 48 --------------------------------- 1 file changed, 48 deletions(-) diff --git a/tests/ops/attention/test_gqa.py b/tests/ops/attention/test_gqa.py index 84bdce2d7..fe077673d 100644 --- a/tests/ops/attention/test_gqa.py +++ b/tests/ops/attention/test_gqa.py @@ -14,7 +14,6 @@ ) from tileops.ops import ( GroupedQueryAttentionBwdOp, - GroupedQueryAttentionDenseFwdOp, GroupedQueryAttentionFwdOp, GroupedQueryAttentionPrefillFwdOp, GroupedQueryAttentionPrefillVarlenFwdOp, @@ -34,53 +33,6 @@ } -def _dense_boundary_inputs(seq_len_q: int = 1, seq_len_kv: int = 4): - q = torch.randn(1, seq_len_q, 4, 8, dtype=torch.float16) - k = torch.randn(1, seq_len_kv, 2, 8, dtype=torch.float16) - return q, k, torch.randn_like(k) - - -@pytest.mark.smoke -@pytest.mark.parametrize("name", ["q", "k", "v"]) -def test_dense_gqa_rejects_non_bshd_inputs(name: str) -> None: - tensors = list(_dense_boundary_inputs()) - index = ("q", "k", "v").index(name) - tensors[index] = tensors[index].squeeze(0) - - with pytest.raises(ValueError, match=rf"{name} must be a rank-4 BSHD tensor"): - GroupedQueryAttentionDenseFwdOp()(*tensors) - - -@pytest.mark.smoke -@pytest.mark.parametrize( - "kwargs, message", - [ - ({"is_causal": True}, "causal dense attention"), - ({"is_causal": False, "pos_encoding_mode": "rope"}, "fused RoPE"), - ], -) -def test_dense_gqa_rejects_bottom_right_modes_when_q_is_longer_than_kv( - kwargs: dict[str, object], message: str -) -> None: - q, k, v = _dense_boundary_inputs(seq_len_q=4, seq_len_kv=1) - - with pytest.raises(ValueError, match=message): - GroupedQueryAttentionDenseFwdOp(**kwargs)(q, k, v) - - -@pytest.mark.smoke -def test_dense_gqa_rejects_rope_tables_shorter_than_the_kv_positions() -> None: - q, k, v = _dense_boundary_inputs(seq_len_q=1, seq_len_kv=4) - rope = torch.randn(3, 4, dtype=torch.float16) - - with pytest.raises(ValueError, match=r"max_position >= 4"): - GroupedQueryAttentionDenseFwdOp( - is_causal=False, - pos_encoding_mode="rope", - rotary_dim=8, - )(q, k, v, rope_cos=rope, rope_sin=rope) - - def _selected_prefill_kernel_cls(op: GroupedQueryAttentionPrefillFwdOp) -> type: """Kernel class selection picks for a uniform, non-FP8 packed prefill call.""" call = op.attention_call(is_fp8=False, is_uniform=True) From f2a994e73ad7f86f5048446340be464d377e26bf Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 20:29:11 +0800 Subject: [PATCH 10/11] [Refactor][Attention] Clarify dense input canonicalization --- src/tileops/ops/attention/gqa.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/src/tileops/ops/attention/gqa.py b/src/tileops/ops/attention/gqa.py index 8c6e20ab3..17a8829d7 100644 --- a/src/tileops/ops/attention/gqa.py +++ b/src/tileops/ops/attention/gqa.py @@ -343,7 +343,7 @@ def _validate_forward_inputs( raise ValueError("rope_cos and rope_sin must have the same shape") @staticmethod - def _resolve_optional_inputs( + def _canonicalize_inputs( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, @@ -353,7 +353,7 @@ def _resolve_optional_inputs( rope_cos: Optional[torch.Tensor], rope_sin: Optional[torch.Tensor], ) -> tuple[Optional[torch.Tensor], ...]: - """Preserve the manifest's eight positional input slots.""" + """Return contiguous tensors in manifest order, preserving None slots.""" return tuple( tensor.contiguous() if tensor is not None else None for tensor in (q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) @@ -380,7 +380,7 @@ def forward( ) -> torch.Tensor: """Validate, normalize, resolve one concrete implementation, and run it.""" self._validate_forward_inputs(q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) - inputs = self._resolve_optional_inputs( + inputs = self._canonicalize_inputs( q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin ) kernel = self._get_kernel(inputs) From 192ec15513e3115e9c848d53bcf1b7ac20ede866 Mon Sep 17 00:00:00 2001 From: superAngGao Date: Tue, 25 Aug 2026 20:29:29 +0800 Subject: [PATCH 11/11] [Style][Attention] Format dense input canonicalization --- src/tileops/ops/attention/gqa.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/tileops/ops/attention/gqa.py b/src/tileops/ops/attention/gqa.py index 17a8829d7..2d12703b6 100644 --- a/src/tileops/ops/attention/gqa.py +++ b/src/tileops/ops/attention/gqa.py @@ -380,9 +380,7 @@ def forward( ) -> torch.Tensor: """Validate, normalize, resolve one concrete implementation, and run it.""" self._validate_forward_inputs(q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) - inputs = self._canonicalize_inputs( - q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin - ) + inputs = self._canonicalize_inputs(q, k, v, q_scale, k_scale, v_scale, rope_cos, rope_sin) kernel = self._get_kernel(inputs) return kernel(*inputs)