Attention Operators¶
Every op on this page is used the same way: construct it once, then call it. The
constructor takes what the kernel is compiled with; the call takes the tensors.
Both are documented under each op — __init__ and forward, where forward is
what runs when you call op(...).
Multi-head attention¶
tileops.ops.attention.mha.MultiHeadAttentionFwdOp
¶
Layout: BSHD.
MHA is the heads_kv == heads specialization of GQA, so route the maintained forward path through the GQA prefill dispatcher.
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
is_causal(bool, default:True) –Manifest
params.is_causal,bool, defaultTrue. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
tileops.ops.attention.mha.MultiHeadAttentionBwdOp
¶
Layout: BSHD.
MHA backward is the heads_kv == heads specialization of GQA backward,
matching the forward path's dispatch through GQA.
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
is_causal(bool, default:True) –Manifest
params.is_causal,bool, defaultTrue. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
k(Tensor) –Input tensor, dtype
same_as(q). -
v(Tensor) –Input tensor, dtype
same_as(q). -
o(Tensor) –Input tensor, dtype
same_as(q). -
do(Tensor) –Input tensor, dtype
same_as(q). -
lse(Tensor) –Input tensor, dtype
float32.
Returns:
-
tuple[Tensor, Tensor, Tensor]–dq,dk,dv, as the manifest declares. Shape rules:dq.shape == (B, S, H, D);dk.shape == (B, S, H, D);dv.shape == (B, S, H, D).
tileops.ops.attention.mha.MultiHeadAttentionDecodeWithKVCacheFwdOp
¶
Layout: BSHD
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
k(Tensor) –Input tensor, dtype
same_as(q). -
v(Tensor) –Input tensor, dtype
same_as(q).
Returns:
-
Tensor–o, as the manifest declares. Shape rules:o.shape == (B, S_q, H, D).
tileops.ops.attention.mha.MultiHeadAttentionDecodePagedWithKVCacheFwdOp
¶
Paged MHA decode with dynamic KV cache. Layout: Q \([batch \times seqlen\_q \times heads \times dim]\) (BSHD);
K, V physical cache [seqlen_kv, heads, dim]; real_seqlen_kv [batch]; block_table [batch, num_pages].
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
page_size(int) –Manifest
params.page_size,int. -
is_causal(bool, default:False) –Manifest
params.is_causal,bool, defaultFalse. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
k(Tensor) –Input tensor, dtype
same_as(q). -
v(Tensor) –Input tensor, dtype
same_as(q). -
real_seqlen_kv(Tensor) –Input tensor, dtype
int32. -
block_table(Tensor) –Input tensor, dtype
int32.
Returns:
-
Tensor–o, as the manifest declares. Shape rules:o.shape == (B, S_q, H, D).
Grouped-query attention¶
tileops.ops.attention.gqa.GroupedQueryAttentionFwdOp
¶
Compatibility square GQA forward wrapper. Public layout: BSHD.
__init__
¶
__init__(
batch,
heads,
heads_kv,
seq_len,
dim,
is_causal=True,
sm_scale=None,
softcap=None,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
is_causal(bool, default:True) –Manifest
params.is_causal,bool, defaultTrue. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
tileops.ops.attention.gqa.GroupedQueryAttentionBwdOp
¶
Layout: BSHD
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
is_causal(bool, default:True) –Manifest
params.is_causal,bool, defaultTrue. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
k(Tensor) –Input tensor, dtype
same_as(q). -
v(Tensor) –Input tensor, dtype
same_as(q). -
o(Tensor) –Input tensor, dtype
same_as(q). -
do(Tensor) –Input tensor, dtype
same_as(q). -
lse(Tensor) –Input tensor, dtype
float32.
Returns:
-
tuple[Tensor, Tensor, Tensor]–dq,dk,dv, as the manifest declares. Shape rules:dq.shape == (B, S, H, D);dk.shape == (B, S, H_kv, D);dv.shape == (B, S, H_kv, D).
tileops.ops.attention.gqa.GroupedQueryAttentionDenseFwdOp
¶
Shape-agnostic dense BSHD grouped-query attention for prefill and contiguous decode.
Let \(g = H / H_{kv}\) and \(r(h) = \lfloor h / g \rfloor\) map query head \(h\) to its KV head. Rectangular attention aligns the two sequences at the bottom right, placing query row \(i\) at key position
Causal attention admits key position \(j\) when \(j \le p_i\); a finite window additionally requires \(p_i - \mathrm{left} \le j \le p_i + \mathrm{right}\). Fused RoPE rotates Q at \(p_i\) and K at \(j\). Causal and fused-RoPE calls therefore require \(S_q \le S_{kv}\).
FP8 inputs are dequantized with per-KV-head scales; each query head uses the scale of its KV group:
With \(\alpha\) = sm_scale (default \(1 / \sqrt{D}\)), the raw score is
\(z = \alpha \, \hat q \cdot \hat k\); with softcap \(c > 0\) it becomes
\(c \tanh(z / c)\) before masking and softmax. Dot products, softmax, and
the weighted V reduction accumulate in FP32; the result is cast to the
configured output dtype.
__init__
¶
__init__(
is_causal=True,
sm_scale=None,
softcap=None,
window_size_left=-1,
window_size_right=-1,
dtype=None,
pos_encoding_mode="none",
rotary_dim=None,
rope_layout="neox",
*,
target=None
)
Build the op. Tensor shapes arrive with each forward call.
Parameters:
-
is_causal(bool, default:True) –Apply the causal mask, bottom-right aligned.
-
sm_scale(Optional[float], default:None) –Score scale \(\alpha\).
Noneresolves to \(1 / \sqrt{D}\) from the call's head dimension. -
softcap(Optional[float], default:None) –Positive cap \(c\) applying \(c \tanh(z / c)\) to raw scores.
Noneor0disables capping. -
window_size_left(int, default:-1) –Keys admitted left of \(p_i\);
-1means unlimited. -
window_size_right(int, default:-1) –Keys admitted right of \(p_i\);
-1means unlimited. -
dtype(Optional[dtype], default:None) –Output dtype. Required (
float16orbfloat16) for FP8 inputs;Noneoutputs the input dtype. -
pos_encoding_mode(str, default:'none') –"none", or"rope"to fuse the rotary embedding into attention. -
rotary_dim(Optional[int], default:None) –Rotated width of each head; even, at most \(D\), default the full head dimension. Valid only with
pos_encoding_mode="rope". -
rope_layout(str, default:'neox') –"neox"(rotate split halves) or"interleaved"(rotate adjacent pairs). -
target(Target, default:None) –Backend target to serve this op, or
Noneto decide from the input device.
Raises:
-
ValueError–A parameter is out of range, or the combination is inconsistent (e.g.
rotary_dimwithout RoPE).
forward
¶
Run dense GQA attention on one batch of BSHD tensors.
Parameters:
-
q(Tensor) –Queries, \([B \times S_q \times H \times D]\);
float16,bfloat16, orfloat8_e4m3fn. -
k(Tensor) –Keys, \([B \times S_{kv} \times H_{kv} \times D]\), same dtype as
q. -
v(Tensor) –Values, \([B \times S_{kv} \times H_{kv} \times D]\), same dtype as
q. -
q_scale(Optional[Tensor], default:None) –FP8 dequantization scales for
q, \([B \times H_{kv}]\),float32. The three scales are required together for FP8 input and invalid otherwise. -
k_scale(Optional[Tensor], default:None) –Scales for
k, \([B \times H_{kv}]\),float32. -
v_scale(Optional[Tensor], default:None) –Scales for
v, \([B \times H_{kv}]\),float32. -
rope_cos(Optional[Tensor], default:None) –RoPE cosine table, \([P \times d_r / 2]\) with \(P \ge S_{kv}\) and \(d_r\) =
rotary_dim, in the output dtype. The two tables are required together withpos_encoding_mode="rope"and invalid otherwise. -
rope_sin(Optional[Tensor], default:None) –RoPE sine table, same shape and dtype as
rope_cos.
Returns:
-
Tensor–Attention output, \([B \times S_q \times H \times D]\), in the
-
Tensor–configured output dtype.
Raises:
-
ValueError–Shapes, dtypes, devices, or optional-input combinations violate the contract above.
tileops.ops.attention.gqa.GroupedQueryAttentionPrefillFwdOp
¶
Canonical packed GQA prefill. Layout: THD.
Dense and square prefill are represented with uniform cu_seqlens. Ragged
prefill uses the same fixed public tensor list. Scale tensors are required
for manifest stability; non-FP8 kernels ignore them.
__init__
¶
__init__(
batch,
heads,
heads_kv,
dim,
max_seqlen_q,
max_seqlen_kv,
dtype=torch.float16,
is_causal=True,
sm_scale=None,
softcap=None,
window_size_left=-1,
window_size_right=-1,
backend="auto",
validate_uniform_cu_seqlens=True,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
dtype(dtype, default:float16) –Element type of
o. The inputs do not determine it: identical float8_e4m3fn q/k/v admit either a float16 or a bfloat16 output, so the caller chooses here. For float16 / bfloat16 inputs it must equal their element type.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16 | float8_e4m3fn. -
k(Tensor) –Input tensor, dtype
same_as(q). -
v(Tensor) –Input tensor, dtype
same_as(q). -
cu_seqlens_q(Tensor) –Input tensor, dtype
int32. -
cu_seqlens_kv(Tensor) –Input tensor, dtype
int32. -
q_scale(Tensor) –Input tensor, dtype
float32. -
k_scale(Tensor) –Input tensor, dtype
float32. -
v_scale(Tensor) –Input tensor, dtype
float32.
Returns:
-
Tensor–o, as the manifest declares. Shape rules:o.shape == (total_q, H, D).
tileops.ops.attention.gqa.GroupedQueryAttentionPrefillVarlenFwdOp
¶
Packed variable-length GQA prefill. Layout: THD.
cu_seqlens_q and cu_seqlens_kv describe packed per-request ranges.
Causal prefill uses bottom-right alignment for each request independently:
key position j is visible to query position i iff
j <= i + (kv_len - q_len).
__init__
¶
__init__(
batch,
heads,
heads_kv,
dim,
max_seqlen_q,
max_seqlen_kv,
is_causal=True,
sm_scale=None,
softcap=None,
validate_inputs=False,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on q, k, v, cu_seqlens_q and cu_seqlens_kv.
tileops.ops.attention.gqa.GroupedQueryAttentionPrefillPagedWithKVCacheFwdOp
¶
Packed GQA prefill with paged KV cache append. Layout: THD.
The current chunk is packed by request. cache_seqlens stores each
request's logical KV length before append. block_table maps logical
page ids to physical pages in k_pages / v_pages.
__init__
¶
__init__(
batch,
heads,
heads_kv,
max_pages_per_req,
page_size,
dim,
is_causal=True,
cache_dtype=None,
sm_scale=None,
softcap=None,
tune=False,
fuse_rope=False,
rope_base=10000.0,
max_position=None,
rotary_dim=None,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
max_pages_per_req(int) –Manifest
params.max_pages_per_req,int. -
page_size(int) –Manifest
params.page_size,int. -
is_causal(bool, default:True) –Manifest
params.is_causal,bool, defaultTrue. -
cache_dtype(Optional[dtype], default:None) –Manifest
params.cache_dtype,dtype | None, defaultNone. -
sm_scale(Optional[float], default:None) –Manifest
params.sm_scale,float | None, defaultNone. -
softcap(Optional[float], default:None) –Manifest
params.softcap,float | None, defaultNone. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
-
fuse_rope(bool, default:False) –Manifest
params.fuse_rope,bool, defaultFalse. -
rope_base(float, default:10000.0) –Manifest
params.rope_base,float, default10000.0. -
max_position(Optional[int], default:None) –Manifest
params.max_position,int | None, defaultNone. -
rotary_dim(Optional[int], default:None) –Manifest
params.rotary_dim,int | None, defaultNone.
forward
¶
forward(
q,
k_new,
v_new,
k_pages,
v_pages,
k_scale,
v_scale,
cu_seqlens_q,
cache_seqlens,
block_table,
max_seqlen_q,
)
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
k_new(Tensor) –Input tensor, dtype
same_as(q). -
v_new(Tensor) –Input tensor, dtype
same_as(q). -
k_pages(Tensor) –Input tensor, dtype
float16 | bfloat16 | float8_e4m3fn. -
v_pages(Tensor) –Input tensor, dtype
same_as(k_pages). -
k_scale(Tensor) –Input tensor, dtype
float32. -
v_scale(Tensor) –Input tensor, dtype
float32. -
cu_seqlens_q(Tensor) –Input tensor, dtype
int32. -
cache_seqlens(Tensor) –Input tensor, dtype
int32. -
block_table(Tensor) –Input tensor, dtype
int32.
Returns:
-
Tensor–o, as the manifest declares. Shape rules:o.shape == (total_q, H, D).
tileops.ops.attention.gqa.GroupedQueryAttentionDecodeWithKVCacheFwdOp
¶
Layout: BSHD
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
k(Tensor) –Input tensor, dtype
same_as(q). -
v(Tensor) –Input tensor, dtype
same_as(q).
Returns:
-
Tensor–o, as the manifest declares. Shape rules:o.shape == (B, H, D).
tileops.ops.attention.gqa.GroupedQueryAttentionDecodePagedWithKVCacheFwdOp
¶
Paged GQA decode with dynamic KV cache. Layout: Q \([batch \times heads \times dim]\) (BHD);
K, V physical cache [seqlen_kv, heads_kv, dim]; real_seqlen_kv [batch]; block_table [batch, num_pages].
__init__
¶
__init__(
batch,
heads,
heads_kv,
seqlen_kv,
dim,
page_size,
sm_scale=None,
softcap=None,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
page_size(int) –Manifest
params.page_size,int. -
sm_scale(Optional[float], default:None) –Manifest
params.sm_scale,float | None, defaultNone. -
softcap(Optional[float], default:None) –Manifest
params.softcap,float | None, defaultNone. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
k(Tensor) –Input tensor, dtype
same_as(q). -
v(Tensor) –Input tensor, dtype
same_as(q). -
real_seqlen_kv(Tensor) –Input tensor, dtype
int32. -
block_table(Tensor) –Input tensor, dtype
int32.
Returns:
-
Tensor–o, as the manifest declares. Shape rules:o.shape == (B, H, D).
tileops.ops.attention.gqa.GroupedQueryAttentionSlidingWindowFwdOp
¶
Fixed-length GQA forward with sliding window attention.
Token at q_pos attends to k_pos when ALL applicable conditions hold
- k_pos <= q_pos (is_causal=True)
- k_pos >= q_pos - window_size_left (window_size_left >= 0)
- k_pos <= q_pos + window_size_right (window_size_right >= 0)
Use window_size_left=-1 / window_size_right=-1 for no restriction.
__init__
¶
__init__(
batch,
heads,
heads_kv,
seq_len,
dim,
is_causal=True,
window_size_left=-1,
window_size_right=-1,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
batch(int) –Batch size.
-
heads(int) –Number of query heads.
-
heads_kv(int) –Number of KV heads (must divide heads evenly).
-
seq_len(int) –Sequence length (same for Q, K, V).
-
dim(int) –Head dimension.
-
is_causal(bool, default:True) –Whether to apply causal masking.
-
window_size_left(int, default:-1) –Left window size (-1 = unlimited).
-
window_size_right(int, default:-1) –Right window size (-1 = unlimited).
-
tune(bool, default:False) –Whether to run autotuning on kernel instantiation.
forward
¶
Run fixed-length GQA sliding window forward.
Parameters:
-
q(Tensor) –Query tensor, shape \([batch \times seq\_len \times heads \times dim]\).
-
k(Tensor) –Key tensor, shape \([batch \times seq\_len \times heads\_kv \times dim]\).
-
v(Tensor) –Value tensor, shape \([batch \times seq\_len \times heads\_kv \times dim]\).
Returns:
-
Tensor–Output tensor, shape \([batch \times seq\_len \times heads \times dim]\).
tileops.ops.attention.gqa.GroupedQueryAttentionSlidingWindowVarlenFwdOp
¶
Variable-length GQA forward with sliding window attention.
Inputs are packed (no padding); per-sample boundaries are given via cu_seqlens arrays. seqlen_q and seqlen_k may differ per sample:
offset = seqlen_k - seqlen_q (per sample, FA3 bottom-right convention)
A token at local q_pos attends to local k_pos when ALL conditions hold
k_pos <= q_pos + offset (is_causal=True) k_pos >= q_pos + offset - window_size_left (window_size_left >= 0) k_pos <= q_pos + offset + window_size_right (window_size_right >= 0)
__init__
¶
__init__(
batch,
heads,
heads_kv,
dim,
is_causal=True,
window_size_left=-1,
window_size_right=-1,
accum_dtype=torch.float32,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
batch(int) –Number of sequences in the batch.
-
heads(int) –Number of query heads.
-
heads_kv(int) –Number of KV heads (must divide heads evenly).
-
dim(int) –Head dimension.
-
is_causal(bool, default:True) –Whether to apply causal masking.
-
window_size_left(int, default:-1) –Left window size (-1 = unlimited).
-
window_size_right(int, default:-1) –Right window size (-1 = unlimited).
-
accum_dtype(dtype, default:float32) –Accumulator data type for intermediate computations.
-
tune(bool, default:False) –Whether to run autotuning on kernel instantiation.
forward
¶
Run variable-length GQA sliding window forward.
Parameters:
-
q(Tensor) –Query tensor, shape \([total\_q \times heads \times dim]\).
-
k(Tensor) –Key tensor, shape \([total\_k \times heads\_kv \times dim]\).
-
v(Tensor) –Value tensor, shape \([total\_k \times heads\_kv \times dim]\).
-
cu_seqlens_q(Tensor) –Cumulative Q lengths, shape \([batch+1]\), dtype int32.
-
cu_seqlens_k(Tensor) –Cumulative K lengths, shape \([batch+1]\), dtype int32.
-
max_seqlen_q(int) –Maximum Q sequence length across the batch.
Returns:
-
Tensor–Output tensor, shape \([total\_q \times heads \times dim]\).
Multi-head latent attention¶
tileops.ops.attention.deepseek_mla.MultiHeadLatentAttentionDecodeWithKVCacheFwdOp
¶
Layout: BSHD
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
pe_dim(int) –Manifest
params.pe_dim,int. -
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
q(Tensor) –Input tensor, dtype
float16 | bfloat16. -
q_pe(Tensor) –Input tensor, dtype
same_as(q). -
k(Tensor) –Input tensor, dtype
same_as(q). -
k_pe(Tensor) –Input tensor, dtype
same_as(q).
Returns:
-
Tensor–o, as the manifest declares. Shape rules:o.shape == (B, H, D).
Native sparse attention¶
tileops.ops.attention.deepseek_nsa.NSACmpFwdVarlenOp
¶
__init__
¶
__init__(
seq_num,
c_seq_len,
heads,
dim_k,
dim_v,
chunk_num,
group,
scale,
bc,
bs,
accum_dtype,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on q, k_cmp, v_cmp, offsets, chunk_offsets and token_indices.
tileops.ops.attention.deepseek_nsa.NSATopkVarlenOp
¶
__init__
¶
__init__(
seq_num,
c_seq_len,
heads,
dim,
chunk_num,
group,
scale,
selected_block_num,
bc,
bs,
accum_dtype,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on q, k_cmp, lse_in, offsets, chunk_offsets and token_indices.
tileops.ops.attention.deepseek_nsa.NSAFwdVarlenOp
¶
__init__
¶
__init__(
batch,
heads,
c_seq_len,
dim,
is_causal,
scale,
block_size,
groups,
selected_blocks,
accum_dtype,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
tune(bool, default:False) –Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on q, k, v, block_indices, block_counts, offsets and token_indices.
DeepSeek sparse attention¶
tileops.ops.attention.deepseek_dsa.DeepSeekSparseAttentionDecodeWithKVCacheFwdOp
¶
Sparse Attention Decode Operation with Key-Value Cache for DeepSeek.
This operation is part of a sparse attention mechanism, designed for use in decoding with key-value (KV) caching.
The layout of the operation is BSHD.
__init__
¶
__init__(
batch,
heads,
seq_len,
seq_len_kv,
dim,
dim_tail,
topk,
stride_kv,
heads_kv,
q_start_index_s,
sm_scale=None,
is_causal=True,
tune=False,
)
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
batch(int) –The batch size.
-
heads(int) –The number of attention heads.
-
seq_len(int) –The length of the input sequence.
-
seq_len_kv(int) –The length of the key-value sequence.
-
dim(int) –The dimension of the attention vectors.
-
dim_tail(int) –The dimension of the tail portion of the attention vectors.
-
topk(int) –The number of top elements to consider in sparse attention.
-
stride_kv(int) –The stride for the key-value sequence.
-
heads_kv(int) –The number of key-value heads.
-
q_start_index_s(int) –The start index for queries in the sequence.
-
sm_scale(Optional[float], default=None, default:None) –Scaling factor for the softmax function.
-
is_causal(bool, default=True, default:True) –Whether the attention is causal (True for causal, False for non-causal).
-
tune(bool, default=False, default:False) –Whether to enable kernel tuning.
forward
¶
Performs the forward pass of the sparse attention operation.
Parameters:
-
q(Tensor) –The query tensor with shape (batch, seq_len, heads, dim + dim_tail).
-
kv(Tensor) –The key-value tensor with shape (batch, seq_len_kv, heads_kv, dim + dim_tail).
-
indices(Tensor) –Indices tensor for sparse attention.
Returns:
-
Tensor–torch.Tensor: The result of applying the sparse attention operation on the input tensors.
Attention indexing¶
tileops.ops.fp8_lightning_indexer.FP8LightningIndexerFwdOp
¶
__init__
¶
Build the op. Shapes and dtype are taken from the first call.
Parameters:
-
clean_logits–Manifest
params.clean_logits,bool, defaultTrue. -
config(Optional[dict], default:None) –Manifest
params.config,dict | None, defaultNone. -
tune–Whether to autotune, applied when a kernel is first built.
forward
¶
Run the op on the inputs the manifest declares.
Parameters:
-
index_q(Tensor) –Input tensor, dtype
bfloat16 | float8_e4m3fn. -
index_k(Tensor) –Input tensor, dtype
bfloat16 | float8_e4m3fn. -
weights(Tensor) –Input tensor, dtype
float32. -
cu_seqlen_ks(Tensor) –Input tensor, dtype
int32. -
cu_seqlen_ke(Tensor) –Input tensor, dtype
int32. -
index_k_scale(Optional[Tensor], default:None) –Input tensor, dtype
float32. Optional.
Returns:
-
Tensor–logits, as the manifest declares.