CoolFace
Apppublic

simon-clmtd/exbert

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
main.py180 linesDownload Raw Back to server
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