Reset all per-block state for a new sequence.
(self)
| 421 | ) # → [q_len, H, D] |
| 422 | |
| 423 | def reset(self) -> None: |
| 424 | """Reset all per-block state for a new sequence.""" |
| 425 | for i in range(self.num_blocks): |
| 426 | self.scale_patch_pages[i].clear() |
| 427 | self.live_window_patch_pages[i].clear() |
| 428 | self.all_special_pages[i].clear() |
| 429 | self.free_patch_pages[i] = list(range(self.max_patch_pages)) |
| 430 | self.free_special_pages[i] = list(range(self.max_patch_pages, self.max_num_pages)) |
| 431 | self.special_token_count[i] = 0 |
| 432 | self.frame_count[i] = 0 |
| 433 | |
| 434 | # ========================================================================= |
| 435 | # Helper methods |
no test coverage detected