CoolFace
Apppublic

Chaitanya182004/nl2sql-api

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes
app.py140 linesDownload Raw Back to root
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    )