Team Ai
Apppublic

KN123/nl2sql-api

sourceHugging Faceapache-2.0updated 2y agoView on Hugging Face
1likes
app.py105 linesDownload Raw Back to root
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)