stjarvie/Generative-SQL-Test-Space
0
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()