CoolFace
Apppublic

namnh113/Code_Summarization

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
app.py151 linesDownload Raw Back to root
1import os2import streamlit as st3from langchain.llms import HuggingFaceHub4from models import return_sum_models5 6class LLM_Langchain():7    def __init__(self):8        st.header('๐Ÿฆœ Code summarization')9        st.warning("Warning: input function needs cleaning and may take long to be processed at first time")10        st.warning("Note: you should not copy the whole function from IDE, the \"\\n\" character needs typing by hand")11        st.info("Reference: [CodeT5](https://arxiv.org/abs/2109.00859), [The Vault](https://arxiv.org/abs/2305.06156), [CodeXGLUE](https://arxiv.org/abs/2102.04664)")12        st.info("About me: namnh113")13        14 15        self.API_KEY = st.sidebar.text_input(16            'API key',17            type='password',18            help="Type in your HuggingFace API key to use this app")19 20            21        self.model_parent = st.sidebar.selectbox(22            label = "Choose language",23            options = ["python", "java", "javascript", "php", "ruby", "go", "cpp"],24            help="Choose languages",25            )26 27        if self.model_parent is None:28            model_name_visibility = True29        else:30            model_name_visibility = False31            32        model_name = return_sum_models(self.model_parent)33        list_model = [model_name]34        if self.model_parent in ["python", "java"]:35            list_model += [model_name+"_v2"]36        if self.model_parent != "cpp":37            list_model += ["Salesforce/codet5-base-multi-sum", f"Salesforce/codet5-base-codexglue-sum-{self.model_parent}"]38        39        self.checkpoint = st.sidebar.selectbox(40            label = "Choose model (namnh113/... is my model)",41            options = list_model,42            help="Model used to predict",43            disabled=model_name_visibility44            )45 46        self.max_new_tokens = st.sidebar.slider(47            label="Token Length",48            min_value=32,49            max_value=248,50            step=4,51            value=128,52            help="Set the max tokens to get accurate results"53            )54 55        self.num_beams = st.sidebar.slider(56            label="num beams",57            min_value=1,58            max_value=10,59            step=1,60            value=2,61            help="Set num beam"62            )63        64        self.top_k = st.sidebar.slider(65            label="top k",66            min_value=1,67            max_value=50,68            step=1,69            value=30,70            help="Set the top_k"71            )72        73        self.top_p = st.sidebar.slider(74            label="top p",75            min_value=0.1,76            max_value=1.0,77            step=0.05,78            value=0.95,79            help="Set the top_p"80            )81 82 83        self.model_kwargs = {84            "max_new_tokens": self.max_new_tokens,85            "top_k": self.top_k,86            "top_p": self.top_p,87            "num_beams": self.num_beams88        }89 90        os.environ['HUGGINGFACEHUB_API_TOKEN'] = self.API_KEY91 92 93    def generate_response(self, input_text):94        95        input_text = "Summarize " + self.model_parent.capitalize() + ": " + input_text96        llm = HuggingFaceHub(97            repo_id = self.checkpoint,98            model_kwargs = self.model_kwargs99        )100 101        return llm(input_text)102    103    104 105    def form_data(self):106        # with st.form('my_form'):107            try:108                if not self.API_KEY.startswith('hf_'):109                    st.warning('Please enter your API key!', icon='โš ')110                111                112                if "messages" not in st.session_state:113                    st.session_state.messages = []114 115                st.write(f"You are using {self.checkpoint} model")116 117                for message in st.session_state.messages:118                    with st.chat_message(message.get('role')):119                        st.write(message.get("content"))120                text = st.chat_input(disabled=False)121                122                if text:123                    st.session_state.messages.append(124                        {125                            "role":"user",126                            "content": text127                        }128                    )129                    with st.chat_message("user"):130                        st.write(text)131                    132                    if text.lower() == "clear":133                        del st.session_state.messages134                        return135                        136                    result = self.generate_response(text)137                    result = result.replace(' * ', '\n* ')138                    st.session_state.messages.append(139                        {140                            "role": "assistant",141                            "content": result142                        }143                    )144                    with st.chat_message('assistant'):145                        st.markdown(result)146                147            except Exception as e:148                st.error(e, icon="๐Ÿšจ")149 150model = LLM_Langchain()151model.form_data()