CoolFace
Modelpublic

MSALab/PerceptionDLM-Base

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
6likes68downloads
cache.py95 linesDownload Raw Back to root
1from dataclasses import dataclass2 3 4@dataclass5class dLLMCacheConfig:6    prompt_interval_steps: int = 17    gen_interval_steps: int = 18    transfer_ratio: float = 0.09    cfg_interval_steps: int = 110 11 12import torch13from collections import defaultdict14 15 16class Singleton(type):17    _instances = {}18 19    def __call__(cls, *args, **kwargs):20        if cls not in cls._instances:21            cls._instances[cls] = super(Singleton, cls).__call__(*args, **kwargs)22        return cls._instances[cls]23 24 25class dLLMCache(metaclass=Singleton):26    gen_interval_steps: int27    prompt_interval_steps: int28    cfg_interval_steps: int29    prompt_length: int30    transfer_ratio: float31    __cache: defaultdict32    __step_counter: defaultdict33 34    @classmethod35    def new_instance(36        cls,37        prompt_interval_steps: int = 1,38        gen_interval_steps: int = 1,39        cfg_interval_steps: int = 1,40        transfer_ratio: float = 0.0,41    ) -> "dLLMCache":42        ins = cls()43        setattr(ins, "prompt_interval_steps", prompt_interval_steps)44        setattr(ins, "gen_interval_steps", gen_interval_steps)45        setattr(ins, "cfg_interval_steps", cfg_interval_steps)46        setattr(ins, "transfer_ratio", transfer_ratio)47        ins.init()48        return ins49 50    def init(self) -> None:51        self.__cache = defaultdict(52            lambda: defaultdict(lambda: defaultdict(lambda: defaultdict(dict)))53        )54        self.__step_counter = defaultdict(lambda: defaultdict(lambda: 0))55 56    def reset_cache(self, prompt_length: int = 0) -> None:57        self.init()58        torch.cuda.empty_cache()59        self.prompt_length = prompt_length60        self.cache_type = "no_cfg"61 62    def set_cache(63        self, layer_id: int, feature_name: str, features: torch.Tensor, cache_type: str64    ) -> None:65        self.__cache[self.cache_type][cache_type][layer_id][feature_name] = {66            0: features67        }68 69    def get_cache(70        self, layer_id: int, feature_name: str, cache_type: str71    ) -> torch.Tensor:72        output = self.__cache[self.cache_type][cache_type][layer_id][feature_name][0]73        return output74 75    def update_step(self, layer_id: int) -> None:76        self.__step_counter[self.cache_type][layer_id] += 177 78    def refresh_gen(self, layer_id: int = 0) -> bool:79        return (self.current_step - 1) % self.gen_interval_steps == 080 81    def refresh_prompt(self, layer_id: int = 0) -> bool:82        return (self.current_step - 1) % self.prompt_interval_steps == 083 84    def refresh_cfg(self, layer_id: int = 0) -> bool:85        return (86            self.current_step - 187        ) % self.cfg_interval_steps == 0 or self.current_step <= 588 89    @property90    def current_step(self) -> int:91        return max(list(self.__step_counter[self.cache_type].values()), default=1)92 93    def __repr__(self):94        return f"USE dLLMCache"95