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)
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='')
Returns nothing: a state draws from a slab it does not allocate.
build_runtime_inputs()
build_runtime_inputs(assignments, buffers, _prefix='', *, staging=None)
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.
bytes_per_block
property bytes_per_block: int
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
Bytes one request’s state occupies on one device, every layer.
data_parallel_degree
data_parallel_degree: int = 1
Degree of data parallelism.
devices
Devices to use for the cache.
devices_per_replica
enable_dp_cross_replica_prefix_copy
enable_dp_cross_replica_prefix_copy: bool = True
enable_prefix_caching
enable_prefix_caching: bool = False
get_symbolic_inputs()
get_symbolic_inputs(namespace='')
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
kv_hash_algo
kv_hash_algo: Literal['ahash64', 'sha256', 'sha256_64'] = 'ahash64'
kv_hash_seed
leaf_kind
leaf_kind: ClassVar[CacheLeafKind] = 'recurrent'
leaves()
leaves(_prefix='')
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.
n_devices
property n_devices: int
Returns the total number of devices.
num_draft_tokens
num_draft_tokens: int = 0
Read only from the leaf a tree takes its pool configuration off.
page_size
page_size: int = 0
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, ...]
The state leaves one request occupies, in flatten order.
replicates_kv_across_tp
property replicates_kv_across_tp: bool
a state’s rows hold one device’s shard of the heads.
-
Type:
-
False
slab_to_bound_views()
slab_to_bound_views(slabs)
Returns each leaf’s rows, keyed where its layers read them.
slab_to_buffer_views()
slab_to_buffer_views(buffers, padded_page_bytes=None, _prefix='')
Returns each leaf’s pages, one state per page.
slab_to_row_views()
slab_to_row_views(slabs)
Converts one replica’s slabs into the rows its kernels index.
speculative_method
speculative_method: Literal['eagle', 'mtp', 'dflash', 'dflash2'] | None = None
staged_dispatch_metadata()
staged_dispatch_metadata(_prefix='')
Nothing: a state cache resolves no attention dispatch.
tensor_parallel_degree
property tensor_parallel_degree: int
unflatten_kv_inputs()
unflatten_kv_inputs(it)
Unflattens the symbolic inputs for this cache.
-
Parameters:
-
Return type:
-
tuple[RecurrentStateInputsPerDevice[TensorValue, BufferValue], …]