Team Ai
Apppublic

modelscope/DiffSynth-Painter

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
14likes
lora.py195 linesDownload Raw Back to models
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