(
self,
key_states: torch.Tensor,
value_states: torch.Tensor,
*args,
**kwargs,
)
| 75 | self.is_initialized = True |
| 76 | |
| 77 | def update( |
| 78 | self, |
| 79 | key_states: torch.Tensor, |
| 80 | value_states: torch.Tensor, |
| 81 | *args, |
| 82 | **kwargs, |
| 83 | ) -> tuple[torch.Tensor, torch.Tensor]: |
| 84 | if not self.is_initialized: |
| 85 | self.lazy_initialization(key_states, value_states) |
| 86 | |
| 87 | kv_length = key_states.shape[-2] |
| 88 | |
| 89 | if self._write_position is not None: |
| 90 | cache_position = torch.arange(kv_length, device=self.device) + self._write_position |
| 91 | else: |
| 92 | cache_position = torch.arange(kv_length, device=self.device) |
| 93 | |
| 94 | try: |
| 95 | self.keys.index_copy_(2, cache_position, key_states) |
| 96 | self.values.index_copy_(2, cache_position, value_states) |
| 97 | except NotImplementedError: |
| 98 | self.keys[:, :, cache_position] = key_states |
| 99 | self.values[:, :, cache_position] = value_states |
| 100 | |
| 101 | return self.keys, self.values |
| 102 | |
| 103 | def get_mask_sizes(self, query_length: int) -> tuple[int, int]: |
| 104 | return self.max_cache_len, 0 |
no test coverage detected