For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo function
fp8_index_naive
def fp8_index_naive[dtype: DType, output_layout: TensorLayout, q_layout: TensorLayout, qs_layout: TensorLayout, k_layout: TensorLayout, ks_layout: TensorLayout, vl_layout: TensorLayout, cro_layout: TensorLayout, //, num_heads: Int, depth: Int](output: TileTensor[.float32, output_layout], q: TileTensor[dtype, q_layout], q_s: TileTensor[.float32, qs_layout], k: TileTensor[dtype, k_layout], k_s: TileTensor[.float32, ks_layout], valid_length: TileTensor[.uint32, vl_layout], cache_row_offsets: TileTensor[.uint32, cro_layout], batch_size: Int, max_seq_len: Int, max_num_keys: Int, ctx: DeviceContext)
Computes the FP8 index/gather score via a two-pass matmul-then-reduce reference path.
Enqueues _index_matmul_max to produce per-head logits followed by
_reduce_logits to sum across heads and apply the per-key scale, serving
as a correctness reference for the optimized tensor-core kernels.
Parameters:
- dtype (
DType): Data type of the query and key tensors. - output_layout (
TensorLayout): Layout of the output score tensor. - q_layout (
TensorLayout): Layout of the query tensor; its head and depth extents must be static and equalnum_headsanddepth. - qs_layout (
TensorLayout): Layout of the per-query scale tensor. - k_layout (
TensorLayout): Layout of the key tensor. - ks_layout (
TensorLayout): Layout of the per-key scale tensor. - vl_layout (
TensorLayout): Layout of the cumulative sequence offsets. - cro_layout (
TensorLayout): Layout of the cache row offsets. - num_heads (
Int): Number of attention heads. - depth (
Int): Per-head feature depth.
Args:
- output (
TileTensor[.float32, output_layout]): Output score tensor of shape[total_seq_len, max_num_keys]. - q (
TileTensor[dtype, q_layout]): Query tensor of shape[total_seq_len, num_heads, depth]. - q_s (
TileTensor[.float32, qs_layout]): Per-query scale tensor of shape[total_seq_len, num_heads]. - k (
TileTensor[dtype, k_layout]): Key tensor of shape[total_keys, 1, depth]. - k_s (
TileTensor[.float32, ks_layout]): Per-key scale tensor of shape[total_keys]. - valid_length (
TileTensor[.uint32, vl_layout]): Cumulative sequence offsets of shape[batch_size + 1]. - cache_row_offsets (
TileTensor[.uint32, cro_layout]): Per-batch row offsets into the paged key cache. - batch_size (
Int): Number of sequences in the batch. - max_seq_len (
Int): Maximum sequence length across the batch. - max_num_keys (
Int): Maximum key count across the batch. - ctx (
DeviceContext): Device context used to enqueue the kernels.
Raises:
When the underlying kernel enqueue reports a device-side error.