yhavinga/rosetta
5
1import time2 3import psutil4import streamlit as st5import torch6from langdetect import detect7from transformers import TextIteratorStreamer8 9from default_texts import default_texts10from generator import GeneratorFactory11 12device = torch.cuda.device_count() - 113 14TRANSLATION_EN_TO_NL = "translation_en_to_nl"15TRANSLATION_NL_TO_EN = "translation_nl_to_en"16 17GENERATOR_LIST = [18 {19 "model_name": "yhavinga/ul2-base-en-nl",20 "desc": "UL2 base en->nl",21 "task": TRANSLATION_EN_TO_NL,22 "split_sentences": False,23 },24 # {25 # "model_name": "yhavinga/ul2-large-en-nl",26 # "desc": "UL2 large en->nl",27 # "task": TRANSLATION_EN_TO_NL,28 # "split_sentences": False,29 # },30 {31 "model_name": "Helsinki-NLP/opus-mt-en-nl",32 "desc": "Opus MT en->nl",33 "task": TRANSLATION_EN_TO_NL,34 "split_sentences": True,35 },36 {37 "model_name": "Helsinki-NLP/opus-mt-nl-en",38 "desc": "Opus MT nl->en",39 "task": TRANSLATION_NL_TO_EN,40 "split_sentences": True,41 },42 # {43 # "model_name": "yhavinga/t5-small-24L-ccmatrix-multi",44 # "desc": "T5 small nl24 ccmatrix nl-en",45 # "task": TRANSLATION_NL_TO_EN,46 # "split_sentences": True,47 # },48 {49 "model_name": "yhavinga/longt5-local-eff-large-nl8-voc8k-ddwn-neddx2-nl-en",50 "desc": "Long t5 large-nl8 nl-en",51 "task": TRANSLATION_NL_TO_EN,52 "split_sentences": False,53 },54 # {55 # "model_name": "yhavinga/byt5-small-ccmatrix-en-nl",56 # "desc": "ByT5 small ccmatrix en->nl",57 # "task": TRANSLATION_EN_TO_NL,58 # "split_sentences": True,59 # },60 # {61 # "model_name": "yhavinga/t5-base-36L-ccmatrix-multi",62 # "desc": "T5 base nl36 ccmatrix en->nl",63 # "task": TRANSLATION_EN_TO_NL,64 # "split_sentences": True,65 # },66 # {67]68 69 70class StreamlitTextIteratorStreamer(TextIteratorStreamer):71 def __init__(72 self, output_placeholder, tokenizer, skip_prompt=False, **decode_kwargs73 ):74 super().__init__(tokenizer, skip_prompt, **decode_kwargs)75 self.output_placeholder = output_placeholder76 self.output_text = ""77 78 def on_finalized_text(self, text: str, stream_end: bool = False):79 self.output_text += text80 self.output_placeholder.markdown(self.output_text, unsafe_allow_html=True)81 super().on_finalized_text(text, stream_end)82 83 84def main():85 st.set_page_config( # Alternate names: setup_page, page, layout86 page_title="Rosetta en/nl", # String or None. Strings get appended with "โข Streamlit".87 layout="wide", # Can be "centered" or "wide". In the future also "dashboard", etc.88 initial_sidebar_state="auto", # Can be "auto", "expanded", "collapsed"89 page_icon="๐", # String, anything supported by st.image, or None.90 )91 92 if "generators" not in st.session_state:93 st.session_state["generators"] = GeneratorFactory(GENERATOR_LIST)94 generators = st.session_state["generators"]95 96 with open("style.css") as f:97 st.markdown(f"<style>{f.read()}</style>", unsafe_allow_html=True)98 99 st.sidebar.image("rosetta.png", width=200)100 st.sidebar.markdown(101 """# Rosetta102 Vertaal van en naar Engels"""103 )104 105 default_text = st.sidebar.radio(106 "Change default text",107 tuple(default_texts.keys()),108 index=0,109 )110 if default_text or "prompt_box" not in st.session_state:111 st.session_state["prompt_box"] = default_texts[default_text]["text"]112 113 # create a left and right column114 left, right = st.columns(2)115 text_area = left.text_area("Enter text", st.session_state.prompt_box, height=500)116 st.session_state["text"] = text_area117 118 # Sidebar parameters119 st.sidebar.title("Parameters:")120 num_beams = st.sidebar.number_input("Num beams", min_value=1, max_value=10, value=1)121 num_beam_groups = st.sidebar.number_input(122 "Num beam groups", min_value=1, max_value=10, value=1123 )124 length_penalty = st.sidebar.number_input(125 "Length penalty", min_value=0.0, max_value=2.0, value=1.0, step=0.1126 )127 diversity_penalty = st.sidebar.number_input(128 "Diversity penalty", min_value=0.0, max_value=2.0, value=0.1, step=0.1129 )130 st.sidebar.markdown(131 """For an explanation of the parameters, head over to the [Huggingface blog post about text generation](https://huggingface.co/blog/how-to-generate)132and the [Huggingface text generation interface doc](https://huggingface.co/transformers/main_classes/model.html?highlight=generate#transformers.generation_utils.GenerationMixin.generate).133"""134 )135 params = {136 "num_beams": num_beams,137 "num_beam_groups": num_beam_groups,138 "diversity_penalty": diversity_penalty if num_beam_groups > 1 else 0.0,139 "length_penalty": length_penalty if num_beams > 1 else 1.0,140 "early_stopping": True,141 }142 143 if left.button("Run"):144 memory = psutil.virtual_memory()145 146 language = detect(st.session_state.text)147 if language == "en":148 task = TRANSLATION_EN_TO_NL149 elif language == "nl":150 task = TRANSLATION_NL_TO_EN151 else:152 left.error(f"Language {language} not supported")153 return154 155 # Num beam groups should be a divisor of num beams156 if num_beams % num_beam_groups != 0:157 left.error("Num beams should be a multiple of num beam groups")158 return159 160 streaming_enabled = num_beams == 1161 if not streaming_enabled:162 left.markdown("*`num_beams > 1` so streaming is disabled*")163 164 for generator in generators.filter(task=task):165 model_container = right.container()166 model_container.markdown(f"๐งฎ **Model `{generator}`**")167 output_placeholder = model_container.empty()168 streamer = (169 StreamlitTextIteratorStreamer(output_placeholder, generator.tokenizer)170 if streaming_enabled171 else None172 )173 time_start = time.time()174 result, params_used = generator.generate(175 text=st.session_state.text, streamer=streamer, **params176 )177 time_end = time.time()178 time_diff = time_end - time_start179 180 if not streaming_enabled:181 right.write(result.replace("\n", " \n"))182 text_line = ", ".join([f"{k}={v}" for k, v in params_used.items()])183 right.markdown(f" ๐ *generated in {time_diff:.2f}s, `{text_line}`*")184 185 st.write(186 f"""187 ---188 *Memory: {memory.total / 10**9:.2f}GB, used: {memory.percent}%, available: {memory.available / 10**9:.2f}GB*189 """190 )191 192 193if __name__ == "__main__":194 main()195 