simon-clmtd/exbert
0
1import argparse2import numpy as np3import connexion4from flask_cors import CORS5from flask import render_template, redirect, send_from_directory6 7import utils.path_fixes as pf8from utils.f import ifnone9 10from model_api import get_details11 12app = connexion.FlaskApp(__name__, static_folder="client/dist", specification_dir=".")13flask_app = app.app14CORS(flask_app)15 16parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter)17parser.add_argument("--debug", action="store_true", help=" Debug mode")18parser.add_argument("--port", default=5051, help="Port to run the app. ")19 20# Flask main routes21@app.route("/")22def hello_world():23 return redirect("client/exBERT.html")24 25# send everything from client as static content26@app.route("/client/<path:path>")27def send_static_client(path):28 """ serves all files from ./client/ to ``/client/<path:path>``29 30 :param path: path from api call31 """32 return send_from_directory(str(pf.CLIENT_DIST), path)33 34# ======================================================================35## CONNEXION API ##36# ======================================================================37def get_model_details(**request):38 """Get important information about a model, like the number of layers and heads39 40 Args:41 request['model']: The model name42 43 Returns:44 {45 status: 200,46 payload: {47 nlayers (int)48 nheads (int)49 }50 }51 """52 mname = request['model']53 deets = get_details(mname)54 55 info = deets.config56 nlayers = info.num_hidden_layers57 nheads = info.num_attention_heads58 59 payload_out = {60 "nlayers": nlayers,61 "nheads": nheads,62 }63 64 return {65 "status": 200,66 "payload": payload_out,67 }68 69def get_attentions_and_preds(**request):70 """For a sentence, at a layer, get the attentions and predictions71 72 Args:73 request['model']: Model name74 request['sentence']: Sentence to get the attentions for75 request['layer']: Which layer to extract from76 77 Returns:78 {79 status: 20080 payload: {81 aa: {82 att: Array((nheads, ntoks, ntoks))83 left: [{84 text (str), 85 topk_words (List[str]),86 topk_probs (List[float])87 }, ...]88 right: [{89 text (str), 90 topk_words (List[str]),91 topk_probs (List[float])92 }, ...]93 }94 }95 }96 """97 model = request["model"]98 details = get_details(model)99 100 sentence = request["sentence"]101 layer = int(request["layer"])102 103 deets = details.from_sentence(sentence)104 105 payload_out = deets.to_json(layer)106 107 return {108 "status": 200,109 "payload": payload_out110 }111 112def update_masked_attention(**request):113 """From tokens and indices of what should be masked, get the attentions and predictions114 115 payload = request['payload']116 117 Args:118 payload['model'] (str): Model name119 payload['tokens'] (List[str]): Tokens to pass through the model120 payload['sentence'] (str): Original sentence the tokens came from121 payload['mask'] (List[int]): Which indices to mask122 payload['layer'] (int): Which layer to extract information from123 124 Returns:125 {126 status: 200127 payload: {128 aa: {129 att: Array((nheads, ntoks, ntoks))130 left: [{131 text (str), 132 topk_words (List[str]),133 topk_probs (List[float])134 }, ...]135 right: [{136 text (str), 137 topk_words (List[str]),138 topk_probs (List[float])139 }, ...]140 }141 }142 }143 """144 payload = request["payload"]145 146 model = payload['model']147 details = get_details(model)148 149 tokens = payload["tokens"]150 sentence = payload["sentence"]151 mask = payload["mask"]152 layer = int(payload["layer"])153 154 MASK = details.tok.mask_token155 mask_tokens = lambda toks, maskinds: [156 t if i not in maskinds else ifnone(MASK, t) for (i, t) in enumerate(toks)157 ]158 159 token_inputs = mask_tokens(tokens, mask)160 161 deets = details.from_tokens(token_inputs, sentence)162 payload_out = deets.to_json(layer)163 164 return {165 "status": 200,166 "payload": payload_out,167 }168 169app.add_api("swagger.yaml")170 171# Setup code172if __name__ != "__main__":173 print("SETTING UP ENDPOINTS")174 175# Then deploy app176else:177 args, _ = parser.parse_known_args()178 print("Initiating app")179 app.run(port=args.port, use_reloader=False, debug=args.debug)180 