Team Ai
Apppublic

dgarzon/SQL-Explainer

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
sql_explainer_core.py160 linesDownload Raw Back to root
1"""Core logic shared by the CLI and Streamlit versions of SQL Explainer Agent."""2 3from __future__ import annotations4 5import os6import textwrap7from functools import lru_cache8 9import torch10from transformers import AutoModelForCausalLM, AutoTokenizer11 12 13DEFAULT_MODEL = os.getenv("LOCAL_MODEL_ID", "Qwen/Qwen2.5-0.5B-Instruct")14DEFAULT_MAX_NEW_TOKENS = int(os.getenv("MAX_NEW_TOKENS", "320"))15MAX_SQL_CHARS = int(os.getenv("MAX_SQL_CHARS", "8000"))16FREE_LOCAL_MODELS = (17    "Qwen/Qwen2.5-0.5B-Instruct",18    "HuggingFaceTB/SmolLM2-360M-Instruct",19    "Qwen/Qwen2.5-1.5B-Instruct",20)21 22 23SYSTEM_PROMPT = """You are SQL Explainer Agent.24 25Your job is to explain SQL queries clearly for analytics or business users.26 27Follow these rules:28- Use practical, plain English. Avoid academic language.29- Do not invent the meaning of tables, columns, or business logic when it is not obvious.30- If something is unclear, say it clearly.31- If the SQL is incomplete, explain only what is visible.32- Be careful with row count changes and explain realistic causes.33- Keep the structure exactly as requested.34- Use the exact section titles shown below.35- Keep statements grounded in the visible SQL.36- Do not add a subject line, greeting, or sign-off in the email section.37- Do not assume a specific SQL dialect unless it is visible in the query.38- Do not claim case sensitivity, function behavior, or NULL behavior that depends on the SQL engine unless it is clearly visible or you label it as dialect-dependent.39- When discussing join duplicates, be cautious: grouping does not automatically remove duplicate impact on counts or sums.40 41Always answer using these sections and headings:421. Overview432. Tables / Data Sources443. Filters and Logic454. Why results may be lower or higher than expected465. Risks / Things to Check476. Short Email Explanation48 49Section guidance:50- Overview: explain what the query does in plain English.51- Tables / Data Sources: list the main tables, views, or CTEs that are visible in the SQL.52- Filters and Logic: explain joins, WHERE filters, GROUP BY, HAVING, DISTINCT, CASE, and any important transformations.53- Why results may be lower or higher than expected: explicitly include one bullet for each of these topics, even if the answer is "Not visible in this SQL" or "Not applicable based on the visible SQL": INNER JOIN vs LEFT JOIN, filters in WHERE, NULL handling, NOT IN behavior, DISTINCT, aggregations, duplicate rows caused by joins.54- Risks / Things to Check: mention ambiguities, missing conditions, risky assumptions, incomplete SQL, possible duplicate amplification, date logic, and edge cases.55- Short Email Explanation: write exactly one short professional paragraph in English that someone could send to a colleague. No bullets. No subject line. No greeting. No sign-off.56"""57 58 59def build_prompt(sql: str) -> str:60    """Build the final user prompt that combines instructions and the SQL query."""61    return textwrap.dedent(62        f"""\63        Explain the following SQL query for a non-technical analytics or business audience.64 65        Format requirements:66        - Use the six section titles exactly as written in the system instructions.67        - In section 4, include exactly seven bullets, one for each required risk topic.68        - In section 6, write exactly one paragraph and nothing else.69 70        SQL:71        ```sql72        {sql}73        ```74        """75    )76 77 78def validate_sql_input(sql: str) -> str:79    """Validate the SQL payload before sending it to the local model."""80    cleaned_sql = sql.strip()81    if not cleaned_sql:82        raise ValueError("No SQL query was provided.")83 84    if len(cleaned_sql) > MAX_SQL_CHARS:85        raise ValueError(86            f"SQL is too long for this demo ({len(cleaned_sql)} characters). "87            f"Keep it under {MAX_SQL_CHARS} characters."88        )89 90    return cleaned_sql91 92 93def get_device_label() -> str:94    """Return a short label for the active inference device."""95    return "GPU" if torch.cuda.is_available() else "CPU"96 97 98def get_free_local_models() -> tuple[str, ...]:99    """Return the curated list of safe local models for the free Space."""100    if DEFAULT_MODEL in FREE_LOCAL_MODELS:101        return FREE_LOCAL_MODELS102    return (DEFAULT_MODEL, *FREE_LOCAL_MODELS)103 104 105@lru_cache(maxsize=2)106def load_model_components(model_id: str) -> tuple[AutoTokenizer, AutoModelForCausalLM]:107    """Load and cache the tokenizer and model once per process."""108    tokenizer = AutoTokenizer.from_pretrained(model_id, use_fast=True)109    if tokenizer.pad_token_id is None:110        tokenizer.pad_token = tokenizer.eos_token111 112    model = AutoModelForCausalLM.from_pretrained(113        model_id,114        torch_dtype=torch.float32,115        low_cpu_mem_usage=True,116    )117    model.eval()118 119    cpu_threads = max(1, min(4, os.cpu_count() or 1))120    torch.set_num_threads(cpu_threads)121 122    return tokenizer, model123 124 125def call_local_model(126    sql: str,127    model_id: str,128    max_new_tokens: int = DEFAULT_MAX_NEW_TOKENS,129) -> str:130    """Generate the SQL explanation locally with Transformers."""131    sql = validate_sql_input(sql)132    tokenizer, model = load_model_components(model_id)133 134    messages = [135        {"role": "system", "content": SYSTEM_PROMPT},136        {"role": "user", "content": build_prompt(sql)},137    ]138    prompt = tokenizer.apply_chat_template(139        messages,140        tokenize=False,141        add_generation_prompt=True,142    )143    inputs = tokenizer(prompt, return_tensors="pt")144 145    with torch.inference_mode():146        output_ids = model.generate(147            **inputs,148            max_new_tokens=max_new_tokens,149            do_sample=False,150            pad_token_id=tokenizer.pad_token_id,151            eos_token_id=tokenizer.eos_token_id,152        )153 154    generated_ids = output_ids[0][inputs["input_ids"].shape[1] :]155    result = tokenizer.decode(generated_ids, skip_special_tokens=True).strip()156    if not result:157        raise RuntimeError("The local model returned an empty response.")158 159    return result160