Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
69 changes: 69 additions & 0 deletions src/tileops/manifest/attention.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,75 @@ 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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

forward enforces two constraints missing from shape_rules:

  • is_causal requires S_q <= S_kv (gqa.py:291)
  • pos_encoding_mode == 'rope' requires S_q <= S_kv (gqa.py:293)

The manifest is the spec, so both belong here, plus the rule stating query-to-key alignment once it is decided.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done. The manifest now includes the causal and fused-RoPE S_q <= S_kv rules and documents bottom-right query-to-key alignment.

- "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)"
- "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
Expand Down
1 change: 1 addition & 0 deletions src/tileops/ops/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
GroupedQueryAttentionBwdOp,
GroupedQueryAttentionDecodePagedWithKVCacheFwdOp,
GroupedQueryAttentionDecodeWithKVCacheFwdOp,
GroupedQueryAttentionDenseFwdOp,
GroupedQueryAttentionFwdOp,
GroupedQueryAttentionPrefillFwdOp,
GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp,
Expand Down
2 changes: 2 additions & 0 deletions src/tileops/ops/attention/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
GroupedQueryAttentionBwdOp,
GroupedQueryAttentionDecodePagedWithKVCacheFwdOp,
GroupedQueryAttentionDecodeWithKVCacheFwdOp,
GroupedQueryAttentionDenseFwdOp,
GroupedQueryAttentionFwdOp,
GroupedQueryAttentionPrefillFwdOp,
GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp,
Expand All @@ -28,6 +29,7 @@
"GroupedQueryAttentionBwdOp",
"GroupedQueryAttentionDecodePagedWithKVCacheFwdOp",
"GroupedQueryAttentionDecodeWithKVCacheFwdOp",
"GroupedQueryAttentionDenseFwdOp",
"GroupedQueryAttentionFwdOp",
"GroupedQueryAttentionPrefillFwdOp",
"GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp",
Expand Down
259 changes: 258 additions & 1 deletion src/tileops/ops/attention/gqa.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
from typing import Dict, Optional
import math
from typing import Callable, Dict, Optional

import torch
import torch.nn.functional as F

from tileops.backend import Target
from tileops.kernels.attention import (
FlashAttnBwdPreprocessKernel,
GQABwdWgmmaPipelinedKernel,
Expand Down Expand Up @@ -40,6 +42,7 @@
"GroupedQueryAttentionBwdOp",
"GroupedQueryAttentionDecodePagedWithKVCacheFwdOp",
"GroupedQueryAttentionDecodeWithKVCacheFwdOp",
"GroupedQueryAttentionDenseFwdOp",
"GroupedQueryAttentionFwdOp",
"GroupedQueryAttentionPrefillFwdOp",
"GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp",
Expand Down Expand Up @@ -128,6 +131,260 @@ def _build_packed_prefill_kernel(
)


class GroupedQueryAttentionDenseFwdOp(Op):
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__(
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 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
# 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")
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(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Why do these two functions use only one parameter but pass so many unused ones?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The full signatures are intentional because they mirror the manifest input slots. _infer_output_shapes is required by the Op abstract contract and manifest parity checks; Dense GQA output shape depends only on q_shape. _validate_dtypes uses the tensor inputs to enforce Q/K/V, scale, and RoPE dtype relations.

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]:
"""Keep this spec-only Op concrete until its roofline is implemented."""
raise NotImplementedError("Dense GQA has no in-tree implementation yet")

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This override can go: Op.eval_roofline (op_base.py:191) already raises NotImplementedError, with a message pointing at roofline.md 4.4.6.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I retained this minimal override after checking the class contract: Op.eval_roofline is abstract, and roofline codegen intentionally skips status: spec-only. Removing the override makes the public class non-instantiable. The override can disappear once an implemented entry supplies generated roofline code.


def _validate_forward_inputs(
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],
) -> 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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

forward runs 26 branches; 21 of them are raise ValueError on the arguments.

  • Two sibling ops in this file already extract that: GroupedQueryAttentionPrefillVarlenFwdOp._validate_forward_inputs (L891) and GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp._validate_forward_inputs (L1178).
  • Here only dtype is extracted (_validate_dtypes); shape, device, and combination checks stay inline. Half-extracted reads worse than not extracting.
  • Suggested shape: _validate_forward_inputs + _resolve_optional_inputs returning the 8-tuple, leaving forward as unpack shapes -> validate -> resolve -> self._get_callable(inputs)(*inputs).

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done. Runtime checks now live in _validate_forward_inputs; contiguous conversion and fixed manifest-slot construction live in _canonicalize_inputs. forward is now validate -> canonicalize -> get kernel -> invoke.

batch_kv, seq_len_kv, heads_kv, dim_kv = k.shape

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

No ndim check on q / k / v.

  • A 3-D q fails here as ValueError: not enough values to unpack, which does not name the argument.
  • Both sibling _validate_forward_inputs methods check ndim explicitly.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done. q, k, and v now receive explicit rank-4 BSHD checks before shape unpacking, with the argument name in each error.

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")

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Blocking: seq_len_q < seq_len_kv is admitted, but nothing defines where query i sits on the key axis.

  • The existing square kernel masks with q_idx >= k_idx, i.e. top-left (kernels/attention/gqa_fwd.py:102).
  • Decode needs bottom-right (i + S_kv - S_q >= j), otherwise it attends to key 0 only.
  • The same offset decides which RoPE position q_i gets, and there is no position/offset input.
  • Neither the docstring nor shape_rules states it, so every target will pick its own. Please fix the alignment here.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fixed. The class contract now defines bottom-right alignment as p_i = i + S_kv - S_q; causal/window masking and fused-RoPE query positions are defined from p_i. The manifest records the same constraints.

if self.pos_encoding_mode == "rope" and seq_len_q > seq_len_kv:
raise ValueError("fused RoPE requires seq_len_q <= seq_len_kv")

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")
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)}")

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'")
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 _canonicalize_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], ...]:
"""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)
)

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._canonicalize_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):
"""Compatibility square GQA forward wrapper. Public layout: BSHD."""

Expand Down
Loading