CoolFace
Apppublic

attenty/gec

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py51 linesDownload Raw Back to root
1from fastapi import FastAPI, HTTPException2import huggingface_hub3import torch4from fairseq.models.bart import BARTModel5import os6import re7 8EXPECTED_TOKEN = os.environ.get("EXPECTED_TOKEN")9REPO_NAME = os.environ.get("REPO_NAME")10REPO_NAME_HUGG = os.environ.get("REPO_NAME_HUGG")11 12app = FastAPI()13 14path_model = huggingface_hub.hf_hub_download(repo_id=f'{REPO_NAME}/{REPO_NAME_HUGG}' , filename='model/checkpoint_best.pt', token=True)15 16files = ['dict.src.txt', 'dict.tgt.txt', 'preprocess.log',17         'train.src-tgt.src.bin', 'train.src-tgt.src.idx', 'train.src-tgt.tgt.bin', 'train.src-tgt.tgt.idx',18         'valid.src-tgt.src.bin', 'valid.src-tgt.src.idx', 'valid.src-tgt.tgt.bin', 'valid.src-tgt.tgt.idx']19 20for file in files:21    path_data = huggingface_hub.hf_hub_download(repo_id=f'{REPO_NAME}/{REPO_NAME_HUGG}' , filename=file, subfolder='gec_data-bin_ptbr', token=True)22 23bart = BARTModel.from_pretrained(24    '/'.join(path_model.split('/')[:-1]),25    checkpoint_file='checkpoint_best.pt',26    data_name_or_path='/'.join(path_data.split('/')[:-1])27)28 29bart.eval()30 31 32def posprocessing(frase: str):33    34    nova_frase = re.sub(r'\s+([.,!?;:])', r'\1', frase)35    36    return nova_frase37 38 39@app.post('/inference')40async def predict(frase: str, token: str):41 42    if token != EXPECTED_TOKEN:43        raise HTTPException(status_code=401, detail="Token inválido")44 45    with torch.no_grad():46 47        result = bart.sample([frase], beam=1)48        49    return posprocessing(result[0])50 51