Central coordinator for anchor-related KV cache interactions.
| 673 | |
| 674 | |
| 675 | class KVCOMMEngine: |
| 676 | """Central coordinator for anchor-related KV cache interactions.""" |
| 677 | |
| 678 | anchors: Dict[str, Any] = {} |
| 679 | anchor_dict: Dict[str, Any] = {} |
| 680 | anchor_len_dict: Dict[str, Any] = {} |
| 681 | anchor_info_dict: Dict[str, Any] = {} |
| 682 | weight_dict: Dict[str, Any] = {} |
| 683 | global_anchor_info_dict: Dict[str, Any] = {} |
| 684 | |
| 685 | _request_lock = threading.Lock() |
| 686 | _request_states: Dict[str, _RequestState] = {} |
| 687 | _active_requests: set[str] = set() |
| 688 | _staged_commits: List[_RequestState] = [] |
| 689 | |
| 690 | def __init__(self, llm: "LLMChat"): |
| 691 | self.llm = llm |
| 692 | self._warning_prefix = "[KVCOMMEngine]" |
| 693 | |
| 694 | def _log_warning(self, message: str) -> None: |
| 695 | logger.opt(colors=True).warning("<yellow>{}</yellow> {}", self._warning_prefix, message) |
| 696 | |
| 697 | @staticmethod |
| 698 | def _stack_cache_tensors(cache: DynamicCache) -> Tuple[torch.Tensor, torch.Tensor]: |
| 699 | return torch.stack(cache.key_cache), torch.stack(cache.value_cache) |
| 700 | |
| 701 | @staticmethod |
| 702 | def _placeholder_length(cache: DynamicCache) -> int: |
| 703 | return cache.key_cache[0].shape[-2] |
| 704 | |
| 705 | def _rotate_segment_caches(self, segment_meta: Dict[str, Any]) -> Tuple[DynamicCache, DynamicCache]: |
| 706 | rotated_placeholder = self.apply_rotary_pos_emb( |
| 707 | segment_meta["ph_cache"], |
| 708 | offset=segment_meta["start"] - segment_meta["drop_num"] + segment_meta["offset_before"], |
| 709 | drop_num=segment_meta["drop_num"], |
| 710 | ) |
| 711 | rotated_prefix = self.apply_rotary_pos_emb( |
| 712 | segment_meta["pf_kv"], |
| 713 | offset=segment_meta["offset_after"], |
| 714 | ) |
| 715 | return rotated_placeholder, rotated_prefix |
| 716 | |
| 717 | @classmethod |
| 718 | def _get_request_state(cls, request_uid: str) -> _RequestState: |
| 719 | """Return or create a request-scoped state container under a lock.""" |
| 720 | if not request_uid: |
| 721 | raise ValueError("request_uid must be provided for scoped anchor updates.") |
| 722 | with cls._request_lock: |
| 723 | state = cls._request_states.get(request_uid) |
| 724 | if state is None: |
| 725 | state = _RequestState( |
| 726 | request_uid, |
| 727 | cls.anchor_dict, |
| 728 | cls.anchor_len_dict, |
| 729 | cls.anchor_info_dict, |
| 730 | cls.weight_dict, |
| 731 | cls.anchors, |
| 732 | cls.global_anchor_info_dict, |