Team Ai
Apppublic

dvwn/nl2sql-api

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes
interactive_mode.py129 linesDownload Raw Back to scripts
1# Path: src/scripts/interactive_mode.py2# Interactive mode: Allows user to manually type questions and see the agent's response3import csv4import json5import pandas as pd6from pathlib import Path7from tabulate import tabulate8from src.database.db_manager import get_db_connection9from src.nl2sql.sql_agent import nl2sql_agent10from src.nl2sql.hf_engine import get_models11 12TEST_CASES_PATH = Path("src/scripts/test_cases.json")13 14def get_query_data(sql_query: str) -> pd.DataFrame:15    """16    Executes a SQL query and returns the results as a DataFrame.17    """18    if not sql_query or sql_query == "N/A":19        return pd.DataFrame()20    21    connection = get_db_connection()22    if not connection:23        return pd.DataFrame()24 25    try:26        df = pd.read_sql_query(sql_query, connection)27        return df28    except Exception as e:29        print(f"Error executing SQL query: {e}")30        return pd.DataFrame()31    finally:32        connection.close()33 34def verify_data(df_gold: pd.DataFrame, df_generated: pd.DataFrame) -> bool:35    """36    Guardrail check:37    Verifies if the generated DataFrame matches the expected gold DataFrame.38    """39    if df_gold.empty and df_generated.empty:40        return False  # Both empty: To catch this as a potential issue41    42    if len(df_gold) != len(df_generated):43        return False44    45    try:46        gold_val = df_gold.fillna("").astype(str).values.tolist()47        gen_val = df_generated.fillna("").astype(str).values.tolist()48 49        # Strip whitespace and convert to tuples (sortable)50        gold_tuples = [tuple(val.strip() for val in row) for row in gold_val]51        gen_tuples = [tuple(val.strip() for val in row) for row in gen_val]52 53        return sorted(gold_tuples) == sorted(gen_tuples)54    except Exception as e:55        print(f"Error during data verification: {e}")56        return False57 58 59def run_interactiveMode(model_id: str):60    """ 61    Automates the interactive Question Answering mode.62    Runs predefined questions through the agent and logs the textual NL response.63    """64    print("\n========= Interactive NL2SQL Mode =========")65    print(f"Running Interactive Question Answering Evaluation on Model: {model_id}")66    #print("Type 'exit' or 'q' to return to the main menu.\n")67 68    if not TEST_CASES_PATH.exists():69        print(f"Error: Could not find test cases at {TEST_CASES_PATH}")70        return71    72    with TEST_CASES_PATH.open("r", encoding="utf-8") as handle:73        test_cases = json.load(handle)74 75    questAns_results = []76 77    for case in test_cases:78        case_id = case.get("id")79        question = case.get("question")80        gold_sql = case.get("gold_sql")81        print(f"\n\nTesting ID {case_id}: {question[:40]}...")82 83        response = nl2sql_agent(user_question=question, model_id=model_id)84 85        # Extract metadata86        status = response.get('status')87        nl_answer = response.get('nl_response', 'N/A')88        sql_query = response.get('query', 'N/A')89        error_msg = response.get('error', '')90        attempts = response.get('attempts', 0)91 92        # Data cross-check93        df_gold = get_query_data(gold_sql)94        df_generated = get_query_data(sql_query)95 96        # Verify accuracy97        is_data_accurate = verify_data(df_gold, df_generated)98 99        questAns_results.append({100            "id": case_id,101            "model_id": model_id,102            "question": question,103            "status": status,104            "data_returned_correct": is_data_accurate,105            "attempts": attempts,106            "nl_response": nl_answer,107            "sql_generated": sql_query,108            "error": error_msg109        })110 111    # Save to CSV for human-readable and easy comparison112    safe_model_name = model_id.replace("/", "_").replace(":", "_").replace(" ", "_")113    output_csv = Path(f"Q&A_report_{safe_model_name}.csv")114 115    keys = questAns_results[0].keys()116    with output_csv.open("w", newline='', encoding="utf-8") as f:117        dict_writer = csv.DictWriter(f, fieldnames=keys)118        dict_writer.writeheader()119        dict_writer.writerows(questAns_results)120 121    print(f"\nInteractive evaluation completed. Results saved to: {output_csv}")122 123if __name__ == "__main__":124    from dotenv import load_dotenv, find_dotenv125    load_dotenv(find_dotenv())126 127    models_to_test = get_models()128    for model in models_to_test:129        run_interactiveMode(model)