File size: 2,236 Bytes
ca9941d
 
 
 
 
dabf498
 
 
 
ca9941d
9a6bae9
dabf498
ca9941d
4cb4101
 
 
 
 
 
 
 
 
 
 
ca9941d
4cb4101
 
 
 
 
 
 
 
 
 
 
 
ca9941d
 
4cb4101
 
 
ca9941d
4cb4101
 
ca9941d
4cb4101
ca9941d
4cb4101
 
ca9941d
4cb4101
 
ca9941d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4cb4101
 
ca9941d
 
 
 
dabf498
4cb4101
 
9a6bae9
4cb4101
ca9941d
 
 
 
 
4cb4101
9a6bae9
dabf498
4cb4101
ca9941d
 
4cb4101
 
ca9941d
4cb4101
ca9941d
 
4cb4101
 
 
 
ca9941d
 
4cb4101
 
 
ca9941d
 
 
4cb4101
 
 
 
ca9941d
 
 
 
4cb4101
 
 
 
ca9941d
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
import os

# Disable Gradio SSR (better for Render)
os.environ["GRADIO_SSR_MODE"] = "False"

import gradio as gr
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
import torch


MODEL_NAME = "gaussalgo/T5-LM-Large-text2sql-spider"


tokenizer = None
model = None

device = "cuda" if torch.cuda.is_available() else "cpu"


def load_model():

    global tokenizer, model

    if model is None:

        print("Loading model...")

        tokenizer = AutoTokenizer.from_pretrained(
            MODEL_NAME
        )

        model = AutoModelForSeq2SeqLM.from_pretrained(
            MODEL_NAME
        )

        model.to(device)

        model.eval()

        print(f"Model ready on {device}")



def generate_sql(question, context):

    try:

        load_model()


        input_text = f"{question} | {context}"


        inputs = tokenizer(
            input_text,
            return_tensors="pt",
            max_length=512,
            truncation=True
        ).to(device)


        with torch.no_grad():

            outputs = model.generate(
                **inputs,
                max_new_tokens=128,
                num_beams=4,
                early_stopping=True
            )


        sql = tokenizer.decode(
            outputs[0],
            skip_special_tokens=True
        )


        return sql


    except Exception as e:

        return f"Error: {str(e)}"




with gr.Blocks() as demo:


    gr.Markdown(
        "# NL2SQL API\nGenerate SQL from Natural Language"
    )


    with gr.Row():

        question = gr.Textbox(
            label="Question",
            placeholder="Example: Find all users"
        )


        context = gr.Textbox(
            label="Database Schema",
            placeholder="Example: users(id,name,email)"
        )


    output = gr.Textbox(
        label="Generated SQL",
        lines=5
    )


    btn = gr.Button(
        "Generate SQL"
    )


    btn.click(
        fn=generate_sql,
        inputs=[
            question,
            context
        ],
        outputs=output
    )



if __name__ == "__main__":

    demo.launch(
        server_name="0.0.0.0",
        server_port=int(
            os.environ.get("PORT",7860)
        ),
        share=False
    )