Team Ai
Apppublic

Fisharp/starcoder-playground

sourceHugging Faceupdated 3y agoView on Hugging Face
6likes
app.py248 linesDownload Raw Back to root
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