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
| 13 | |
| 14 | |
| 15 | class 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 | |
| 31 | class BaseKVCache(BaseLayer): |
no outgoing calls