hugging-apps/echo-memory
0
1import torch2 3 4 5class GeneralLoRALoader:6 def __init__(self, device="cpu", torch_dtype=torch.float32):7 self.device = device8 self.torch_dtype = torch_dtype9 10 11 def get_name_dict(self, lora_state_dict):12 lora_name_dict = {}13 for key in lora_state_dict:14 if ".lora_B." not in key:15 continue16 keys = key.split(".")17 if len(keys) > keys.index("lora_B") + 2:18 keys.pop(keys.index("lora_B") + 1)19 keys.pop(keys.index("lora_B"))20 if keys[0] == "diffusion_model":21 keys.pop(0)22 keys.pop(-1)23 target_name = ".".join(keys)24 lora_name_dict[target_name] = (key, key.replace(".lora_B.", ".lora_A."))25 return lora_name_dict26 27 28 def load(self, model: torch.nn.Module, state_dict_lora, alpha=1.0):29 updated_num = 030 lora_name_dict = self.get_name_dict(state_dict_lora)31 for name, module in model.named_modules():32 if name in lora_name_dict:33 weight_up = state_dict_lora[lora_name_dict[name][0]].to(device=self.device, dtype=self.torch_dtype)34 weight_down = state_dict_lora[lora_name_dict[name][1]].to(device=self.device, dtype=self.torch_dtype)35 if len(weight_up.shape) == 4:36 weight_up = weight_up.squeeze(3).squeeze(2)37 weight_down = weight_down.squeeze(3).squeeze(2)38 weight_lora = alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3)39 else:40 weight_lora = alpha * torch.mm(weight_up, weight_down)41 state_dict = module.state_dict()42 state_dict["weight"] = state_dict["weight"].to(device=self.device, dtype=self.torch_dtype) + weight_lora43 module.load_state_dict(state_dict)44 updated_num += 145 print(f"{updated_num} tensors are updated by LoRA.")46 