bigcode/bigcode-editor
54
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 