CoolFace
Apppublic

karthiknitt/text_to_sql_with_context

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