modelscope/DiffSynth-Painter
14
1import torch2from .sd_unet import SDUNet3from .sdxl_unet import SDXLUNet4from .sd_text_encoder import SDTextEncoder5from .sdxl_text_encoder import SDXLTextEncoder, SDXLTextEncoder26from .sd3_dit import SD3DiT7from .hunyuan_dit import HunyuanDiT8 9 10 11class LoRAFromCivitai:12 def __init__(self):13 self.supported_model_classes = []14 self.lora_prefix = []15 self.renamed_lora_prefix = {}16 self.special_keys = {}17 18 19 def convert_state_dict(self, state_dict, lora_prefix="lora_unet_", alpha=1.0):20 renamed_lora_prefix = self.renamed_lora_prefix.get(lora_prefix, "")21 state_dict_ = {}22 for key in state_dict:23 if ".lora_up" not in key:24 continue25 if not key.startswith(lora_prefix):26 continue27 weight_up = state_dict[key].to(device="cuda", dtype=torch.float16)28 weight_down = state_dict[key.replace(".lora_up", ".lora_down")].to(device="cuda", dtype=torch.float16)29 if len(weight_up.shape) == 4:30 weight_up = weight_up.squeeze(3).squeeze(2).to(torch.float32)31 weight_down = weight_down.squeeze(3).squeeze(2).to(torch.float32)32 lora_weight = alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3)33 else:34 lora_weight = alpha * torch.mm(weight_up, weight_down)35 target_name = key.split(".")[0].replace(lora_prefix, renamed_lora_prefix).replace("_", ".") + ".weight"36 for special_key in self.special_keys:37 target_name = target_name.replace(special_key, self.special_keys[special_key])38 state_dict_[target_name] = lora_weight.cpu()39 return state_dict_40 41 42 def load(self, model, state_dict_lora, lora_prefix, alpha=1.0, model_resource=None):43 state_dict_model = model.state_dict()44 state_dict_lora = self.convert_state_dict(state_dict_lora, lora_prefix=lora_prefix, alpha=alpha)45 if model_resource == "diffusers":46 state_dict_lora = model.__class__.state_dict_converter().from_diffusers(state_dict_lora)47 elif model_resource == "civitai":48 state_dict_lora = model.__class__.state_dict_converter().from_civitai(state_dict_lora)49 if len(state_dict_lora) > 0:50 print(f" {len(state_dict_lora)} tensors are updated.")51 for name in state_dict_lora:52 state_dict_model[name] += state_dict_lora[name].to(53 dtype=state_dict_model[name].dtype, device=state_dict_model[name].device)54 model.load_state_dict(state_dict_model)55 56 57 def match(self, model, state_dict_lora):58 for lora_prefix, model_class in zip(self.lora_prefix, self.supported_model_classes):59 if not isinstance(model, model_class):60 continue61 state_dict_model = model.state_dict()62 for model_resource in ["diffusers", "civitai"]:63 try:64 state_dict_lora_ = self.convert_state_dict(state_dict_lora, lora_prefix=lora_prefix, alpha=1.0)65 converter_fn = model.__class__.state_dict_converter().from_diffusers if model_resource == "diffusers" \66 else model.__class__.state_dict_converter().from_civitai67 state_dict_lora_ = converter_fn(state_dict_lora_)68 if len(state_dict_lora_) == 0:69 continue70 for name in state_dict_lora_:71 if name not in state_dict_model:72 break73 else:74 return lora_prefix, model_resource75 except:76 pass77 return None78 79 80 81class SDLoRAFromCivitai(LoRAFromCivitai):82 def __init__(self):83 super().__init__()84 self.supported_model_classes = [SDUNet, SDTextEncoder]85 self.lora_prefix = ["lora_unet_", "lora_te_"]86 self.special_keys = {87 "down.blocks": "down_blocks",88 "up.blocks": "up_blocks",89 "mid.block": "mid_block",90 "proj.in": "proj_in",91 "proj.out": "proj_out",92 "transformer.blocks": "transformer_blocks",93 "to.q": "to_q",94 "to.k": "to_k",95 "to.v": "to_v",96 "to.out": "to_out",97 "text.model": "text_model",98 "self.attn.q.proj": "self_attn.q_proj",99 "self.attn.k.proj": "self_attn.k_proj",100 "self.attn.v.proj": "self_attn.v_proj",101 "self.attn.out.proj": "self_attn.out_proj",102 "input.blocks": "model.diffusion_model.input_blocks",103 "middle.block": "model.diffusion_model.middle_block",104 "output.blocks": "model.diffusion_model.output_blocks",105 }106 107 108class SDXLLoRAFromCivitai(LoRAFromCivitai):109 def __init__(self):110 super().__init__()111 self.supported_model_classes = [SDXLUNet, SDXLTextEncoder, SDXLTextEncoder2]112 self.lora_prefix = ["lora_unet_", "lora_te1_", "lora_te2_"]113 self.renamed_lora_prefix = {"lora_te2_": "2"}114 self.special_keys = {115 "down.blocks": "down_blocks",116 "up.blocks": "up_blocks",117 "mid.block": "mid_block",118 "proj.in": "proj_in",119 "proj.out": "proj_out",120 "transformer.blocks": "transformer_blocks",121 "to.q": "to_q",122 "to.k": "to_k",123 "to.v": "to_v",124 "to.out": "to_out",125 "text.model": "conditioner.embedders.0.transformer.text_model",126 "self.attn.q.proj": "self_attn.q_proj",127 "self.attn.k.proj": "self_attn.k_proj",128 "self.attn.v.proj": "self_attn.v_proj",129 "self.attn.out.proj": "self_attn.out_proj",130 "input.blocks": "model.diffusion_model.input_blocks",131 "middle.block": "model.diffusion_model.middle_block",132 "output.blocks": "model.diffusion_model.output_blocks",133 "2conditioner.embedders.0.transformer.text_model.encoder.layers": "text_model.encoder.layers"134 }135 136 137 138class GeneralLoRAFromPeft:139 def __init__(self):140 self.supported_model_classes = [SDUNet, SDXLUNet, SD3DiT, HunyuanDiT]141 142 143 def convert_state_dict(self, state_dict, alpha=1.0, device="cuda", torch_dtype=torch.float16):144 state_dict_ = {}145 for key in state_dict:146 if ".lora_B." not in key:147 continue148 weight_up = state_dict[key].to(device=device, dtype=torch_dtype)149 weight_down = state_dict[key.replace(".lora_B.", ".lora_A.")].to(device=device, dtype=torch_dtype)150 if len(weight_up.shape) == 4:151 weight_up = weight_up.squeeze(3).squeeze(2)152 weight_down = weight_down.squeeze(3).squeeze(2)153 lora_weight = alpha * torch.mm(weight_up, weight_down).unsqueeze(2).unsqueeze(3)154 else:155 lora_weight = alpha * torch.mm(weight_up, weight_down)156 keys = key.split(".")157 keys.pop(keys.index("lora_B") + 1)158 keys.pop(keys.index("lora_B"))159 target_name = ".".join(keys)160 state_dict_[target_name] = lora_weight.cpu()161 return state_dict_162 163 164 def load(self, model, state_dict_lora, lora_prefix="", alpha=1.0, model_resource=""):165 state_dict_model = model.state_dict()166 for name, param in state_dict_model.items():167 torch_dtype = param.dtype168 device = param.device169 break170 state_dict_lora = self.convert_state_dict(state_dict_lora, alpha=alpha, device=device, torch_dtype=torch_dtype)171 if len(state_dict_lora) > 0:172 print(f" {len(state_dict_lora)} tensors are updated.")173 for name in state_dict_lora:174 state_dict_model[name] += state_dict_lora[name].to(175 dtype=state_dict_model[name].dtype, device=state_dict_model[name].device)176 model.load_state_dict(state_dict_model)177 178 179 def match(self, model, state_dict_lora):180 for model_class in self.supported_model_classes:181 if not isinstance(model, model_class):182 continue183 state_dict_model = model.state_dict()184 try:185 state_dict_lora_ = self.convert_state_dict(state_dict_lora, alpha=1.0)186 if len(state_dict_lora_) == 0:187 continue188 for name in state_dict_lora_:189 if name not in state_dict_model:190 break191 else:192 return "", ""193 except:194 pass195 return None