MCPcopy Create free account
hub / github.com/FastMAS/KVCOMM / KVCOMMEngine

Class KVCOMMEngine

KVCOMM/llm/kvcomm_engine.py:675–1296  ·  view source on GitHub ↗

Central coordinator for anchor-related KV cache interactions.

Source from the content-addressed store, hash-verified

673
674
675class 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,

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected