CoolFace
Apppublic

forestcalled/text-generation-webui

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
models_settings.py274 linesDownload Raw Back to modules
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