Team Ai
Apppublic

thangved/text2sql

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
main.py69 linesDownload Raw Back to root
1import torch2 3from fastapi import FastAPI4from pydantic import BaseModel5from transformers import T5ForConditionalGeneration, T5Tokenizer6from fastapi.middleware.cors import CORSMiddleware7 8app = FastAPI()9app.add_middleware(10    CORSMiddleware,11    allow_origins=['*'],12    allow_credentials=True,13    allow_methods=["*"],14    allow_headers=["*"],15)16 17device = torch.device("cuda" if torch.cuda.is_available() else "cpu")18 19 20model = T5ForConditionalGeneration.from_pretrained(21    "thangved/text2sql").to(device)  # type: ignore22tokenizer = T5Tokenizer.from_pretrained("t5-small")23 24 25def predict(context, question):26    inputs = tokenizer(f"query for: {question}? ",27                       f"tables: {context}",28                       max_length=200,29                       padding="max_length",30                       truncation=True,31                       pad_to_max_length=True,32                       add_special_tokens=True)33 34    input_ids = torch.tensor(35        inputs["input_ids"], dtype=torch.long).to(device).unsqueeze(0)36    attention_mask = torch.tensor(37        inputs["attention_mask"], dtype=torch.long).to(device).unsqueeze(0)38 39    outputs = model.generate(40        input_ids=input_ids, attention_mask=attention_mask, max_length=128)41    answer = tokenizer.decode(42        outputs.flatten(), skip_special_tokens=True)  # type: ignore43    return answer44 45 46class Text2SqlReq(BaseModel):47    context: str48    question: str49 50 51class Text2SqlRes(BaseModel):52    answer: str53 54 55class StatusRes(BaseModel):56    status: int57 58 59@app.post('/text2sql', summary='Text 2 SQL', tags=['Text 2 SQL'], response_model=Text2SqlRes)60async def text2sql(body: Text2SqlReq):61    answer = predict(body.context, body.question)62 63    return Text2SqlRes(answer=answer)64 65 66@app.get('/status', summary='Check server status', tags=['Status'], response_model=StatusRes)67async def status():68    return StatusRes(status=200)69