MCPcopy Create free account
hub / github.com/apple/axlearn / KVState

Class KVState

axlearn/common/kv_cache/base_kv_cache.py:15–28  ·  view source on GitHub ↗

Represents key/value projections. Fields: k_proj: [batch, source_length, num_kv_heads, per_head_dim], Projected key tensor. v_proj: [batch, source_length, num_kv_heads, per_head_dim], Projected value tensor. key_positions: [batch, source_length], Positions of the keys in

Source from the content-addressed store, hash-verified

13
14
15class KVState(NamedTuple):
16 """Represents key/value projections.
17
18 Fields:
19 k_proj: [batch, source_length, num_kv_heads, per_head_dim], Projected key tensor.
20 v_proj: [batch, source_length, num_kv_heads, per_head_dim], Projected value tensor.
21 key_positions: [batch, source_length], Positions of the keys in the batch.
22 page_indices: [batch, max_pages_per_request], optional page indices for batched requests.
23 """
24
25 k_proj: Tensor
26 v_proj: Tensor
27 key_positions: Tensor
28 page_indices: Optional[Tensor] = None
29
30
31class BaseKVCache(BaseLayer):

Calls

no outgoing calls