Team Ai
Apppublic

bigcode/bigcode-editor

sourceHugging Faceupdated 3y agoView on Hugging Face
54likes
app.py151 linesDownload Raw Back to root
1import sys2from typing import List3import traceback4import os5import base646import json7import pprint8 9from huggingface_hub import Repository10from text_generation import Client11 12from requests.exceptions import ReadTimeout13 14PORT = 786015 16# TODO: implement maximum length (currently, each iteration is limited by the slider-specified max length, but this can be iterated, or long code entered into the editor, to get really long documents17# if os.path.exists('unlock'):18#     # create an 'unlock' file (not checked into Git) locally to get full context lengths19#     MAX_LENGTH = 819220# else:21#     # set to a shorter value to prevent long contexts and make the demo more efficient22#     MAX_LENGTH = 102423# TRUNCATION_MESSAGE = f'warning: This demo is limited to {MAX_LENGTH} tokens in the document for efficiency.'24TRUNCATION_MESSAGE = f'TODO'25 26HF_TOKEN = os.environ.get("HF_TOKEN", None)27API_URL = os.environ.get("API_URL")28 29with open("./HHH_prompt.txt", "r") as f:30    HHH_PROMPT = f.read() + "\n\n"31 32# used by the model33FIM_PREFIX = "<fim_prefix>"34FIM_MIDDLE = "<fim_middle>"35FIM_SUFFIX = "<fim_suffix>"36END_OF_TEXT = "<|endoftext|>"37 38# used to mark infill locations in the editor39FIM_INDICATOR = "<infill>"40 41client = Client(42    API_URL, headers={"Authorization": f"Bearer {HF_TOKEN}"},43)44 45from fastapi import FastAPI, Request46from fastapi.staticfiles import StaticFiles47from fastapi.responses import FileResponse, StreamingResponse48app = FastAPI(docs_url=None, redoc_url=None)49app.mount("/static", StaticFiles(directory="static"), name="static")50 51@app.head("/")52@app.get("/")53def index() -> FileResponse:54    return FileResponse(path="static/index.html", media_type="text/html")55 56def generate(prefix, suffix=None, temperature=0.9, max_new_tokens=256, top_p=0.95, repetition_penalty=1.0):57    # TODO: deduplicate code between this and `infill`58    temperature = float(temperature)59    if temperature < 1e-2:60        temperature = 1e-261    top_p = float(top_p)62    63    generate_kwargs = dict(64        temperature=temperature,65        max_new_tokens=max_new_tokens,66        top_p=top_p,67        repetition_penalty=repetition_penalty,68        do_sample=True,69        seed=42,70    )71 72    fim_mode = suffix is not None73    74    if suffix is not None:75        prompt = f"{FIM_PREFIX}{prefix}{FIM_SUFFIX}{suffix}{FIM_MIDDLE}"76    else:77        prompt = prefix78    output = client.generate(prompt, **generate_kwargs)79    generated_text = output.generated_text80    # TODO: set this based on stop reason from client.generate81    truncated = False82    while generated_text.endswith(END_OF_TEXT):83        generated_text = generated_text[:-len(END_OF_TEXT)]84    generation = {85        'truncated': truncated,86    }87    if fim_mode:88        generation['type'] = 'infill'89        generation['text'] = prefix + generated_text + suffix90        generation['parts'] = [prefix, suffix]91        generation['infills'] = [generated_text]92    else:93        generation['type'] = 'generate'94        generation['text'] = prompt + generated_text95        generation['parts'] = [prompt]96    return generation97 98@app.get('/generate')99async def generate_maybe(info: str):100    # info is a base64-encoded, url-escaped json string (since GET doesn't support a body, and POST leads to CORS issues)101    # fix padding, following https://stackoverflow.com/a/9956217/1319683102    info = base64.urlsafe_b64decode(info + '=' * (4 - len(info) % 4)).decode('utf-8')103    form = json.loads(info)104    prompt = form['prompt']105    length_limit = int(form['length'])106    temperature = float(form['temperature'])107    try:108        generation = generate(prompt, temperature=temperature, max_new_tokens=length_limit, top_p=0.95, repetition_penalty=1.0)109        if generation['truncated']:110            message = TRUNCATION_MESSAGE 111        else:112            message = ''113        return {'result': 'success', 'type': 'generate', 'prompt': prompt, 'text': generation['text'], 'message': message}114    except ReadTimeout as e:115        print(e)116        return {'result': 'error', 'type': 'generate', 'prompt': prompt, 'message': f'Request timed out.'}117    except Exception as e:118        traceback.print_exception(*sys.exc_info())119        return {'result': 'error', 'type': 'generate', 'prompt': prompt, 'message': f'Error: {e}.'}120 121@app.get('/infill')122async def infill_maybe(info: str):123    # info is a base64-encoded, url-escaped json string (since GET doesn't support a body, and POST leads to CORS issues)124    # fix padding, following https://stackoverflow.com/a/9956217/1319683125    info = base64.urlsafe_b64decode(info + '=' * (4 - len(info) % 4)).decode('utf-8')126    form = json.loads(info)127    length_limit = int(form['length'])128    temperature = float(form['temperature'])129    try:130        if len(form['parts']) > 2:131            return {'result': 'error', 'text': ''.join(form['parts']), 'type': 'infill', 'message': f"error: Only a single <infill> token is supported!"}132        elif len(form['parts']) == 1:133            return {'result': 'error', 'text': ''.join(form['parts']), 'type': 'infill', 'message': f"error: Must have an <infill> token present!"}134        prefix, suffix = form['parts']135        generation = generate(prefix, suffix=suffix, temperature=temperature, max_new_tokens=length_limit, top_p=0.95, repetition_penalty=1.0)136        generation['result'] = 'success'137        if generation['truncated']:138            generation['message'] = TRUNCATION_MESSAGE139        else:140            generation['message'] = ''141        return generation142    except ReadTimeout as e:143        print(e)144        return {'result': 'error', 'type': 'generate', 'prompt': prompt, 'message': f'Request timed out.'}145    except Exception as e:146        traceback.print_exception(*sys.exc_info())147        return {'result': 'error', 'type': 'infill', 'message': f'Error: {e}.'}148 149if __name__ == "__main__":150    app.run(host='0.0.0.0', port=PORT, threaded=False)151