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

Python class

RecurrentStateParams

RecurrentStateParams​

class max.nn.kv_cache.RecurrentStateParams(regions, devices, data_parallel_degree=1, page_size=0, kv_connector_config=<factory>, enable_prefix_caching=False, enable_dp_cross_replica_prefix_copy=True, kv_hash_algo='ahash64', kv_hash_seed=None, speculative_method=None, num_draft_tokens=0)

source

Bases: CacheLeafParamInterface

A cache leaf whose entry is a state rather than a span of tokens.

One fixed-size value carrying every token before it, drawn from the same slab, prefix index and eviction order as the attention caches.

A cache may hold more than one – a speculative pair keeps a state for the target and one for the draft. They are told apart by their regions’ leaf ids, which name both the pool entry a state binds and the symbolic dim it declares, so each state chooses its own.

Parameters:

  • regions (tuple[RecurrentStateRegion, ...])
  • devices (Sequence[DeviceRef])
  • data_parallel_degree (int)
  • page_size (int)
  • kv_connector_config (KVConnectorConfigInterface)
  • enable_prefix_caching (bool)
  • enable_dp_cross_replica_prefix_copy (bool)
  • kv_hash_algo (Literal['ahash64', 'sha256', 'sha256_64'])
  • kv_hash_seed (bytes | None)
  • speculative_method (Literal['eagle', 'mtp', 'dflash', 'dflash2'] | None)
  • num_draft_tokens (int)

allocate_buffers()​

allocate_buffers(total_num_pages, _prefix='')

source

Returns nothing: a state draws from a slab it does not allocate.

Parameters:

  • total_num_pages (int)
  • _prefix (str)

Return type:

list[KVCacheBufferInterface]

build_runtime_inputs()​

build_runtime_inputs(assignments, buffers, _prefix='', *, staging=None)

source

Gathers this forward’s state rows, replica-major.

buffers is unused: a state’s pool is staged in the assignment alongside the rows that address it.

Parameters:

  • assignments (Sequence[KVCacheAssignments])
  • buffers (Sequence[KVCacheBufferInterface])
  • _prefix (str)
  • staging (GraphInputStagingInterface | None)

Return type:

tuple[RecurrentStateInputsPerDevice[Buffer, Buffer], …]

bytes_per_block​

property bytes_per_block: int

source

Returns zero because a state’s page is a per-request cost, not a cost per token.

bytes_per_state​

property bytes_per_state: int

source

Bytes one request’s state occupies on one device, every layer.

data_parallel_degree​

data_parallel_degree: int = 1

source

Degree of data parallelism.

devices​

devices: Sequence[DeviceRef]

source

Devices to use for the cache.

devices_per_replica​

property devices_per_replica: Sequence[Sequence[DeviceRef]]

source

enable_dp_cross_replica_prefix_copy​

enable_dp_cross_replica_prefix_copy: bool = True

source

enable_prefix_caching​

enable_prefix_caching: bool = False

source

get_symbolic_inputs()​

get_symbolic_inputs(namespace='')

source

Returns the symbolic inputs for the state leaves.

namespace is unused: the region ids are already distinct.

Parameters:

namespace (str)

Return type:

tuple[RecurrentStateInputsPerDevice[TensorType, BufferType], …]

kv_connector_config​

kv_connector_config: KVConnectorConfigInterface

source

kv_hash_algo​

kv_hash_algo: Literal['ahash64', 'sha256', 'sha256_64'] = 'ahash64'

source

kv_hash_seed​

kv_hash_seed: bytes | None = None

source

leaf_kind​

leaf_kind: ClassVar[CacheLeafKind] = 'recurrent'

source

leaves()​

leaves(_prefix='')

source

Returns one pool leaf per state leaf, each one state wide.

A region marked scratch lands in the scratch group rather than the recurrent one, so the same tree can declare a published state and the per-request scratch that rides beside it.

_prefix is unused, so the pool key and the name a layer asks for stay the same string.

Parameters:

_prefix (str)

Return type:

Mapping[str, KVLeafRegion]

n_devices​

property n_devices: int

source

Returns the total number of devices.

num_draft_tokens​

num_draft_tokens: int = 0

source

Read only from the leaf a tree takes its pool configuration off.

page_size​

page_size: int = 0

source

Tokens a block covers, set only by a cache with no attention leaf.

Zero otherwise: the attention leaves beside a state declare the pool, and no page covers zero tokens, so nothing real collides with it.

regions​

regions: tuple[RecurrentStateRegion, ...]

source

The state leaves one request occupies, in flatten order.

replicates_kv_across_tp​

property replicates_kv_across_tp: bool

source

a state’s rows hold one device’s shard of the heads.

Type:

False

slab_to_bound_views()​

slab_to_bound_views(slabs)

source

Returns each leaf’s rows, keyed where its layers read them.

Parameters:

slabs (Sequence[Buffer])

Return type:

Mapping[str, list[Buffer]]

slab_to_buffer_views()​

slab_to_buffer_views(buffers, padded_page_bytes=None, _prefix='')

source

Returns each leaf’s pages, one state per page.

Parameters:

Return type:

KVCacheBufferInterface

slab_to_row_views()​

slab_to_row_views(slabs)

source

Converts one replica’s slabs into the rows its kernels index.

Parameters:

slabs (Sequence[Buffer]) – That replica’s [num_huge_blocks, huge_page_bytes] uint8 slab, one per device.

Returns:

One entry per leaf, holding a view per device in the order slabs came in.

Return type:

dict[str, list[Buffer]]

speculative_method​

speculative_method: Literal['eagle', 'mtp', 'dflash', 'dflash2'] | None = None

source

staged_dispatch_metadata()​

staged_dispatch_metadata(_prefix='')

source

Nothing: a state cache resolves no attention dispatch.

Parameters:

_prefix (str)

Return type:

Mapping[str, DispatchMetadataSpec]

tensor_parallel_degree​

property tensor_parallel_degree: int

source

unflatten_kv_inputs()​

unflatten_kv_inputs(it)

source

Unflattens the symbolic inputs for this cache.

Parameters:

it (Iterator[Any])

Return type:

tuple[RecurrentStateInputsPerDevice[TensorValue, BufferValue], …]