(
&mut self,
block_idx: usize,
mut k: Tensor,
mut v: Tensor,
limit: usize,
)
| 182 | } |
| 183 | |
| 184 | fn process_kv_inner( |
| 185 | &mut self, |
| 186 | block_idx: usize, |
| 187 | mut k: Tensor, |
| 188 | mut v: Tensor, |
| 189 | limit: usize, |
| 190 | ) -> Result<(Tensor, Tensor)> { |
| 191 | if self.use_kv_cache { |
| 192 | if let Some((cache_k, cache_v)) = &self.kvs[block_idx] { |
| 193 | // tensor shape is (batch, num_heads, seq_len, head_dim) |
| 194 | // cat() already returns a contiguous tensor — no need for .contiguous() |
| 195 | k = Tensor::cat(&[cache_k, &k], 2)?; |
| 196 | v = Tensor::cat(&[cache_v, &v], 2)?; |
| 197 | |
| 198 | let k_seq_len = k.dims()[2]; |
| 199 | if k_seq_len > limit { |
| 200 | k = k.narrow(2, k_seq_len - limit, limit)?.contiguous()?; |
| 201 | } |
| 202 | let v_seq_len = v.dims()[2]; |
| 203 | if v_seq_len > limit { |
| 204 | v = v.narrow(2, v_seq_len - limit, limit)?.contiguous()?; |
| 205 | } |
| 206 | } |
| 207 | self.kvs[block_idx] = Some((k.clone(), v.clone())) |
| 208 | } |
| 209 | Ok((k, v)) |
| 210 | } |
| 211 | |
| 212 | /// Directly set the KV cache for a layer (used for voice prompt injection). |
| 213 | pub fn set_kv(&mut self, block_idx: usize, k: Tensor, v: Tensor) { |
no test coverage detected