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
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:
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:
num_kv_rows
def num_kv_rows(self) -> Int
Returns the total number of virtual rows in this KV cache view.
Returns:
num_scale_rows
def num_scale_rows(self) -> Int
Total virtual rows in the scale pool, as num_kv_rows is for values.
Returns:
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:
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:
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_block
for 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 ofnum_pagesuint32s starting at any validfirst_lut_idxstays in bounds (seePagedKVCacheManagerfor the allocation-side padding).base_kv_row % base_alignment == 0holds at runtime (typicallymask.start_column_alignment[...]()). Fornum_pages > 1,base_alignmentmust be at leastpage_size, required sotok_in_block_idx == 0and the SIMDmultiply-addcollapses to amultiply. Largerbase_alignmentvalues let us pick a wider SIMD chunk (chunk * page_sizemust dividebase_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 thatbase_kv_row % base_alignment == 0at runtime; must be at leastpage_sizewhennum_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 toFalse). - is_leader (
Bool): Whenpair_ctaisTrue, whether this CTA is the leader half (defaults toTrue).
Args:
- batch_idx (
UInt32): Index of the request in the batch. - base_kv_row (
UInt32): Base virtual row of theBN-row tile; must satisfybase_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:
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:
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) intma_dtypeelements. - tile_stride (
Int): Row stride in elements in global memory. Defaults totile_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 toSelf.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:
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:
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:
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 ishead_dim_idx // quantization_granularity.
Returns:
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:
empty_cache
def empty_cache(self) -> Bool
Returns true if the cache_lengths for all requests is 0, false otherwise.
Returns:
max_prompt_length
def max_prompt_length(self) -> UInt32
Returns the maximum sequence length across all batches of the current request.
Returns:
max_context_length
def max_context_length(self) -> UInt32
Returns the maximum cache length used across all batches of the current request.
Returns:
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]