(
self,
key_states: torch.Tensor,
value_states: torch.Tensor,
layer_idx: int,
cache_kwargs: Optional[Dict[str, Any]] = None,
)
| 1294 | ) |
| 1295 | |
| 1296 | def update( |
| 1297 | self, |
| 1298 | key_states: torch.Tensor, |
| 1299 | value_states: torch.Tensor, |
| 1300 | layer_idx: int, |
| 1301 | cache_kwargs: Optional[Dict[str, Any]] = None, |
| 1302 | ) -> Tuple[torch.Tensor]: |
| 1303 | cache_position = cache_kwargs.get("cache_position") |
| 1304 | k_out = self.key_cache[layer_idx] |
| 1305 | v_out = self.value_cache[layer_idx] |
| 1306 | |
| 1307 | # assume this only happens in prefill phase when prompt length > sliding_window_size (= max_cache_len) |
| 1308 | if cache_position.shape[0] > self.max_cache_len: |
| 1309 | k_out = key_states[:, :, -self.max_cache_len :, :] |
| 1310 | v_out = value_states[:, :, -self.max_cache_len :, :] |
| 1311 | # Assumption: caches are all zeros at this point, `+=` is equivalent to `=` but compile-friendly |
| 1312 | self.key_cache[layer_idx] += k_out |
| 1313 | self.value_cache[layer_idx] += v_out |
| 1314 | # we should return the whole states instead of k_out, v_out to take the whole prompt |
| 1315 | # into consideration when building kv cache instead of just throwing away tokens outside of the window |
| 1316 | return key_states, value_states |
| 1317 | |
| 1318 | slicing = torch.ones(self.max_cache_len, dtype=torch.long, device=value_states.device).cumsum(0) |
| 1319 | cache_position = cache_position.clamp(0, self.max_cache_len - 1) |
| 1320 | to_shift = cache_position >= self.max_cache_len - 1 |
| 1321 | indices = (slicing + to_shift[-1].int() - 1) % self.max_cache_len |
| 1322 | |
| 1323 | k_out = k_out[:, :, indices] |
| 1324 | v_out = v_out[:, :, indices] |
| 1325 | |
| 1326 | try: |
| 1327 | k_out.index_copy_(2, cache_position.to(k_out.device), key_states.to(k_out.device)) |
| 1328 | v_out.index_copy_(2, cache_position.to(v_out.device), value_states.to(v_out.device)) |
| 1329 | except NotImplementedError: |
| 1330 | # The operator 'aten::index_copy.out' is not currently implemented for the MPS device. |
| 1331 | k_out[:, :, cache_position] = key_states |
| 1332 | v_out[:, :, cache_position] = value_states |
| 1333 | |
| 1334 | # `_.zero()` followed by `+=` is equivalent `=`, but compile-friendly (without graph breaks due to assignment) |
| 1335 | self.key_cache[layer_idx].zero_() |
| 1336 | self.value_cache[layer_idx].zero_() |
| 1337 | |
| 1338 | self.key_cache[layer_idx] += k_out |
| 1339 | self.value_cache[layer_idx] += v_out |
| 1340 | |
| 1341 | return k_out, v_out |
| 1342 | |
| 1343 | def get_max_length(self) -> Optional[int]: |
| 1344 | # in theory there is no limit because the sliding window size is fixed no matter how long the sentence is |
nothing calls this directly
no outgoing calls
no test coverage detected