forestcalled/text-generation-webui
0
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 