Updates the cache with the new `key_states` and `value_states` for the layer `layer_idx`. Parameters: key_states (`torch.Tensor`): The new key states to cache. value_states (`torch.Tensor`): The new value states to cache.
(
self,
key_states: torch.Tensor,
value_states: torch.Tensor,
layer_idx: int,
cache_kwargs: Optional[Dict[str, Any]] = None,
)
| 395 | return len(self.key_cache) |
| 396 | |
| 397 | def update( |
| 398 | self, |
| 399 | key_states: torch.Tensor, |
| 400 | value_states: torch.Tensor, |
| 401 | layer_idx: int, |
| 402 | cache_kwargs: Optional[Dict[str, Any]] = None, |
| 403 | ) -> Tuple[torch.Tensor, torch.Tensor]: |
| 404 | """ |
| 405 | Updates the cache with the new `key_states` and `value_states` for the layer `layer_idx`. |
| 406 | |
| 407 | Parameters: |
| 408 | key_states (`torch.Tensor`): |
| 409 | The new key states to cache. |
| 410 | value_states (`torch.Tensor`): |
| 411 | The new value states to cache. |
| 412 | layer_idx (`int`): |
| 413 | The index of the layer to cache the states for. |
| 414 | cache_kwargs (`Dict[str, Any]`, `optional`): |
| 415 | Additional arguments for the cache subclass. No additional arguments are used in `DynamicCache`. |
| 416 | |
| 417 | Return: |
| 418 | A tuple containing the updated key and value states. |
| 419 | """ |
| 420 | # Update the number of seen tokens |
| 421 | if layer_idx == 0: |
| 422 | self._seen_tokens += key_states.shape[-2] |
| 423 | |
| 424 | # Update the cache |
| 425 | if len(self.key_cache) <= layer_idx: |
| 426 | # There may be skipped layers, fill them with empty lists |
| 427 | for _ in range(len(self.key_cache), layer_idx): |
| 428 | self.key_cache.append([]) |
| 429 | self.value_cache.append([]) |
| 430 | self.key_cache.append(key_states) |
| 431 | self.value_cache.append(value_states) |
| 432 | elif len(self.key_cache[layer_idx]) == 0: # fills previously skipped layers; checking for tensor causes errors |
| 433 | self.key_cache[layer_idx] = key_states |
| 434 | self.value_cache[layer_idx] = value_states |
| 435 | else: |
| 436 | self.key_cache[layer_idx] = torch.cat([self.key_cache[layer_idx], key_states], dim=-2) |
| 437 | self.value_cache[layer_idx] = torch.cat([self.value_cache[layer_idx], value_states], dim=-2) |
| 438 | |
| 439 | return self.key_cache[layer_idx], self.value_cache[layer_idx] |
| 440 | |
| 441 | def get_seq_length(self, layer_idx: Optional[int] = 0) -> int: |
| 442 | """Returns the sequence length of the cached states. A layer index can be optionally passed.""" |