Team Ai
Apppublic

dkAmulet/sql-query-optimizer

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
inference.py278 linesDownload Raw Back to root
1#!/usr/bin/env python32"""3Baseline inference script for SQL Query Optimizer Environment.4 5Required environment variables:6  API_BASE_URL  - Base URL for the LLM API7  MODEL_NAME    - Model identifier8  HF_TOKEN      - API key (no default)9 10Usage:11  set API_BASE_URL=https://api.openai.com/v112  set MODEL_NAME=gpt-4o-mini13  set HF_TOKEN=sk-...14  python inference.py15"""16from __future__ import annotations17 18import json19import os20import sys21import time22import traceback23from typing import Dict, Optional24 25# ── Environment variables (required format) ───────────────────────────────────26API_BASE_URL = os.getenv("API_BASE_URL", "https://api.openai.com/v1")27MODEL_NAME   = os.getenv("MODEL_NAME",   "gpt-4o-mini")28HF_TOKEN     = os.getenv("HF_TOKEN")29 30TEMPERATURE  = 0.131MAX_TOKENS   = 102432 33SYSTEM_PROMPT = """\34You are an expert database engineer specialising in SQL query optimisation.35Rewrite slow SQL queries to be more efficient while returning the EXACT same result set.36 37Common optimisations:38- Replace SELECT * with only the required columns.39- Convert IN (SELECT ...) subqueries to explicit JOINs.40- Replace correlated subqueries in SELECT list with JOIN + GROUP BY.41- Exploit available indexes listed in the schema.42 43Return ONLY the optimised SQL, no markdown fences, no explanation.44"""45 46# ── Imports with error handling ───────────────────────────────────────────────47try:48    from openai import OpenAI49except ImportError as e:50    print(f"[ERROR] Failed to import openai: {e}")51    print("[ERROR] Run: pip install openai>=2.7.2")52    sys.exit(1)53 54try:55    from env import SQLQueryOptimizerEnv56    from models import SQLAction57    from tasks import TASK_ORDER, TASKS58except ImportError as e:59    print(f"[ERROR] Failed to import environment modules: {e}")60    print("[ERROR] Make sure pydantic and other deps are installed.")61    sys.exit(1)62 63 64# ── LLM helpers ───────────────────────────────────────────────────────────────65 66def build_user_message(obs_dict: dict, feedback: str) -> str:67    lines = [68        f"Task ({obs_dict['difficulty']}): {obs_dict['description']}",69        "",70        "=== Database Schema ===",71        obs_dict["schema_ddl"],72        "",73        "=== Slow Query to Optimise ===",74        obs_dict["slow_query"],75    ]76    if feedback:77        lines += [78            "",79            "=== Your Previous Attempt ===",80            obs_dict.get("current_query", ""),81            "",82            "=== Grader Feedback ===",83            feedback,84            "",85            "Fix the issues above and return an improved query.",86        ]87    else:88        lines += ["", "Write an optimised version of the slow query."]89    return "\n".join(lines)90 91 92def call_llm(client: OpenAI, user_message: str) -> str:93    try:94        response = client.chat.completions.create(95            model=MODEL_NAME,96            messages=[97                {"role": "system", "content": SYSTEM_PROMPT},98                {"role": "user",   "content": user_message},99            ],100            temperature=TEMPERATURE,101            max_tokens=MAX_TOKENS,102            stream=False,103        )104        text = (response.choices[0].message.content or "").strip()105        # Strip markdown fences if present106        for fence in ("```sql", "```SQL", "```"):107            if text.startswith(fence):108                text = text[len(fence):]109        if text.endswith("```"):110            text = text[:-3]111        return text.strip()112    except Exception as e:113        print(f"[WARN] LLM call failed: {e}")114        return ""115 116 117# ── Fallback query per task (used if LLM fails) ───────────────────────────────118FALLBACK_QUERIES = {119    "select_star_removal": (120        "SELECT user_id, username, email FROM users WHERE is_active = 1"121    ),122    "subquery_to_join": (123        "SELECT o.order_id, o.user_id, o.total_amount "124        "FROM orders o JOIN users u ON u.user_id = o.user_id "125        "WHERE u.country = 'USA' AND u.is_active = 1 AND o.status = 'delivered'"126    ),127    "aggregation_optimization": (128        "SELECT c.name AS category_name, "129        "SUM(oi.quantity * oi.unit_price) AS total_revenue "130        "FROM categories c "131        "JOIN products p ON p.category_id = c.category_id "132        "JOIN order_items oi ON oi.product_id = p.product_id "133        "GROUP BY c.category_id, c.name "134        "HAVING SUM(oi.quantity * oi.unit_price) > 1000 "135        "ORDER BY total_revenue DESC"136    ),137}138 139 140# ── Per-task runner ───────────────────────────────────────────────────────────141 142def run_task(env: SQLQueryOptimizerEnv, client: Optional[OpenAI],143             task_id: str) -> float:144    task_meta = TASKS[task_id]145 146    print(f"[START] task_id={task_id} difficulty={task_meta['difficulty']} "147          f"max_steps={task_meta['max_steps']}")148 149    try:150        obs = env.reset(task_id)151    except Exception as e:152        print(f"[END] task_id={task_id} best_reward=0.0 error=RESET_FAILED")153        return 0.0154 155    obs_dict  = obs.model_dump()156    feedback  = ""157    best_reward: float = 0.0158 159    while True:160        step_num = obs.step_number + 1161 162        # Try LLM first, fall back to hardcoded optimal query163        optimised = ""164        if client is not None:165            optimised = call_llm(client, build_user_message(obs_dict, feedback))166 167        if not optimised:168            optimised = FALLBACK_QUERIES.get(task_id,169                        "SELECT 1")  # last resort170            print(f"[WARN] Using fallback query for {task_id}")171 172        try:173            result  = env.step(SQLAction(optimized_query=optimised))174            reward  = result.reward175 176            print(f"[STEP] task_id={task_id} step={step_num} "177                  f"reward={reward.value:.4f} "178                  f"validity={reward.breakdown.validity:.2f} "179                  f"correctness={reward.breakdown.correctness:.2f} "180                  f"performance={reward.breakdown.performance:.2f} "181                  f"style={reward.breakdown.style:.2f} "182                  f"done={result.done}")183 184            if reward.value > best_reward:185                best_reward = reward.value186 187            obs      = result.observation188            obs_dict = obs.model_dump()189            feedback = reward.feedback190 191            if result.done:192                break193 194        except Exception as e:195            print(f"[STEP] task_id={task_id} step={step_num} "196                  f"reward=0.0 error=STEP_FAILED detail={e}")197            break198 199    print(f"[END] task_id={task_id} best_reward={best_reward:.4f}")200    return best_reward201 202 203# ── Main ──────────────────────────────────────────────────────────────────────204 205def main() -> Dict[str, float]:206    print(f"[INFO] SQL Query Optimizer Baseline Inference")207    print(f"[INFO] API_BASE_URL={API_BASE_URL}")208    print(f"[INFO] MODEL_NAME={MODEL_NAME}")209    print(f"[INFO] HF_TOKEN={'set' if HF_TOKEN else 'NOT SET'}")210 211    # Build OpenAI client — handle missing token gracefully212    client: Optional[OpenAI] = None213    try:214        if HF_TOKEN:215            client = OpenAI(api_key=HF_TOKEN, base_url=API_BASE_URL)216            print("[INFO] OpenAI client initialised successfully")217        else:218            print("[WARN] HF_TOKEN not set — will use fallback queries only")219    except Exception as e:220        print(f"[WARN] Could not initialise OpenAI client: {e} — using fallbacks")221 222    # Initialise environment223    try:224        env = SQLQueryOptimizerEnv()225    except Exception as e:226        print(f"[ERROR] Failed to initialise environment: {e}")227        traceback.print_exc()228        sys.exit(1)229 230    results: Dict[str, float] = {}231    t0 = time.time()232 233    try:234        for task_id in TASK_ORDER:235            try:236                results[task_id] = run_task(env, client, task_id)237            except Exception as e:238                print(f"[ERROR] Task {task_id} failed unexpectedly: {e}")239                traceback.print_exc()240                results[task_id] = 0.0241    finally:242        try:243            env.close()244        except Exception:245            pass246 247    elapsed = time.time() - t0248    overall = sum(results.values()) / max(len(results), 1)249 250    print(f"\n[SUMMARY] overall_average={overall:.4f} elapsed_seconds={elapsed:.1f}")251    for task_id, score in results.items():252        print(f"[SUMMARY] task={task_id} score={score:.4f}")253 254    try:255        with open("baseline_results.json", "w") as fh:256            json.dump({257                "task_scores":     results,258                "overall_average": overall,259                "model":           MODEL_NAME,260                "elapsed_seconds": round(elapsed, 1),261            }, fh, indent=2)262        print("[INFO] Results saved to baseline_results.json")263    except Exception as e:264        print(f"[WARN] Could not save results file: {e}")265 266    return results267 268 269if __name__ == "__main__":270    try:271        results = main()272        # Exit 0 even if all scores are 0 — don't fail the pipeline273        sys.exit(0)274    except Exception as e:275        print(f"[ERROR] Unhandled exception: {e}")276        traceback.print_exc()277        sys.exit(1)278