MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / update

Method update

deepspeed/utils/static_cache.py:77–101  ·  view source on GitHub ↗
(
        self,
        key_states: torch.Tensor,
        value_states: torch.Tensor,
        *args,
        **kwargs,
    )

Source from the content-addressed store, hash-verified

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

Calls 1

lazy_initializationMethod · 0.95

Tested by

no test coverage detected