KN123/nl2sql-api
1
1from fastapi import FastAPI, HTTPException
2from fastapi.middleware.cors import CORSMiddleware
3from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
4from typing import List, Dict
5import time
6import datetime
7import uvicorn
8
9model = AutoModelForSeq2SeqLM.from_pretrained("KN123/nl2sql")
10tokenizer = AutoTokenizer.from_pretrained("KN123/nl2sql")
11
12def get_prompt(tables, question):
13 prompt = f"""convert question and table into SQL query. tables: {tables}. question: {question}"""
14 # print(prompt)
15 return prompt
16
17def prepare_input(question: str, tables: Dict[str, List[str]]):
18 tables = [f"""{table_name}({",".join(tables[table_name])})""" for table_name in tables]
19 # print(tables)
20 tables = ", ".join(tables)
21 # print(tables)
22 prompt = get_prompt(tables, question)
23 # print(prompt)
24 input_ids = tokenizer(prompt, max_length=512, return_tensors="pt").input_ids
25 # print(input_ids)
26 return input_ids
27
28def inference(question: str, tables: Dict[str, List[str]]) -> str:
29 input_data = prepare_input(question=question, tables=tables)
30 input_data = input_data.to(model.device)
31 outputs = model.generate(inputs=input_data, num_beams=10, top_k=10, max_length=512)
32 # print("Outputs", outputs)
33 result = tokenizer.decode(token_ids=outputs[0], skip_special_tokens=True)
34 return result
35
36app = FastAPI()
37app.add_middleware(
38 CORSMiddleware,
39 allow_origins=["*"], # Allows all origins
40 allow_credentials=True,
41 allow_methods=["GET", "POST", "PUT", "DELETE", "OPTIONS"], # Allows all methods
42 allow_headers=["*"], # Allows all headers
43)
44
45@app.get("/")
46def home():
47 return {
48 "message" : "Hello there! Everything is working fine!",
49 "api-version": "1.0.0",
50 "role": "nl2sql",
51 "description": "This api can be used to convert natural language to SQL given the human prompt, tables and the attributes."
52 }
53
54@app.get("/test-generate")
55def generate(text:str):
56 start = time.time()
57 res = inference("how many people with name jui and age less than 25", {
58 "people_name":["id","name"], "people_age": ["people_id","age"]
59 })
60 end = time.time()
61 total_time_taken = end - start
62 current_utc_datetime = datetime.datetime.now(datetime.timezone.utc)
63 current_date = datetime.date.today()
64 timezone_name = time.tzname[time.daylight]
65 print(res)
66 return {
67 "api_response": f"{res}",
68 "time_taken(s)": f"{total_time_taken}",
69 "request_details": {
70 "utc_datetime": f"{current_utc_datetime}",
71 "current_date": f"{current_date}",
72 "timezone_name": f"{timezone_name}"
73 }
74 }
75
76@app.post("/generate")
77def generate(request_body:Dict):
78 if 'text' not in request_body or 'tables' not in request_body:
79 raise HTTPException(status_code=400, detail="Missing 'text' or 'tables' in request body")
80
81 prompt = request_body['text']
82 tables = request_body['tables']
83
84 start = time.time()
85 res = inference(prompt, tables)
86 end = time.time()
87 total_time_taken = end - start
88 current_utc_datetime = datetime.datetime.now(datetime.timezone.utc)
89 current_date = datetime.date.today()
90 timezone_name = time.tzname[time.daylight]
91 print(res)
92 return {
93 "api_response": f"{res}",
94 "time_taken(s)": f"{total_time_taken}",
95 "request_details": {
96 "utc_datetime": f"{current_utc_datetime}",
97 "current_date": f"{current_date}",
98 "timezone_name": f"{timezone_name}"
99 }
100 }
101
102
103
104if __name__ == "__main__":
105 uvicorn.run(app, host="127.0.0.1", port=8000)