| 1357 | return (self.self_attention_cache.key_cache[layer_idx][0, 0].any(dim=-1)).sum() |
| 1358 | |
| 1359 | def reset(self): |
| 1360 | if hasattr(self.self_attention_cache, "reset"): |
| 1361 | self.self_attention_cache.reset() |
| 1362 | if hasattr(self.cross_attention_cache, "reset"): |
| 1363 | self.cross_attention_cache.reset() |
| 1364 | elif not hasattr(self.self_attention_cache, "reset") and not hasattr(self.cross_attention_cache, "reset"): |
| 1365 | raise ValueError( |
| 1366 | "Neither self nor cross-attention cache have valid `.reset()` methods. `.reset()` should " |
| 1367 | "only be called on compatible cache classes, such as `StaticCache` or `SlidingWindowCache`. " |
| 1368 | f"Got {self.self_attention_cache.__str__()} for the self attention cache and " |
| 1369 | f"{self.cross_attention_cache.__str__()} for the cross attention cache." |
| 1370 | ) |
| 1371 | for layer_idx in self.is_updated: |
| 1372 | self.is_updated[layer_idx] = False |
| 1373 | |
| 1374 | def reorder_cache(self, beam_idx: torch.LongTensor): |
| 1375 | """Reorders the cache for beam search, given the selected beam indices.""" |