IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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 struct

PagedKVCache

struct PagedKVCache[dtype_: DType, kv_params_: KVCacheStaticParams, page_size: Int, blocks_origin: MutOrigin, blocks_engine: TensorEngine, cache_lengths_origin: ImmOrigin, cache_lengths_engine_: TensorEngine, lookup_table_origin: ImmOrigin, lookup_table_engine: TensorEngine, scales_origin: MutOrigin, scales_engine: TensorEngine = DefaultEngine, *, scale_dtype_: Optional[DType] = None, quantization_granularity_: Int = Int(1)]

The PagedKVCache is a wrapper around the KVCache blocks for a given layer. It is used to access the KVCache blocks for PagedAttention.

Note: This struct represents a 4D view of a 6D PagedKVCacheCollection tensor. The compile-time layout has UNKNOWN_VALUE for stride[0] because the actual stride depends on num_layers from the parent tensor, which is only known at runtime. This ensures offset calculations use the correct runtime strides rather than incorrect compile-time values.

Parameters​

  • ​dtype_ (DType): The dtype of the kv-cache.
  • ​kv_params_ (KVCacheStaticParams): The kv-cache static parameters.
  • ​page_size (Int): The size of the page.
  • ​blocks_origin (MutOrigin): Origin of the KV cache blocks buffer.
  • ​blocks_engine (TensorEngine): Engine policy of the KV cache blocks buffer.
  • ​cache_lengths_origin (ImmOrigin): Origin of the cache lengths buffer.
  • ​cache_lengths_engine_ (TensorEngine): Engine policy of the cache lengths buffer.
  • ​lookup_table_origin (ImmOrigin): Origin of the lookup table buffer.
  • ​lookup_table_engine (TensorEngine): Engine policy of the lookup table buffer.
  • ​scales_origin (MutOrigin): Origin of the quantization scales buffer.
  • ​scales_engine (TensorEngine): Engine policy of the quantization scales buffer.
  • ​scale_dtype_ (Optional[DType]): Dtype of the quantization scales (if quantization enabled).
  • ​quantization_granularity_ (Int): Block size used for quantization (e.g. 128).

Fields​

  • ​blocks (PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].blocks_tt_type):
  • ​cache_lengths (PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].cache_lengths_tt_type):
  • ​lookup_table (PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].lookup_table_tt_type):
  • ​max_seq_length (UInt32):
  • ​max_cache_length (UInt32):
  • ​scales (OptionalReg[TileTensor[PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].scale_dtype, Layout[TypeList[Int64, ComptimeInt[page_size], ComptimeInt[kv_params_.num_heads], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)]](), TypeList[Int64, ComptimeInt[(kv_params_ * ceildiv(kv_params_.head_size, quantization_granularity_))], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)], ComptimeInt[Int(1)]]()], scales_origin, Engine=scales_engine]]):
  • ​scales_lookup_table (PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].lookup_table_tt_type):

Implemented traits​

AnyType, Copyable, Deinitable, DevicePassable, ImplicitlyCopyable, KVCacheT, Movable, RegisterPassable, TrivialRegisterPassable

comptime members​

blocks_layout​

comptime blocks_layout = Layout(PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].blocks_shape, PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].blocks_strides)

blocks_shape​

comptime blocks_shape = IntTuple(Int(-1), page_size, kv_params_, kv_params_)

blocks_strides​

comptime blocks_strides = IntTuple(Int(-1), Int((mul kv_params_.head_size, kv_params_.num_heads)), kv_params_, Int(1))

blocks_tt_layout​

comptime blocks_tt_layout = Layout[TypeList[Int64, ComptimeInt[page_size], ComptimeInt[kv_params_.num_heads], ComptimeInt[kv_params_.head_size]](), TypeList[Int64, ComptimeInt[(kv_params_ * kv_params_)], ComptimeInt[kv_params_.head_size], ComptimeInt[Int(1)]]()]

blocks_tt_type​

comptime blocks_tt_type = TileTensor[PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].dtype, Layout[TypeList[Int64, ComptimeInt[page_size], ComptimeInt[kv_params_.num_heads], ComptimeInt[kv_params_.head_size]](), TypeList[Int64, ComptimeInt[(kv_params_ * kv_params_)], ComptimeInt[kv_params_.head_size], ComptimeInt[Int(1)]]()], blocks_origin, Engine=blocks_engine]

cache_lengths_engine​

comptime cache_lengths_engine = cache_lengths_engine_

cache_lengths_tt_layout​

comptime cache_lengths_tt_layout = Layout[TypeList[Int64](), TypeList[ComptimeInt[Int(1)]]()]

cache_lengths_tt_type​

comptime cache_lengths_tt_type = TileTensor[.uint32, Layout[TypeList[Int64](), TypeList[ComptimeInt[Int(1)]]()], cache_lengths_origin, Engine=PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].cache_lengths_engine]

device_type​

comptime device_type = PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_]

dtype​

comptime dtype = dtype_

Engine​

comptime Engine = blocks_engine

head_dim_granularity​

comptime head_dim_granularity = ceildiv(kv_params_.head_size, PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].quantization_granularity)

kv_params​

comptime kv_params = kv_params_

lookup_table_tt_layout​

comptime lookup_table_tt_layout = Layout[TypeList[Int64, Int64](), TypeList[Int64, ComptimeInt[Int(1)]]()]

lookup_table_tt_type​

comptime lookup_table_tt_type = TileTensor[.uint32, Layout[TypeList[Int64, Int64](), TypeList[Int64, ComptimeInt[Int(1)]]()], lookup_table_origin, Engine=lookup_table_engine]

page_size_​

comptime page_size_ = page_size

quantization_enabled​

comptime quantization_enabled = (scale_dtype_ isnot NoneType(None))

quantization_granularity​

comptime quantization_granularity = quantization_granularity_

scale_dtype​

comptime scale_dtype = scale_dtype_.or_else(dtype_)

scales_tt_layout​

comptime scales_tt_layout = Layout[TypeList[Int64, ComptimeInt[page_size], ComptimeInt[kv_params_.num_heads], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)]](), TypeList[Int64, ComptimeInt[(kv_params_ * ceildiv(kv_params_.head_size, quantization_granularity_))], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)], ComptimeInt[Int(1)]]()]

scales_tt_type​

comptime scales_tt_type = TileTensor[PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].scale_dtype, Layout[TypeList[Int64, ComptimeInt[page_size], ComptimeInt[kv_params_.num_heads], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)]](), TypeList[Int64, ComptimeInt[(kv_params_ * ceildiv(kv_params_.head_size, quantization_granularity_))], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)], ComptimeInt[Int(1)]]()], scales_origin, Engine=scales_engine]

Methods​

__init__​

def __init__(blocks: TileTensor[Self.dtype, Layout[TypeList[Int64, ComptimeInt[page_size], ComptimeInt[kv_params_.num_heads], ComptimeInt[kv_params_.head_size]](), TypeList[Int64, ComptimeInt[(kv_params_ * kv_params_)], ComptimeInt[kv_params_.head_size], ComptimeInt[Int(1)]]()], blocks_origin, Engine=blocks_engine], cache_lengths: TileTensor[.uint32, Layout[TypeList[Int64](), TypeList[ComptimeInt[Int(1)]]()], cache_lengths_origin, Engine=Self.cache_lengths_engine], lookup_table: TileTensor[.uint32, Layout[TypeList[Int64, Int64](), TypeList[Int64, ComptimeInt[Int(1)]]()], lookup_table_origin, Engine=lookup_table_engine], max_seq_length: UInt32, max_cache_length: UInt32, scales: OptionalReg[TileTensor[Self.scale_dtype, Layout[TypeList[Int64, ComptimeInt[page_size], ComptimeInt[kv_params_.num_heads], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)]](), TypeList[Int64, ComptimeInt[(kv_params_ * ceildiv(kv_params_.head_size, quantization_granularity_))], ComptimeInt[ceildiv(kv_params_.head_size, quantization_granularity_)], ComptimeInt[Int(1)]]()], scales_origin, Engine=scales_engine]] = None, scales_lookup_table: OptionalReg[TileTensor[.uint32, Layout[TypeList[Int64, Int64](), TypeList[Int64, ComptimeInt[Int(1)]]()], lookup_table_origin, Engine=lookup_table_engine]] = None) -> Self

get_type_name​

static def get_type_name() -> String

Returns:

String

max_tile_size​

static def max_tile_size() -> Int

Returns the maximum tile size for the KVCache.

Returns:

Int

cache_lengths_nd​

def cache_lengths_nd(self) -> Self.cache_lengths_tt_type

Returns:

Self.cache_lengths_tt_type

cache_length​

def cache_length(self, batch_idx: Int) -> Int

Returns the length of the cache for a given batch index.

Returns:

Int

get_tma_row​

def get_tma_row(self, encoded_index: Int32) -> Int32

Convert an encoded sparse index to a physical TMA row.

The encoded index is physical_block * page_size + offset. This method decomposes it and returns physical_block * stride + offset where stride is the distance (in rows) between consecutive physical blocks in the flattened memory view.

Returns:

Int32

num_kv_rows​

def num_kv_rows(self) -> Int

Returns the total number of virtual rows in this KV cache view.

Returns:

Int

num_scale_rows​

def num_scale_rows(self) -> Int

Total virtual rows in the scale pool, as num_kv_rows is for values.

Returns:

Int

scale_row_idx​

def scale_row_idx(self, batch_idx: UInt32, start_tok_idx: UInt32) -> UInt32

Returns the row idx of a token's scale in the scale pool.

NOT row_idx: the block comes from scales_lookup_table and the stride from the scale pool, both of which may differ from the value pool's -- see the scales_lookup_table field comment.

Returns:

UInt32

scale_tma_coords​

def scale_tma_coords(self, batch_idx: UInt32, start_tok_idx: UInt32) -> Tuple[Int32, Int32]

The (row_in_block, block) coordinate of a token's scale tile.

What :func:create_paged_scale_tma_tile is addressed by. Split rather than folded into one row index because a fold is bounded by the whole pool and a TMA coordinate is signed 32-bit; see that function.

Args:

  • ​batch_idx (UInt32): Batch entry whose scales lookup table is read.
  • ​start_tok_idx (UInt32): First token of the tile, within the entry.

Returns:

Tuple[Int32, Int32]: The row within the block, and the block.

kv_tma_coords​

def kv_tma_coords(self, batch_idx: UInt32, tok_idx: UInt32) -> Tuple[Int32, Int32]

The (row_in_block, block) coordinate of a token's KV tile.

What :meth:create_paged_tma_tile is addressed by, and the same split :meth:scale_tma_coords makes for the scale pool. :meth:row_idx folds these two into block * stride + row, and that product is bounded by the whole slab rather than by this leaf's share of it: a shared-slab allocator hands a small-page leaf millions of blocks, and a TMA coordinate is signed 32-bit. Split, each coordinate is bounded by something that does not track total cache memory -- the block count, and a block's own rows.

Args:

  • ​batch_idx (UInt32): Batch entry whose lookup table is read.
  • ​tok_idx (UInt32): Token within the entry.

Returns:

Tuple[Int32, Int32]: The row within the block, and the block.

row_idx​

def row_idx(self, batch_idx: UInt32, tok_idx: UInt32) -> UInt32

Returns the row idx when viewing the memory as a matrix.

Returns:

UInt32

populate​

def populate[BN: Int, base_alignment: Int, pair_cta: Bool = False, is_leader: Bool = True](self, batch_idx: UInt32, base_kv_row: UInt32) -> PagedRowIndices[BN, Self.page_size_, pair_cta, is_leader]

SIMD LUT-load the num_pages block indices in one shot.

Computes `result.rows[i] = lookup_table[batch, first_lut_idx+i]

  • stride + tok_in_blockfor allnum_pagesentries using one (or a small fixed number of) alignedld.global.v{N}.u32` loads from the lookup table row.

Invariants:

  • self.lookup_table.dim[1] is large enough that a SIMD read of num_pages uint32s starting at any valid first_lut_idx stays in bounds (see PagedKVCacheManager for the allocation-side padding).
  • base_kv_row % base_alignment == 0 holds at runtime (typically mask.start_column_alignment[...]()). For num_pages > 1, base_alignment must be at least page_size, required so tok_in_block_idx == 0 and the SIMD multiply-add collapses to a multiply. Larger base_alignment values let us pick a wider SIMD chunk (chunk * page_size must divide base_alignment).

The per-load width chunk is the largest power of two that divides both num_pages and base_alignment / page_size, capped at 8. With base_alignment == BN (the historical contract), this matches the previous behaviour: chunk = min(num_pages & -num_pages, 8). With looser alignments (e.g. ChunkedMask providing only page_size alignment when BN > page_size), the chunk degrades to 1 (scalar loads).

Parameters:

  • ​BN (Int): Tile row count of the V sub-tile to populate indices for.
  • ​base_alignment (Int): Comptime promise that base_kv_row % base_alignment == 0 at runtime; must be at least page_size when num_pages > 1. Larger values enable wider SIMD LUT loads.
  • ​pair_cta (Bool): Whether this CTA is one of a pair sharing the K tile (defaults to False).
  • ​is_leader (Bool): When pair_cta is True, whether this CTA is the leader half (defaults to True).

Args:

  • ​batch_idx (UInt32): Index of the request in the batch.
  • ​base_kv_row (UInt32): Base virtual row of the BN-row tile; must satisfy base_kv_row % base_alignment == 0.

Returns:

PagedRowIndices[BN, Self.page_size_, pair_cta, is_leader]

create_tma_tile​

def create_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, BK: Int = padded_depth[dtype_, swizzle_mode, kv_params_.head_size](), fold_chunks: Int = Int(1), row_major: Bool = False](self, ctx: DeviceContext) -> TMATensorTile[Self.dtype, Int(3), _padded_shape[Int(3), Self.dtype, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode](), _ragged_shape[Int(3), Self.dtype, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode]()]

Creates a TMA tile for this KV cache.

Returns:

TMATensorTile[Self.dtype, Int(3), _padded_shape[Int(3), Self.dtype, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode](), _ragged_shape[Int(3), Self.dtype, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode]()]

create_paged_tma_tile​

def create_paged_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, BK: Int = padded_depth[dtype_, swizzle_mode, kv_params_.head_size]()](self, ctx: DeviceContext) -> TMATensorTile[Self.dtype, Int(4), Index[Int, Int, Int, Int](Int(1), BN, Int(1), BK)]

Builds a total_blocks x page_size x heads x depth KV descriptor.

The split counterpart of :meth:create_tma_tile, addressed by :meth:kv_tma_coords rather than by :meth:row_idx. The flat descriptor folds the block into the row coordinate, and that product spans the whole shared slab, so it leaves signed 32 bits while the slab is still an ordinary size. Here the block is its own coordinate and neither one tracks total cache memory. The scale pool retired the same ceiling the same way; see :func:create_paged_scale_tma_tile.

page_size is the extent rather than the block pitch: only a page of a block's num_layers * page_size rows belong to this layer, and they are contiguous, so the layer rides in ptr and the rest is stride. Declaring the pitch instead would run the last block's extent past the allocation by the layer's own offset.

A tile never straddles a block, which is what lets the block be constant across one copy: the consumers assert page_size == 0 or page_size % BN == 0 and route any other page to a scalar kernel.

Parameters:

  • ​swizzle_mode (TensorMapSwizzle): TMA swizzle for the innermost dimension.
  • ​BN (Int): Rows per tile.
  • ​BK (Int): Depth per tile, padded to the swizzle granularity.

Args:

  • ​ctx (DeviceContext): Device context used to create the TMA descriptor.

Returns:

TMATensorTile[Self.dtype, Int(4), Index[Int, Int, Int, Int](Int(1), BN, Int(1), BK)]: The TMA descriptor.

create_index_scale_tma_tile​

def create_index_scale_tma_tile[TILE: Int](self, ctx: DeviceContext) -> TMATensorTile[Self.scale_dtype, Int(2), Index[Int, Int](Int(1), flat_scale_window[scale_dtype_.or_else(dtype_), TILE]())]

Creates a flat TMA descriptor over the scale pool.

Returns:

TMATensorTile[Self.scale_dtype, Int(2), Index[Int, Int](Int(1), flat_scale_window[scale_dtype_.or_else(dtype_), TILE]())]

create_gather4_tma_tile​

def create_gather4_tma_tile[*, tile_height: Int = Int(4), tile_width: Int, tile_stride: Int = tile_width, swizzle_mode: TensorMapSwizzle = TensorMapSwizzle.SWIZZLE_NONE, tma_dtype: DType = PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].dtype, l2_promotion: TensorMapL2Promotion = TensorMapL2Promotion.NONE](self, ctx: DeviceContext) -> TMATensorTile[tma_dtype, Int(2), IndexList(tile_height, _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]

Creates a 2D TMA gather4 descriptor for this KV cache.

The descriptor views the KV cache as a flat 2D matrix of [num_kv_rows, tile_width] and is configured for gather4 operations that load 4 non-contiguous rows per TMA instruction. The box width is derived from the swizzle mode; for SWIZZLE_NONE it equals tile_width.

When tma_dtype differs from Self.dtype, the underlying data pointer is bitcast to tma_dtype at descriptor creation time.

Parameters:

  • ​tile_height (Int): Number of rows in the tile. Must be a multiple of 4. Defaults to 4 for backward compatibility.
  • ​tile_width (Int): Number of elements per row to load (box width) in tma_dtype elements.
  • ​tile_stride (Int): Row stride in elements in global memory. Defaults to tile_width. Use a larger value when the global row is wider than the portion to load.
  • ​swizzle_mode (TensorMapSwizzle): TMA swizzle mode for shared memory access pattern. Defaults to SWIZZLE_NONE.
  • ​tma_dtype (DType): The data type used for the TMA descriptor. Defaults to Self.dtype. When different, the pointer is bitcast.
  • ​l2_promotion (TensorMapL2Promotion): L2 cache promotion hint for TMA loads. Defaults to NONE.

Args:

  • ​ctx (DeviceContext): The CUDA device context used to create the TMA descriptor.

Returns:

TMATensorTile[tma_dtype, Int(2), IndexList(tile_height, _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[tma_dtype, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]: A TMATensorTile with box width derived from the swizzle mode.

create_rope_tma_tile​

def create_rope_tma_tile[swizzle_mode: TensorMapSwizzle, *, BN: Int, BK: Int, padded_depth: Int](self, ctx: DeviceContext, out tma: TMATensorTile[.bfloat16, Int(3), _padded_shape[Int(3), DType.bfloat16, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode](), _ragged_shape[Int(3), DType.bfloat16, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode]()])

Creates a BF16 TMA tile for the rope portion of the per-tensor rope-aware KV cache.

In the per-tensor rope-aware layout each token row is: padded_depth FP8 bytes (content) | BK BF16 elements (rope) Total row bytes = padded_depth + BK * 2.

The TMA descriptor points at the rope data by offsetting blocks.ptr by padded_depth bytes, then reinterpreting as BF16. The global memory stride dimension (last dim of gmem_shape) is the total row size expressed in BF16 units: (padded_depth + BK * 2) // 2.

Returns:

TMATensorTile[.bfloat16, Int(3), _padded_shape[Int(3), DType.bfloat16, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode](), _ragged_shape[Int(3), DType.bfloat16, IndexList(BN, Int(1), BK, __list_literal__=NoneType(None)), swizzle_mode]()]

create_rope_gather4_tma_tile​

def create_rope_gather4_tma_tile[*, tile_height: Int = Int(4), tile_width: Int, padded_depth: Int, swizzle_mode: TensorMapSwizzle = TensorMapSwizzle.SWIZZLE_NONE, l2_promotion: TensorMapL2Promotion = TensorMapL2Promotion.NONE](self, ctx: DeviceContext) -> TMATensorTile[.bfloat16, Int(2), IndexList(tile_height, _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]

Creates a BF16 gather4 TMA descriptor for the rope portion of the KV cache.

For the per-tensor rope-aware layout each token row is stored as padded_depth FP8 bytes (content) followed by BF16 rope elements. The total row width in BF16 units is (padded_depth + tile_width * 2) // 2.

This method offsets blocks.ptr by padded_depth bytes, reinterprets as BF16, and creates a gather4 TMA descriptor whose row stride is the full row width in BF16 elements.

Returns:

TMATensorTile[.bfloat16, Int(2), IndexList(tile_height, _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None)), IndexList(Int(1), _gather4_box_width[DType.bfloat16, tile_width, swizzle_mode](), __list_literal__=NoneType(None))]

load​

def load[width: Int, output_dtype: DType = PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].dtype](self, bs: Int, head_idx: Int, tok_idx: Int, head_dim_idx: Int) -> SIMD[output_dtype, width]

Loads an element from the given index.

Returns:

SIMD[output_dtype, width]

store​

def store(self, bs: Int, head_idx: Int, tok_idx: Int, head_dim_idx: Int, val: SIMD[Self.dtype])

Stores an element at the given index.

Skips the write when the LUT entry for (bs, tok_idx // page_size) is the unassigned-slot sentinel, i.e. when the resolved block_idx is outside [0, total_num_blocks). The cache manager fills LUT columns past a request's allocated block count with the sentinel value total_num_pages (see cache_manager.py's lut_table_np.fill(self._total_num_pages)) so that SIMD over-reads of the LUT row are safe, but the value of the sentinel times the page stride lands one page past the end of the cache buffer. Without this guard a sentinel-resolved store corrupts whatever device allocation happens to sit immediately after the KV cache.

load_scale​

def load_scale[width: Int](self, bs: Int, head_idx: Int, tok_idx: Int, head_dim_idx: Int) -> SIMD[Self.scale_dtype, width]

Loads a quantization scale from the given index.

Parameters:

  • ​width (Int): SIMD vector width of the returned scale values in elements.

Args:

  • ​bs (Int): Index of the request in the batch, in [0, num_requests).
  • ​head_idx (Int): Attention head index in [0, kv_params.num_heads).
  • ​tok_idx (Int): Token position within the request's sequence.
  • ​head_dim_idx (Int): Starting element offset within the head dimension, in [0, kv_params.head_size); the scale slot is head_dim_idx // quantization_granularity.

Returns:

SIMD[Self.scale_dtype, width]

store_scale​

def store_scale[scales_dtype: DType = PagedKVCache[dtype_, kv_params_, page_size, blocks_origin, blocks_engine, cache_lengths_origin, cache_lengths_engine_, lookup_table_origin, lookup_table_engine, scales_origin, scales_engine, scale_dtype_=scale_dtype_, quantization_granularity_=quantization_granularity_].scale_dtype, width: Int = Int(1)](self, bs: Int, head_idx: Int, tok_idx: Int, head_dim_idx: Int, scales: SIMD[scales_dtype, width])

Stores the quantization scales at the given index.

load_quantized​

def load_quantized[width: Int](self, bs: Int, head_idx: Int, tok_idx: Int, head_dim_idx: Int) -> SIMD[Self.dtype, width]

Loads a quantized element from the given index.

Returns:

SIMD[Self.dtype, width]

empty_cache​

def empty_cache(self) -> Bool

Returns true if the cache_lengths for all requests is 0, false otherwise.

Returns:

Bool

max_prompt_length​

def max_prompt_length(self) -> UInt32

Returns the maximum sequence length across all batches of the current request.

Returns:

UInt32

max_context_length​

def max_context_length(self) -> UInt32

Returns the maximum cache length used across all batches of the current request.

Returns:

UInt32

block_paged_ptr​

def block_paged_ptr[tile_size: Int](self, batch_idx: Int, start_tok_idx: Int, head_idx: Int, head_dim_idx: Int = Int(0)) -> Pointer[Scalar[Self.dtype], MutAnyOrigin]

Returns:

Pointer[Scalar[Self.dtype], MutAnyOrigin]

block_paged_storage​

def block_paged_storage[tile_size: Int](self, batch_idx: Int, start_tok_idx: Int, head_idx: Int, head_dim_idx: Int = Int(0)) -> blocks_engine.StorageType[True, MutUnsafeAnyOrigin, Self.dtype, MutAnyOrigin, AddressSpace.GENERIC]

Offsets the blocks handle to a paged block through the engine.

The pointer-returning block_paged_ptr drops the engine, so callers that need an engine-carrying tile take this instead.

Parameters:

  • ​tile_size (Int): Tile size in rows used to compute the paged block.

Args:

  • ​batch_idx (Int): Batch index of the request.
  • ​start_tok_idx (Int): Starting token index within the batch.
  • ​head_idx (Int): KV head index.
  • ​head_dim_idx (Int): Index along the head dimension (defaults to 0).

Returns:

blocks_engine.StorageType[True, MutUnsafeAnyOrigin, Self.dtype, MutAnyOrigin, AddressSpace.GENERIC]: The blocks storage handle advanced to the paged block.

scales_block_paged_ptr​

def scales_block_paged_ptr(self, batch_idx: Int, start_tok_idx: Int, head_idx: Int, head_dim_idx: Int = Int(0)) -> Pointer[Scalar[Self.scale_dtype], MutAnyOrigin]

Returns a pointer to the scales block at the requested indices.

Returns:

Pointer[Scalar[Self.scale_dtype], MutAnyOrigin]

scales_raw_ptr​

def scales_raw_ptr(self) -> Pointer[Scalar[Self.scale_dtype], MutAnyOrigin]

Returns the base pointer to the scales tensor, or a dangling pointer if scales are not set.

Returns:

Pointer[Scalar[Self.scale_dtype], MutAnyOrigin]

Was this page helpful?