MSALab/PerceptionDLM
1369
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 