Spaces:
Sleeping
Sleeping
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
) |