MCPcopy Create free account
hub / github.com/Cornell-RelaxML/qtip / update

Method update

model/cache_utils.py:397–439  ·  view source on GitHub ↗

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,
    )

Source from the content-addressed store, hash-verified

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."""

Callers 1

from_batch_splitsMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected