CoolFace
Apppublic

Limour/llama-python-streamingllm

sourceHugging Facegpl-3.0updated 2y agoView on Hugging Face
1likes
btn_com.py138 linesDownload Raw Back to mods
1def init(cfg):2    chat_template = cfg['chat_template']3    model = cfg['model']4    gr = cfg['gr']5    lock = cfg['session_lock']6 7    with gr.Row():8        cfg['btn_vo'] = gr.Button("旁白")9        cfg['btn_rag'] = gr.Button("RAG")10        cfg['btn_retry'] = gr.Button("Retry")11        cfg['btn_stop'] = gr.Button("Stop")12        cfg['btn_reset'] = gr.Button("Reset")13        cfg['btn_debug'] = gr.Button("Debug")14        cfg['btn_submit_vo_suggest'] = gr.Button("Submit&旁白&建议", variant="primary")15        cfg['btn_submit'] = gr.Button("Submit")16        cfg['btn_suggest'] = gr.Button("建议")17        cfg['btn_status_bar'] = gr.Button("状态")18 19    cfg['btn_stop_status'] = True20 21    # ========== 流式输出函数 ==========22    def btn_com(_n_keep, _n_discard,23                _temperature, _repeat_penalty, _frequency_penalty,24                _presence_penalty, _repeat_last_n, _top_k,25                _top_p, _min_p, _typical_p,26                _tfs_z, _mirostat_mode, _mirostat_eta,27                _mirostat_tau, _role, _max_tokens):28        # ========== 初始化输出模版 ==========29        t_bot = chat_template(_role)30        completion_tokens = []  # 有可能多个 tokens 才能构成一个 utf-8 编码的文字31        history = ''32        # ========== 流式输出 ==========33        for token in model.generate_t(34                tokens=t_bot,35                n_keep=_n_keep,36                n_discard=_n_discard,37                im_start=chat_template.im_start_token,38                top_k=_top_k,39                top_p=_top_p,40                min_p=_min_p,41                typical_p=_typical_p,42                temp=_temperature,43                repeat_penalty=_repeat_penalty,44                repeat_last_n=_repeat_last_n,45                frequency_penalty=_frequency_penalty,46                presence_penalty=_presence_penalty,47                tfs_z=_tfs_z,48                mirostat_mode=_mirostat_mode,49                mirostat_tau=_mirostat_tau,50                mirostat_eta=_mirostat_eta,51        ):52            # ========== eos or nlnl 说明eos了 ==========53            if token in chat_template.eos or token == chat_template.nlnl:54                t_bot.extend(completion_tokens)55                print('token in eos', token)56                break57            # ========== 避免不完整的utf-8编码 ==========58            completion_tokens.append(token)59            all_text = model.str_detokenize(completion_tokens)60            if not all_text:61                continue62            t_bot.extend(completion_tokens)63            # ========== 流式输出 ==========64            history += all_text65            yield history66            # ========== \n role \n 结构说明eos了 ==========67            tmp = chat_template.eos_in_role(history, t_bot)68            if tmp:69                tmp -= 1  # 最后一个token并未进入kv_cache70                if tmp:71                    model.venv_pop_token(tmp)72                break73            # ========== \n\n 结构说明eos了 ==========74            tmp = chat_template.eos_in_nlnl(history, t_bot)75            if tmp:76                tmp -= 1  # 最后一个token并未进入kv_cache77                if tmp:78                    model.venv_pop_token(tmp)79                break80            # ========== 过长 or 按下了stop按钮 ==========81            if len(t_bot) > _max_tokens or cfg['btn_stop_status']:82                break83            completion_tokens = []84        # ========== 查看末尾的换行符 ==========85        print('history', repr(history))86        # ========== 给 kv_cache 加上输出结束符 ==========87        model.eval_t(chat_template.im_end_nl, _n_keep, _n_discard, chat_template.im_start_token)88        t_bot.extend(chat_template.im_end_nl)89 90    cfg['btn_com'] = btn_com91 92    btn_start_or_finish_outputs = [cfg['btn_submit'], cfg['btn_vo'], cfg['btn_rag'],93                                   cfg['btn_suggest'], cfg['btn_retry'],94                                   cfg['btn_submit_vo_suggest'],95                                   cfg['btn_status_bar']]96 97    def btn_start_or_finish(finish):98        tmp = gr.update(interactive=finish)99        tmp = (tmp,) * len(btn_start_or_finish_outputs)100 101        def _inner():102            with lock:103                if cfg['session_active'] != finish:104                    raise RuntimeError('任务中断!请稍等或Reset,如已Reset,请忽略。')105                cfg['session_active'] = not cfg['session_active']106                yield tmp107                if finish and cfg['btn_stop_status']:108                    raise RuntimeError('Stop或Reset被按下,任务已中断!如非您所为,可能他人正在使用中!')109                cfg['btn_stop_status'] = finish110 111        return _inner112 113    cfg['btn_concurrency'] = {114        'trigger_mode': 'once',115        'concurrency_id': 'btn_com',116        'concurrency_limit': 1117    }118 119    cfg['btn_start'] = {120        'fn': btn_start_or_finish(False),121        'outputs': btn_start_or_finish_outputs122    }123    cfg['btn_start'].update(cfg['btn_concurrency'])124 125    cfg['btn_finish'] = {126        'fn': btn_start_or_finish(True),127        'outputs': btn_start_or_finish_outputs128    }129    cfg['btn_finish'].update(cfg['btn_concurrency'])130 131    cfg['setting'] = [cfg[x] for x in ('setting_n_keep', 'setting_n_discard',132                                       'setting_temperature', 'setting_repeat_penalty', 'setting_frequency_penalty',133                                       'setting_presence_penalty', 'setting_repeat_last_n', 'setting_top_k',134                                       'setting_top_p', 'setting_min_p', 'setting_typical_p',135                                       'setting_tfs_z', 'setting_mirostat_mode', 'setting_mirostat_eta',136                                       'setting_mirostat_tau', 'role_usr', 'role_char',137                                       'rag', 'setting_max_tokens')]138