namnh113/Code_Summarization
1
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()