Fisharp/starcoder-playground
6
1import sys2import os3import logging as log4from typing import Generator5 6import gradio as gr7from gradio.themes.utils import sizes8from text_generation import Client9from src.request import StarCoderRequest, StarCoderRequestConfig10 11from src.utils import (12 get_file_as_string,13 get_sections,14 get_url_from_env_or_default_path,15 preview16)17from constants import (18 FIM_MIDDLE,19 FIM_PREFIX,20 FIM_SUFFIX,21 END_OF_TEXT,22 MIN_TEMPERATURE,23)24from settings import (25 DEFAULT_PORT,26 DEFAULT_STARCODER_API_PATH,27 DEFAULT_STARCODER_BASE_API_PATH,28)29 30HF_TOKEN = os.environ.get("HF_TOKEN", None)31# Gracefully exit the app if the HF_TOKEN is not set,32# printing to system `errout` the error (instead of raising an exception)33# and the expected behavior34if not HF_TOKEN:35 ERR_MSG = """36 Please set the HF_TOKEN environment variable with your Hugging Face API token.37 You can get one by signing up at https://huggingface.co/join and then visiting38 https://huggingface.co/settings/tokens."""39 print(ERR_MSG, file=sys.stderr)40 # gr.errors.GradioError(ERR_MSG)41 # gr.close_all(verbose=False)42 sys.exit(1)43 44API_URL_STAR = get_url_from_env_or_default_path("STARCODER_API", DEFAULT_STARCODER_API_PATH)45API_URL_BASE = get_url_from_env_or_default_path("STARCODER_BASE_API", DEFAULT_STARCODER_BASE_API_PATH)46 47preview("StarCoder Model URL", API_URL_STAR)48preview("StarCoderBase Model URL", API_URL_BASE)49preview("HF Token", HF_TOKEN, ofuscate=True)50 51_styles = get_file_as_string("styles.css")52_script = get_file_as_string("community-btn.js")53_sharing_icon_svg = get_file_as_string("community-icon.svg")54_loading_icon_svg = get_file_as_string("loading-icon.svg")55 56# Loads the whole content of the ./README.md file57# slicing/unpacking its different sections into their proper variables58readme_file_content = get_file_as_string("README.md", path='./')59(60 manifest,61 description,62 disclaimer,63 formats,64) = get_sections(readme_file_content, "---", up_to=4)65 66theme = gr.themes.Monochrome(67 primary_hue="indigo",68 secondary_hue="blue",69 neutral_hue="slate",70 radius_size=sizes.radius_sm,71 font=[72 gr.themes.GoogleFont("IBM Plex Sans", [400, 600]),73 "ui-sans-serif",74 "system-ui",75 "sans-serif",76 ],77 text_size=sizes.text_lg,78)79 80HEADERS = {81 "Authorization": f"Bearer {HF_TOKEN}",82}83client_star = Client(API_URL_STAR, headers=HEADERS)84client_base = Client(API_URL_BASE, headers=HEADERS)85 86def get_tokens_collector(request: StarCoderRequest) -> Generator[str, None, None]:87 88 model_client = client_star if request.settings.version == "StarCoder" else client_base89 stream = model_client.generate_stream(request.prompt, **request.settings.kwargs())90 for response in stream:91 # print(response.token.id, response.token.text)92 # if token.text != END_OF_TEXT:93 if response.token.id != 0:94 yield response.token.text95 96def get_tokens_accumulator(request: StarCoderRequest) -> Generator[str, None, None]:97 # start with the prefix (if in fim_mode)98 output = request.prefix if request.fim_mode else request.prompt99 for token in get_tokens_collector(request=request):100 output += token101 yield output102 # after the last token, append the suffix (if in fim_mode)103 if request.fim_mode:104 output += request.suffix105 yield output106 # Append an extra line at the end107 yield output + '\n'108 109def get_tokens_linker(request: StarCoderRequest) -> str:110 return "".join(list(get_tokens_collector(request)))111 112def generate(113 prompt: str,114 temperature = 0.9,115 max_new_tokens = 256,116 top_p = 0.95,117 repetition_penalty = 1.0,118 version = "StarCoder",119 ) -> Generator[str, None, None]:120 request = StarCoderRequest(121 prompt=prompt,122 settings=StarCoderRequestConfig(123 version=version,124 temperature=temperature,125 max_new_tokens=max_new_tokens,126 top_p=top_p,127 repetition_penalty=repetition_penalty,128 )129 )130 yield from get_tokens_accumulator(request)131 132def process_example(133 prompt: str,134 temperature = 0.9,135 max_new_tokens = 256,136 top_p = 0.95,137 repetition_penalty = 1.0,138 version = "StarCoder",139 ) -> Generator[str, None, None]:140 request = StarCoderRequest(141 prompt=prompt,142 settings=StarCoderRequestConfig(143 version=version,144 temperature=temperature,145 max_new_tokens=max_new_tokens,146 top_p=top_p,147 repetition_penalty=repetition_penalty,148 )149 )150 yield from get_tokens_linker(request)151 152# todo: move it into the README too153examples = [154 "X_train, y_train, X_test, y_test = train_test_split(X, y, test_size=0.1)\n\n# Train a logistic regression model, predict the labels on the test set and compute the accuracy score",155 "// Returns every other value in the array as a new array.\nfunction everyOther(arr) {",156 "def alternating(list1, list2):\n results = []\n for i in range(min(len(list1), len(list2))):\n results.append(list1[i])\n results.append(list2[i])\n if len(list1) > len(list2):\n <FILL_HERE>\n else:\n results.extend(list2[i+1:])\n return results",157]158 159with gr.Blocks(theme=theme, analytics_enabled=False, css=_styles) as demo:160 with gr.Column():161 gr.Markdown(description)162 with gr.Row():163 with gr.Column():164 instruction = gr.Textbox(165 placeholder="Enter your code here",166 label="Code",167 elem_id="q-input",168 )169 submit = gr.Button("Generate", variant="primary")170 output = gr.Code(elem_id="q-output", lines=30)171 with gr.Row():172 with gr.Column():173 with gr.Accordion("Advanced settings", open=False):174 with gr.Row():175 column_1, column_2 = gr.Column(), gr.Column()176 with column_1:177 temperature = gr.Slider(178 label="Temperature",179 value=0.2,180 minimum=0.0,181 maximum=1.0,182 step=0.05,183 interactive=True,184 info="Higher values produce more diverse outputs",185 )186 max_new_tokens = gr.Slider(187 label="Max new tokens",188 value=256,189 minimum=0,190 maximum=8192,191 step=64,192 interactive=True,193 info="The maximum numbers of new tokens",194 )195 with column_2:196 top_p = gr.Slider(197 label="Top-p (nucleus sampling)",198 value=0.90,199 minimum=0.0,200 maximum=1,201 step=0.05,202 interactive=True,203 info="Higher values sample more low-probability tokens",204 )205 repetition_penalty = gr.Slider(206 label="Repetition penalty",207 value=1.2,208 minimum=1.0,209 maximum=2.0,210 step=0.05,211 interactive=True,212 info="Penalize repeated tokens",213 )214 with gr.Column():215 version = gr.Dropdown(216 ["StarCoderBase", "StarCoder"],217 value="StarCoder",218 label="Version",219 info="",220 )221 gr.Markdown(disclaimer)222 with gr.Group(elem_id="share-btn-container"):223 community_icon = gr.HTML(_sharing_icon_svg, visible=True)224 loading_icon = gr.HTML(_loading_icon_svg, visible=True)225 share_button = gr.Button(226 "Share to community", elem_id="share-btn", visible=True227 )228 gr.Examples(229 examples=examples,230 inputs=[instruction],231 cache_examples=False,232 fn=process_example,233 outputs=[output],234 )235 gr.Markdown(formats)236 237 submit.click(238 generate,239 inputs=[instruction, temperature, max_new_tokens, top_p, repetition_penalty, version],240 outputs=[output],241 # preprocess=False,242 max_batch_size=8,243 show_progress=True244 )245 share_button.click(None, [], [], _js=_script)246 247demo.queue(concurrency_count=16).launch(debug=True, server_port=DEFAULT_PORT)248 