thangved/text2sql
0
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 