CoolFace
Apppublic

forestcalled/text-generation-webui

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
LoRA.py198 linesDownload Raw Back to modules
1from pathlib import Path2 3import torch4from peft import PeftModel5from transformers import is_torch_xpu_available6 7import modules.shared as shared8from modules.logging_colors import logger9from modules.models import reload_model10 11 12def add_lora_to_model(lora_names):13    if 'GPTQForCausalLM' in shared.model.__class__.__name__ or shared.args.loader == 'AutoGPTQ':14        add_lora_autogptq(lora_names)15    elif shared.model.__class__.__name__ in ['ExllamaModel', 'ExllamaHF'] or shared.args.loader == 'ExLlama':16        add_lora_exllama(lora_names)17    elif shared.model.__class__.__name__ in ['Exllamav2Model', 'Exllamav2HF'] or shared.args.loader == ['ExLlamav2', 'ExLlamav2_HF']:18        add_lora_exllamav2(lora_names)19    else:20        add_lora_transformers(lora_names)21 22 23def get_lora_path(lora_name):24    p = Path(lora_name)25    if p.exists():26        lora_name = p.parts[-1]27 28    return Path(f"{shared.args.lora_dir}/{lora_name}")29 30 31def add_lora_exllama(lora_names):32 33    try:34        from exllama.lora import ExLlamaLora35    except:36        try:37            from repositories.exllama.lora import ExLlamaLora38        except:39            logger.error("Could not find the file repositories/exllama/lora.py. Make sure that exllama is cloned inside repositories/ and is up to date.")40            return41 42    if len(lora_names) == 0:43        if shared.model.__class__.__name__ == 'ExllamaModel':44            shared.model.generator.lora = None45        else:46            shared.model.lora = None47 48        shared.lora_names = []49        return50    else:51        if len(lora_names) > 1:52            logger.warning('ExLlama can only work with 1 LoRA at the moment. Only the first one in the list will be loaded.')53 54        lora_path = get_lora_path(lora_names[0])55        lora_config_path = lora_path / "adapter_config.json"56        for file_name in ["adapter_model.safetensors", "adapter_model.bin"]:57            file_path = lora_path / file_name58            if file_path.is_file():59                lora_adapter_path = file_path60 61        logger.info("Applying the following LoRAs to {}: {}".format(shared.model_name, ', '.join([lora_names[0]])))62        if shared.model.__class__.__name__ == 'ExllamaModel':63            lora = ExLlamaLora(shared.model.model, str(lora_config_path), str(lora_adapter_path))64            shared.model.generator.lora = lora65        else:66            lora = ExLlamaLora(shared.model.ex_model, str(lora_config_path), str(lora_adapter_path))67            shared.model.lora = lora68 69        shared.lora_names = [lora_names[0]]70        return71 72 73def add_lora_exllamav2(lora_names):74 75    from exllamav2 import ExLlamaV2Lora76 77    if isinstance(shared.model.loras, list):78        for lora in shared.model.loras:79            lora.unload()80 81    if len(lora_names) > 0:82        logger.info("Applying the following LoRAs to {}: {}".format(shared.model_name, ', '.join(lora_names)))83        shared.model.loras = []84        for lora_name in lora_names:85            lora_path = get_lora_path(lora_name)86            if shared.model.__class__.__name__ == 'Exllamav2Model':87                lora = ExLlamaV2Lora.from_directory(shared.model.model, str(lora_path))88            else:89                lora = ExLlamaV2Lora.from_directory(shared.model.ex_model, str(lora_path))90 91            shared.model.loras.append(lora)92 93        shared.lora_names = lora_names94    else:95        shared.lora_names = []96        shared.model.loras = None97 98 99def add_lora_autogptq(lora_names):100    '''101    Adapted from https://github.com/Ph0rk0z/text-generation-webui-testing102    '''103 104    try:105        from auto_gptq import get_gptq_peft_model106        from auto_gptq.utils.peft_utils import GPTQLoraConfig107    except:108        logger.error("This version of AutoGPTQ does not support LoRA. You need to install from source or wait for a new release.")109        return110 111    if len(lora_names) == 0:112        reload_model()113 114        shared.lora_names = []115        return116    else:117        if len(lora_names) > 1:118            logger.warning('AutoGPTQ can only work with 1 LoRA at the moment. Only the first one in the list will be loaded.')119        if not shared.args.no_inject_fused_attention:120            logger.warning('Fused Atttention + AutoGPTQ may break Lora loading. Disable it.')121 122        peft_config = GPTQLoraConfig(123            inference_mode=True,124        )125 126        lora_path = get_lora_path(lora_names[0])127        logger.info("Applying the following LoRAs to {}: {}".format(shared.model_name, ', '.join([lora_names[0]])))128        shared.model = get_gptq_peft_model(shared.model, peft_config, lora_path)129        shared.lora_names = [lora_names[0]]130        return131 132 133def add_lora_transformers(lora_names):134    prior_set = set(shared.lora_names)135    added_set = set(lora_names) - prior_set136    removed_set = prior_set - set(lora_names)137 138    # If no LoRA needs to be added or removed, exit139    if len(added_set) == 0 and len(removed_set) == 0:140        return141 142    # Add a LoRA when another LoRA is already present143    if len(removed_set) == 0 and len(prior_set) > 0 and "__merged" not in shared.model.peft_config.keys():144        logger.info(f"Adding the LoRA(s) named {added_set} to the model")145        for lora in added_set:146            shared.model.load_adapter(get_lora_path(lora), lora)147 148        if len(lora_names) > 1:149            merge_loras()150 151        shared.lora_names = lora_names152        return153 154    # If any LoRA needs to be removed, start over155    if len(removed_set) > 0:156        shared.model = shared.model.unload()157 158    if len(lora_names) > 0:159        params = {}160        if not shared.args.cpu:161            if shared.args.load_in_4bit or shared.args.load_in_8bit:162                params['peft_type'] = shared.model.dtype163            else:164                params['dtype'] = shared.model.dtype165                if hasattr(shared.model, "hf_device_map"):166                    params['device_map'] = {"base_model.model." + k: v for k, v in shared.model.hf_device_map.items()}167 168        logger.info("Applying the following LoRAs to {}: {}".format(shared.model_name, ', '.join(lora_names)))169        shared.model = PeftModel.from_pretrained(shared.model, get_lora_path(lora_names[0]), adapter_name=lora_names[0], **params)170        for lora in lora_names[1:]:171            shared.model.load_adapter(get_lora_path(lora), lora)172 173        if len(lora_names) > 1:174            merge_loras()175 176        if not shared.args.load_in_8bit and not shared.args.cpu:177            shared.model.half()178            if not hasattr(shared.model, "hf_device_map"):179                if torch.backends.mps.is_available():180                    device = torch.device('mps')181                    shared.model = shared.model.to(device)182                elif is_torch_xpu_available():183                    device = torch.device("xpu:0")184                    shared.model = shared.model.to(device)185                else:186                    shared.model = shared.model.cuda()187 188    shared.lora_names = lora_names189 190 191def merge_loras():192    if len(list({shared.model.peft_config[adapter].r for adapter in shared.model.peft_config.keys()})) > 1:193        logger.warning("The loaded LoRAs cannot be merged, as they have dissimilar ranks. Only the first one will be active.")194        return195 196    shared.model.add_weighted_adapter(shared.lora_names, [1] * len(shared.lora_names), "__merged")197    shared.model.set_adapter("__merged")198