CoolFace
Apppublic

stinkypudding/code_assistant

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py293 linesDownload Raw Back to root
1"""This script refers to the dialogue example of streamlit, the interactive2generation code of chatglm2 and transformers.3 4We mainly modified part of the code logic to adapt to the5generation of our model.6Please refer to these links below for more information:7    1. streamlit chat example:8        https://docs.streamlit.io/knowledge-base/tutorials/build-conversational-apps9    2. chatglm2:10        https://github.com/THUDM/ChatGLM2-6B11    3. transformers:12        https://github.com/huggingface/transformers13Please run with the command `streamlit run path/to/web_demo.py14    --server.address=0.0.0.0 --server.port 7860`.15Using `python path/to/web_demo.py` may cause unknown problems.16"""17# isort: skip_file18import copy19import warnings20from dataclasses import asdict, dataclass21from typing import Callable, List, Optional22 23import streamlit as st24import torch25from torch import nn26from transformers.generation.utils import (LogitsProcessorList,27                                           StoppingCriteriaList)28from transformers.utils import logging29 30from transformers import AutoTokenizer, AutoModelForCausalLM  # isort: skip31 32logger = logging.get_logger(__name__)33model_name_or_path="https://huggingface.co/stinkypudding/code_assistant"34 35@dataclass36class GenerationConfig:37    # this config is used for chat to provide more diversity38    max_length: int = 3276839    top_p: float = 0.840    temperature: float = 0.841    do_sample: bool = True42    repetition_penalty: float = 1.00543 44 45@torch.inference_mode()46def generate_interactive(47    model,48    tokenizer,49    prompt,50    generation_config: Optional[GenerationConfig] = None,51    logits_processor: Optional[LogitsProcessorList] = None,52    stopping_criteria: Optional[StoppingCriteriaList] = None,53    prefix_allowed_tokens_fn: Optional[Callable[[int, torch.Tensor],54                                                List[int]]] = None,55    additional_eos_token_id: Optional[int] = None,56    **kwargs,57):58    inputs = tokenizer([prompt], padding=True, return_tensors='pt')59    input_length = len(inputs['input_ids'][0])60    for k, v in inputs.items():61        inputs[k] = v.cuda()62    input_ids = inputs['input_ids']63    _, input_ids_seq_length = input_ids.shape[0], input_ids.shape[-1]64    if generation_config is None:65        generation_config = model.generation_config66    generation_config = copy.deepcopy(generation_config)67    model_kwargs = generation_config.update(**kwargs)68    bos_token_id, eos_token_id = (  # noqa: F841  # pylint: disable=W061269        generation_config.bos_token_id,70        generation_config.eos_token_id,71    )72    if isinstance(eos_token_id, int):73        eos_token_id = [eos_token_id]74    if additional_eos_token_id is not None:75        eos_token_id.append(additional_eos_token_id)76    has_default_max_length = kwargs.get(77        'max_length') is None and generation_config.max_length is not None78    if has_default_max_length and generation_config.max_new_tokens is None:79        warnings.warn(80            f"Using 'max_length''s default \81                ({repr(generation_config.max_length)}) \82                to control the generation length. "83            'This behaviour is deprecated and will be removed from the \84                config in v5 of Transformers -- we'85            ' recommend using `max_new_tokens` to control the maximum \86                length of the generation.',87            UserWarning,88        )89    elif generation_config.max_new_tokens is not None:90        generation_config.max_length = generation_config.max_new_tokens + \91            input_ids_seq_length92        if not has_default_max_length:93            logger.warn(  # pylint: disable=W490294                f"Both 'max_new_tokens' (={generation_config.max_new_tokens}) "95                f"and 'max_length'(={generation_config.max_length}) seem to "96                "have been set. 'max_new_tokens' will take precedence. "97                'Please refer to the documentation for more information. '98                '(https://huggingface.co/docs/transformers/main/'99                'en/main_classes/text_generation)',100                UserWarning,101            )102 103    if input_ids_seq_length >= generation_config.max_length:104        input_ids_string = 'input_ids'105        logger.warning(106            f'Input length of {input_ids_string} is {input_ids_seq_length}, '107            f"but 'max_length' is set to {generation_config.max_length}. "108            'This can lead to unexpected behavior. You should consider'109            " increasing 'max_new_tokens'.")110 111    # 2. Set generation parameters if not already defined112    logits_processor = logits_processor if logits_processor is not None \113        else LogitsProcessorList()114    stopping_criteria = stopping_criteria if stopping_criteria is not None \115        else StoppingCriteriaList()116 117    logits_processor = model._get_logits_processor(118        generation_config=generation_config,119        input_ids_seq_length=input_ids_seq_length,120        encoder_input_ids=input_ids,121        prefix_allowed_tokens_fn=prefix_allowed_tokens_fn,122        logits_processor=logits_processor,123    )124 125    stopping_criteria = model._get_stopping_criteria(126        generation_config=generation_config,127        stopping_criteria=stopping_criteria)128    logits_warper = model._get_logits_warper(generation_config)129 130    unfinished_sequences = input_ids.new(input_ids.shape[0]).fill_(1)131    scores = None132    while True:133        model_inputs = model.prepare_inputs_for_generation(134            input_ids, **model_kwargs)135        # forward pass to get next token136        outputs = model(137            **model_inputs,138            return_dict=True,139            output_attentions=False,140            output_hidden_states=False,141        )142 143        next_token_logits = outputs.logits[:, -1, :]144 145        # pre-process distribution146        next_token_scores = logits_processor(input_ids, next_token_logits)147        next_token_scores = logits_warper(input_ids, next_token_scores)148 149        # sample150        probs = nn.functional.softmax(next_token_scores, dim=-1)151        if generation_config.do_sample:152            next_tokens = torch.multinomial(probs, num_samples=1).squeeze(1)153        else:154            next_tokens = torch.argmax(probs, dim=-1)155 156        # update generated ids, model inputs, and length for next step157        input_ids = torch.cat([input_ids, next_tokens[:, None]], dim=-1)158        model_kwargs = model._update_model_kwargs_for_generation(159            outputs, model_kwargs, is_encoder_decoder=False)160        unfinished_sequences = unfinished_sequences.mul(161            (min(next_tokens != i for i in eos_token_id)).long())162 163        output_token_ids = input_ids[0].cpu().tolist()164        output_token_ids = output_token_ids[input_length:]165        for each_eos_token_id in eos_token_id:166            if output_token_ids[-1] == each_eos_token_id:167                output_token_ids = output_token_ids[:-1]168        response = tokenizer.decode(output_token_ids)169 170        yield response171        # stop when each sentence is finished172        # or if we exceed the maximum length173        if unfinished_sequences.max() == 0 or stopping_criteria(174                input_ids, scores):175            break176 177 178def on_btn_click():179    del st.session_state.messages180 181 182@st.cache_resource183def load_model():184    model = (AutoModelForCausalLM.from_pretrained(185        model_name_or_path,186        trust_remote_code=True).to(torch.bfloat16).cuda())187    tokenizer = AutoTokenizer.from_pretrained(model_name_or_path,188                                              trust_remote_code=True)189    return model, tokenizer190 191 192def prepare_generation_config():193    with st.sidebar:194        max_length = st.slider('Max Length',195                               min_value=8,196                               max_value=32768,197                               value=32768)198        top_p = st.slider('Top P', 0.0, 1.0, 0.8, step=0.01)199        temperature = st.slider('Temperature', 0.0, 1.0, 0.7, step=0.01)200        st.button('Clear Chat History', on_click=on_btn_click)201 202    generation_config = GenerationConfig(max_length=max_length,203                                         top_p=top_p,204                                         temperature=temperature)205 206    return generation_config207 208 209user_prompt = '<|im_start|>user\n{user}<|im_end|>\n'210robot_prompt = '<|im_start|>assistant\n{robot}<|im_end|>\n'211cur_query_prompt = '<|im_start|>user\n{user}<|im_end|>\n\212    <|im_start|>assistant\n'213 214 215def combine_history(prompt):216    messages = st.session_state.messages217    meta_instruction = ('You are a helpful, honest, '218                        'and harmless AI assistant.')219    total_prompt = f'<s><|im_start|>system\n{meta_instruction}<|im_end|>\n'220    for message in messages:221        cur_content = message['content']222        if message['role'] == 'user':223            cur_prompt = user_prompt.format(user=cur_content)224        elif message['role'] == 'robot':225            cur_prompt = robot_prompt.format(robot=cur_content)226        else:227            raise RuntimeError228        total_prompt += cur_prompt229    total_prompt = total_prompt + cur_query_prompt.format(user=prompt)230    return total_prompt231 232 233def main():234    st.title('internlm2_5-7b-chat-assistant')235 236    # torch.cuda.empty_cache()237    print('load model begin.')238    model, tokenizer = load_model()239    print('load model end.')240 241    generation_config = prepare_generation_config()242 243    # Initialize chat history244    if 'messages' not in st.session_state:245        st.session_state.messages = []246 247    # Display chat messages from history on app rerun248    for message in st.session_state.messages:249        with st.chat_message(message['role'], avatar=message.get('avatar')):250            st.markdown(message['content'])251 252    # Accept user input253    if prompt := st.chat_input('What is up?'):254        # Display user message in chat message container255 256        with st.chat_message('user', avatar='user'):257 258            st.markdown(prompt)259        real_prompt = combine_history(prompt)260        # Add user message to chat history261        st.session_state.messages.append({262            'role': 'user',263            'content': prompt,264            'avatar': 'user'265        })266 267        with st.chat_message('robot', avatar='assistant'):268 269            message_placeholder = st.empty()270            for cur_response in generate_interactive(271                    model=model,272                    tokenizer=tokenizer,273                    prompt=real_prompt,274                    additional_eos_token_id=92542,275                    device='cuda:0',276                    **asdict(generation_config),277            ):278                # Display robot response in chat message container279                message_placeholder.markdown(cur_response + 'โ–Œ')280            message_placeholder.markdown(cur_response)281        # Add robot response to chat history282        st.session_state.messages.append({283            'role': 'robot',284            'content': cur_response,  # pylint: disable=undefined-loop-variable285            'avatar': 'assistant',286        })287        torch.cuda.empty_cache()288 289 290if __name__ == '__main__':291    main()292 293