karthiknitt/text_to_sql_with_context
0
1import gradio as gr2from transformers import AutoModelForCausalLM, AutoTokenizer3import sqlite34import numpy as np5import faiss6import torch7import json8import os9 10# Load the model and tokenizer directly from the Hugging Face hub11model_name_or_path = "meta-llama/Meta-Llama-3-8B-Instruct"12 13# Get the Hugging Face token from environment variables14hf_token = os.getenv("HUGGINGFACE_HUB_TOKEN")15tokenizer = AutoTokenizer.from_pretrained(model_name_or_path,token=hf_token)16model = AutoModelForCausalLM.from_pretrained(model_name_or_path,token=hf_token)17 18# Functions provided19def generate_response(prompt, model):20 encoded_input = tokenizer(prompt, return_tensors="pt", add_special_tokens=True)21 model_inputs = encoded_input.to('cpu') # Ensure the model runs on CPU22 generated_ids = model.generate(**model_inputs, max_new_tokens=100, do_sample=True, pad_token_id=tokenizer.eos_token_id, eos_token_id=tokenizer.eos_token_id)23 decoded_output = tokenizer.batch_decode(generated_ids)24 return decoded_output[0].replace(prompt, "")25 26def get_embedding(text):27 inputs = tokenizer(text, return_tensors="pt").input_ids.to("cpu") # Ensure the model runs on CPU28 with torch.no_grad():29 outputs = model(inputs)30 embedding = outputs[0].mean(dim=1).squeeze()31 return embedding.cpu().numpy()32 33def find_relevant_columns(user_query, column_descriptions):34 column_embeddings = []35 column_keys = []36 37 for table, columns in column_descriptions.items():38 for column, (data_type, description) in columns.items():39 column_key = f"{table}.{column}"40 column_keys.append(column_key)41 embedding = get_embedding(description)42 column_embeddings.append(embedding)43 44 column_embeddings = np.array(column_embeddings)45 d = column_embeddings.shape[1]46 index = faiss.IndexFlatL2(d)47 index.add(column_embeddings)48 49 query_embedding = get_embedding(user_query)50 k = 451 D, I = index.search(np.array([query_embedding]), k)52 relevant_columns = [column_keys[i] for i in I[0]]53 return relevant_columns54 55def construct_create_statement(relevant_columns, column_descriptions):56 tables = {}57 for col in relevant_columns:58 table, column = col.split('.')59 if table not in tables:60 tables[table] = []61 tables[table].append(column)62 create_statements = []63 for table, columns in tables.items():64 column_defs = ', '.join([f"{col} {column_descriptions[table][col][0]}" for col in columns])65 statement = f"CREATE TABLE {table} ({column_defs});"66 create_statements.append(statement)67 return create_statements68 69def generate_sql_query(user_query, relevant_columns, column_descriptions):70 system_prompt = "You are a helpful AI assistant for SQL queries. Generate only SQL query and no other word to answer the following question based on the context schema:"71 relevant_schema = construct_create_statement(relevant_columns, column_descriptions)72 enhanced_query = f"system{system_prompt}userquestion: {user_query} context: {' '.join(relevant_schema)}assistant"73 sql_query = generate_response(enhanced_query, model)74 return sql_query75 76def submit_query(user_input, column_desc_input):77 column_descriptions = json.loads(column_desc_input)78 relevant_columns = find_relevant_columns(user_input, column_descriptions)79 sql_query = generate_sql_query(user_input, relevant_columns, column_descriptions)80 return sql_query81 82# Gradio interface83with gr.Blocks() as demo:84 with gr.Row():85 user_input = gr.Textbox(placeholder="Enter your SQL query prompt here", label="SQL Query Prompt")86 column_desc_input = gr.Textbox(placeholder='Enter your column descriptions JSON here', label='Column Descriptions JSON')87 submit_button = gr.Button("Generate SQL Query")88 output = gr.Textbox(label="Generated SQL Query")89 submit_button.click(submit_query, inputs=[user_input, column_desc_input], outputs=[output])90 91demo.queue().launch(debug=True, share=True, inline=False)92 