forestcalled/text-generation-webui
0
1import json2import re3from pathlib import Path4 5import yaml6 7from modules import chat, loaders, metadata_gguf, shared, ui8 9 10def get_fallback_settings():11 return {12 'wbits': 'None',13 'groupsize': 'None',14 'desc_act': False,15 'model_type': 'None',16 'max_seq_len': 2048,17 'n_ctx': 2048,18 'rope_freq_base': 0,19 'compress_pos_emb': 1,20 'truncation_length': shared.settings['truncation_length'],21 'skip_special_tokens': shared.settings['skip_special_tokens'],22 'custom_stopping_strings': shared.settings['custom_stopping_strings'],23 }24 25 26def get_model_metadata(model):27 model_settings = {}28 29 # Get settings from models/config.yaml and models/config-user.yaml30 settings = shared.model_config31 for pat in settings:32 if re.match(pat.lower(), model.lower()):33 for k in settings[pat]:34 model_settings[k] = settings[pat][k]35 36 path = Path(f'{shared.args.model_dir}/{model}/config.json')37 if path.exists():38 hf_metadata = json.loads(open(path, 'r').read())39 else:40 hf_metadata = None41 42 if 'loader' not in model_settings:43 if hf_metadata is not None and 'quip_params' in hf_metadata:44 model_settings['loader'] = 'QuIP#'45 else:46 loader = infer_loader(model, model_settings)47 if 'wbits' in model_settings and type(model_settings['wbits']) is int and model_settings['wbits'] > 0:48 loader = 'AutoGPTQ'49 50 model_settings['loader'] = loader51 52 # GGUF metadata53 if model_settings['loader'] in ['llama.cpp', 'llamacpp_HF', 'ctransformers']:54 path = Path(f'{shared.args.model_dir}/{model}')55 if path.is_file():56 model_file = path57 else:58 model_file = list(path.glob('*.gguf'))[0]59 60 metadata = metadata_gguf.load_metadata(model_file)61 if 'llama.context_length' in metadata:62 model_settings['n_ctx'] = metadata['llama.context_length']63 if 'llama.rope.scale_linear' in metadata:64 model_settings['compress_pos_emb'] = metadata['llama.rope.scale_linear']65 if 'llama.rope.freq_base' in metadata:66 model_settings['rope_freq_base'] = metadata['llama.rope.freq_base']67 if 'tokenizer.chat_template' in metadata:68 template = metadata['tokenizer.chat_template']69 eos_token = metadata['tokenizer.ggml.tokens'][metadata['tokenizer.ggml.eos_token_id']]70 bos_token = metadata['tokenizer.ggml.tokens'][metadata['tokenizer.ggml.bos_token_id']]71 template = template.replace('eos_token', "'{}'".format(eos_token))72 template = template.replace('bos_token', "'{}'".format(bos_token))73 74 template = re.sub(r'raise_exception\([^)]*\)', "''", template)75 model_settings['instruction_template'] = 'Custom (obtained from model metadata)'76 model_settings['instruction_template_str'] = template77 78 else:79 # Transformers metadata80 if hf_metadata is not None:81 metadata = json.loads(open(path, 'r').read())82 if 'max_position_embeddings' in metadata:83 model_settings['truncation_length'] = metadata['max_position_embeddings']84 model_settings['max_seq_len'] = metadata['max_position_embeddings']85 86 if 'rope_theta' in metadata:87 model_settings['rope_freq_base'] = metadata['rope_theta']88 89 if 'rope_scaling' in metadata and type(metadata['rope_scaling']) is dict and all(key in metadata['rope_scaling'] for key in ('type', 'factor')):90 if metadata['rope_scaling']['type'] == 'linear':91 model_settings['compress_pos_emb'] = metadata['rope_scaling']['factor']92 93 if 'quantization_config' in metadata:94 if 'bits' in metadata['quantization_config']:95 model_settings['wbits'] = metadata['quantization_config']['bits']96 if 'group_size' in metadata['quantization_config']:97 model_settings['groupsize'] = metadata['quantization_config']['group_size']98 if 'desc_act' in metadata['quantization_config']:99 model_settings['desc_act'] = metadata['quantization_config']['desc_act']100 101 # Read AutoGPTQ metadata102 path = Path(f'{shared.args.model_dir}/{model}/quantize_config.json')103 if path.exists():104 metadata = json.loads(open(path, 'r').read())105 if 'bits' in metadata:106 model_settings['wbits'] = metadata['bits']107 if 'group_size' in metadata:108 model_settings['groupsize'] = metadata['group_size']109 if 'desc_act' in metadata:110 model_settings['desc_act'] = metadata['desc_act']111 112 # Try to find the Jinja instruct template113 path = Path(f'{shared.args.model_dir}/{model}') / 'tokenizer_config.json'114 if path.exists():115 metadata = json.loads(open(path, 'r').read())116 if 'chat_template' in metadata:117 template = metadata['chat_template']118 for k in ['eos_token', 'bos_token']:119 if k in metadata:120 value = metadata[k]121 if type(value) is dict:122 value = value['content']123 124 template = template.replace(k, "'{}'".format(value))125 126 template = re.sub(r'raise_exception\([^)]*\)', "''", template)127 model_settings['instruction_template'] = 'Custom (obtained from model metadata)'128 model_settings['instruction_template_str'] = template129 130 if 'instruction_template' not in model_settings:131 model_settings['instruction_template'] = 'Alpaca'132 133 if model_settings['instruction_template'] != 'Custom (obtained from model metadata)':134 model_settings['instruction_template_str'] = chat.load_instruction_template(model_settings['instruction_template'])135 136 # Ignore rope_freq_base if set to the default value137 if 'rope_freq_base' in model_settings and model_settings['rope_freq_base'] == 10000:138 model_settings.pop('rope_freq_base')139 140 # Apply user settings from models/config-user.yaml141 settings = shared.user_config142 for pat in settings:143 if re.match(pat.lower(), model.lower()):144 for k in settings[pat]:145 model_settings[k] = settings[pat][k]146 147 return model_settings148 149 150def infer_loader(model_name, model_settings):151 path_to_model = Path(f'{shared.args.model_dir}/{model_name}')152 if not path_to_model.exists():153 loader = None154 elif (path_to_model / 'quantize_config.json').exists() or ('wbits' in model_settings and type(model_settings['wbits']) is int and model_settings['wbits'] > 0):155 loader = 'ExLlama_HF'156 elif (path_to_model / 'quant_config.json').exists() or re.match(r'.*-awq', model_name.lower()):157 loader = 'AutoAWQ'158 elif len(list(path_to_model.glob('*.gguf'))) > 0:159 loader = 'llama.cpp'160 elif re.match(r'.*\.gguf', model_name.lower()):161 loader = 'llama.cpp'162 elif re.match(r'.*rwkv.*\.pth', model_name.lower()):163 loader = 'RWKV'164 elif re.match(r'.*exl2', model_name.lower()):165 loader = 'ExLlamav2_HF'166 elif re.match(r'.*-hqq', model_name.lower()):167 return 'HQQ'168 else:169 loader = 'Transformers'170 171 return loader172 173 174def update_model_parameters(state, initial=False):175 '''176 UI: update the command-line arguments based on the interface values177 '''178 elements = ui.list_model_elements() # the names of the parameters179 gpu_memories = []180 181 for i, element in enumerate(elements):182 if element not in state:183 continue184 185 value = state[element]186 if element.startswith('gpu_memory'):187 gpu_memories.append(value)188 continue189 190 if initial and element in shared.provided_arguments:191 continue192 193 # Setting null defaults194 if element in ['wbits', 'groupsize', 'model_type'] and value == 'None':195 value = vars(shared.args_defaults)[element]196 elif element in ['cpu_memory'] and value == 0:197 value = vars(shared.args_defaults)[element]198 199 # Making some simple conversions200 if element in ['wbits', 'groupsize', 'pre_layer']:201 value = int(value)202 elif element == 'cpu_memory' and value is not None:203 value = f"{value}MiB"204 205 if element in ['pre_layer']:206 value = [value] if value > 0 else None207 208 setattr(shared.args, element, value)209 210 found_positive = False211 for i in gpu_memories:212 if i > 0:213 found_positive = True214 break215 216 if not (initial and vars(shared.args)['gpu_memory'] != vars(shared.args_defaults)['gpu_memory']):217 if found_positive:218 shared.args.gpu_memory = [f"{i}MiB" for i in gpu_memories]219 else:220 shared.args.gpu_memory = None221 222 223def apply_model_settings_to_state(model, state):224 '''225 UI: update the state variable with the model settings226 '''227 model_settings = get_model_metadata(model)228 if 'loader' in model_settings:229 loader = model_settings.pop('loader')230 231 # If the user is using an alternative loader for the same model type, let them keep using it232 if not (loader == 'AutoGPTQ' and state['loader'] in ['GPTQ-for-LLaMa', 'ExLlama', 'ExLlama_HF', 'ExLlamav2', 'ExLlamav2_HF']) and not (loader == 'llama.cpp' and state['loader'] in ['llamacpp_HF', 'ctransformers']):233 state['loader'] = loader234 235 for k in model_settings:236 if k in state:237 if k in ['wbits', 'groupsize']:238 state[k] = str(model_settings[k])239 else:240 state[k] = model_settings[k]241 242 return state243 244 245def save_model_settings(model, state):246 '''247 Save the settings for this model to models/config-user.yaml248 '''249 if model == 'None':250 yield ("Not saving the settings because no model is loaded.")251 return252 253 with Path(f'{shared.args.model_dir}/config-user.yaml') as p:254 if p.exists():255 user_config = yaml.safe_load(open(p, 'r').read())256 else:257 user_config = {}258 259 model_regex = model + '$' # For exact matches260 if model_regex not in user_config:261 user_config[model_regex] = {}262 263 for k in ui.list_model_elements():264 if k == 'loader' or k in loaders.loaders_and_params[state['loader']]:265 user_config[model_regex][k] = state[k]266 267 shared.user_config = user_config268 269 output = yaml.dump(user_config, sort_keys=False)270 with open(p, 'w') as f:271 f.write(output)272 273 yield (f"Settings for `{model}` saved to `{p}`.")274 