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

Method update

model/cache_utils.py:1296–1341  ·  view source on GitHub ↗
(
        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

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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected