Chaitanya182004/nl2sql-api
0
1import os2 3# Disable Gradio SSR (better for Render)4os.environ["GRADIO_SSR_MODE"] = "False"5 6import gradio as gr7from transformers import AutoTokenizer, AutoModelForSeq2SeqLM8import torch9 10 11MODEL_NAME = "gaussalgo/T5-LM-Large-text2sql-spider"12 13 14tokenizer = None15model = None16 17device = "cuda" if torch.cuda.is_available() else "cpu"18 19 20def load_model():21 22 global tokenizer, model23 24 if model is None:25 26 print("Loading model...")27 28 tokenizer = AutoTokenizer.from_pretrained(29 MODEL_NAME30 )31 32 model = AutoModelForSeq2SeqLM.from_pretrained(33 MODEL_NAME34 )35 36 model.to(device)37 38 model.eval()39 40 print(f"Model ready on {device}")41 42 43 44def generate_sql(question, context):45 46 try:47 48 load_model()49 50 51 input_text = f"{question} | {context}"52 53 54 inputs = tokenizer(55 input_text,56 return_tensors="pt",57 max_length=512,58 truncation=True59 ).to(device)60 61 62 with torch.no_grad():63 64 outputs = model.generate(65 **inputs,66 max_new_tokens=128,67 num_beams=4,68 early_stopping=True69 )70 71 72 sql = tokenizer.decode(73 outputs[0],74 skip_special_tokens=True75 )76 77 78 return sql79 80 81 except Exception as e:82 83 return f"Error: {str(e)}"84 85 86 87 88with gr.Blocks() as demo:89 90 91 gr.Markdown(92 "# NL2SQL API\nGenerate SQL from Natural Language"93 )94 95 96 with gr.Row():97 98 question = gr.Textbox(99 label="Question",100 placeholder="Example: Find all users"101 )102 103 104 context = gr.Textbox(105 label="Database Schema",106 placeholder="Example: users(id,name,email)"107 )108 109 110 output = gr.Textbox(111 label="Generated SQL",112 lines=5113 )114 115 116 btn = gr.Button(117 "Generate SQL"118 )119 120 121 btn.click(122 fn=generate_sql,123 inputs=[124 question,125 context126 ],127 outputs=output128 )129 130 131 132if __name__ == "__main__":133 134 demo.launch(135 server_name="0.0.0.0",136 server_port=int(137 os.environ.get("PORT",7860)138 ),139 share=False140 )