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

AttnKeyInterface

AttnKeyInterface​

class max.nn.kv_cache.AttnKeyInterface

source

Bases: object

Common base for resolved attention keys.

dispatch_metadata_buffers()​

dispatch_metadata_buffers(staging, *, name, devices, max_cache_valid_length)

source

One buffer per tensor-parallel shard, staged when it can be.

A device-resident key with somewhere to stage is one host write fanned out to every shard, so it travels in the forward’s own transfer. Otherwise each shard gets a buffer of its own – for a host-resident key there is no transfer to fold in anyway.

Parameters:

  • staging (GraphInputStagingInterface | None)
  • name (str)
  • devices (Sequence[Device])
  • max_cache_valid_length (int)

Return type:

tuple[Buffer, …]

dispatch_metadata_spec()​

classmethod dispatch_metadata_spec()

source

Returns the shape, dtype and residence of this kernel’s metadata.

Return type:

DispatchMetadataSpec

pack_into()​

pack_into(into, max_cache_valid_length)

source

Writes this dispatch shape’s values into into.

into is shaped by dispatch_metadata_spec(). Writing rather than allocating is what lets a caller hand over staging it already owns – see GraphInputStager.

max_cache_valid_length is the runtime cache length; it is supplied here rather than stored so the identity is independent of it.

Parameters:

Return type:

None

pack_into_buffer()​

pack_into_buffer(device, max_cache_valid_length)

source

Packs this into a freshly allocated dispatch-metadata buffer.

For callers with nowhere to stage it. A host-resident spec stays on the host; a device-resident one is copied over, which is the copy staging exists to fold into the rest of a forward’s.

Parameters:

  • device (Device)
  • max_cache_valid_length (int)

Return type:

Buffer