Team Ai
Apppublic

stjarvie/Generative-SQL-Test-Space

sourceHugging Faceopenrailupdated 3y agoView on Hugging Face
0likes
app.py49 linesDownload Raw Back to root
1import torch2from peft import PeftModel, PeftConfig3from transformers import AutoModelForCausalLM, AutoTokenizer4 5peft_model_id = f"stjarvie/bloom-1b7-sql-generation"6config = PeftConfig.from_pretrained(peft_model_id)7model = AutoModelForCausalLM.from_pretrained(8    config.base_model_name_or_path,9    return_dict=True,10    load_in_8bit=True,11    device_map="auto",12)13tokenizer = AutoTokenizer.from_pretrained(config.base_model_name_or_path)14 15# Load the Lora model16model = PeftModel.from_pretrained(model, peft_model_id)17 18 19def generate_prompt_inference(question: str, schema: str) -> str:20  prompt = f"### Question:\n{question}\n\n### Table Schema:\n{schema}\n\n### SQL Query:\n "21  return prompt22 23 24def make_inference(question, schema):25    batch = tokenizer(26        generate_prompt_inference(question, schema),27        return_tensors="pt",28    )29 30    with torch.cuda.amp.autocast():31        output_tokens = model.generate(**batch, max_new_tokens=200)32 33    return tokenizer.decode(output_tokens[0], skip_special_tokens=True)34 35 36if __name__ == "__main__":37    # make a gradio interface38    import gradio as gr39 40    gr.Interface(41        make_inference,42        [43            gr.inputs.Textbox(lines=2, label="Question"),44            gr.inputs.Textbox(lines=5, label="Table Schema"),45        ],46        gr.outputs.Textbox(label="Ad"),47        title="Generative-SQL-AI",48        description="This is a tool that generates SQL given a question and related SQL table schema..",49    ).launch()