Team Ai
Apppublic

Tsah00/sql-env

sourceHugging Facemitupdated 6mo agoView on Hugging Face
0likes
inference.py356 linesDownload Raw Back to root
1"""2inference.py - Baseline inference script for the SQL Query Learning Environment.3 4MANDATORY STDOUT FORMAT:5  [START] task=<task_name> env=<benchmark> model=<model_name>6  [STEP]  step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>7  [END]   success=<true|false> steps=<n> score=<score> rewards=<r1,r2,...,rn>8 9Environment variables:10  API_BASE_URL   The API endpoint for the LLM (default: HF router)11  MODEL_NAME     The model identifier (default: Qwen/Qwen2.5-72B-Instruct)12  HF_TOKEN       Hugging Face / API key13 14Usage:15  python inference.py16"""17 18from __future__ import annotations19 20import os21import sys22import textwrap23from typing import List, Optional24 25from openai import OpenAI26 27# ── Config ──────────────────────────────────────────────────────────────────28 29API_KEY = (30    os.getenv("HF_TOKEN")31    or os.getenv("OPENAI_API_KEY")32    or os.getenv("API_KEY")33)34API_BASE_URL = os.getenv("API_BASE_URL") or "https://router.huggingface.co/v1"35MODEL_NAME = os.getenv("MODEL_NAME") or "Qwen/Qwen2.5-72B-Instruct"36 37BENCHMARK = "sql_env"38MAX_STEPS = 1039TEMPERATURE = 0.140MAX_TOKENS = 51241 42# All 9 tasks across 3 difficulty tiers43ALL_TASKS = [44    ("easy_1",   "easy"),45    ("easy_2",   "easy"),46    ("easy_3",   "easy"),47    ("medium_1", "medium"),48    ("medium_2", "medium"),49    ("medium_3", "medium"),50    ("hard_1",   "hard"),51    ("hard_2",   "hard"),52    ("hard_3",   "hard"),53]54 55SYSTEM_PROMPT = textwrap.dedent("""56You are a data analyst at an e-commerce company. Business stakeholders57(marketing, finance, CRM, merchandising) submit ad-hoc data requests and58you fulfil them by writing SQL queries against the company database.59 60RULES:61- Return ONLY the SQL query — no explanation, no markdown fences, no comments.62- Use standard SQLite syntax.63- Use column aliases to match the expected column names exactly as stated in the task.64- Do NOT use SELECT * — select only the required columns.65- Aim for the simplest correct query; avoid unnecessary subqueries or CROSS JOINs.66""").strip()67 68 69# ── Structured Logging (MANDATORY FORMAT) ───────────────────────────────────70 71def log_start(task: str, env: str, model: str) -> None:72    print(f"[START] task={task} env={env} model={model}", flush=True)73 74 75def log_step(step: int, action: str, reward: float, done: bool,76             error: Optional[str]) -> None:77    # Sanitize action string: remove newlines, limit length78    action_clean = action.replace("\n", " ").replace("\r", "").strip()79    error_val = error if error else "null"80    done_val = str(done).lower()81    print(82        f"[STEP] step={step} action={action_clean} "83        f"reward={reward:.2f} done={done_val} error={error_val}",84        flush=True,85    )86 87 88def log_end(success: bool, steps: int, score: float,89            rewards: List[float]) -> None:90    rewards_str = ",".join(f"{r:.2f}" for r in rewards)91    print(92        f"[END] success={str(success).lower()} steps={steps} "93        f"score={score:.2f} rewards={rewards_str}",94        flush=True,95    )96 97 98# ── LLM Query ──────────────────────────────────────────────────────────────99 100def get_sql_from_llm(101    client: OpenAI,102    schema_info: str,103    task_description: str,104    expected_columns: List[str],105    previous_attempts: List[dict],106) -> str:107    """Ask the LLM to produce a SQL query for the given task."""108    attempts_text = ""109    if previous_attempts:110        last = previous_attempts[-1]111        attempts_text = (112            f"\nPREVIOUS ATTEMPT:\n"113            f"  Query: {last['query']}\n"114            f"  Reward: {last['reward']}\n"115            f"  Feedback: {last['message']}\n"116            f"  Error: {last.get('error', 'none')}\n"117            f"\nImprove on this attempt.\n"118        )119 120    user_prompt = (121        f"DATABASE SCHEMA:\n{schema_info}\n\n"122        f"TASK:\n{task_description}\n\n"123        f"EXPECTED OUTPUT COLUMNS: {', '.join(expected_columns)}\n"124        f"{attempts_text}\n"125        f"SQL QUERY:"126    )127 128    try:129        completion = client.chat.completions.create(130            model=MODEL_NAME,131            messages=[132                {"role": "system", "content": SYSTEM_PROMPT},133                {"role": "user", "content": user_prompt},134            ],135            temperature=TEMPERATURE,136            max_tokens=MAX_TOKENS,137            stream=False,138        )139        text = (completion.choices[0].message.content or "").strip()140        # Strip markdown fences if present141        if text.startswith("```"):142            lines = text.split("\n")143            lines = [l for l in lines if not l.startswith("```")]144            text = "\n".join(lines).strip()145        return text if text else _fallback_query(task_description)146    except Exception as exc:147        print(f"[DEBUG] LLM request failed: {exc}", flush=True)148        return _fallback_query(task_description)149 150 151def _fallback_query(description: str) -> str:152    """153    Deterministic rule-based fallback when LLM API is unavailable.154    Pattern-matches task descriptions to known reference queries.155    """156    d = description.lower()157 158    # Easy159    if "usa" in d or "united states" in d:160        return "SELECT name, email FROM customers WHERE country = 'USA'"161    if "count" in d and "completed" in d:162        return "SELECT COUNT(*) AS total_completed FROM orders WHERE status = 'completed'"163    if "top 5" in d and "expensive" in d:164        return "SELECT name, category, price FROM products ORDER BY price DESC LIMIT 5"165 166    # Hard (checked before medium to avoid substring collisions)167    if "above" in d and "average" in d:168        return (169            "WITH customer_totals AS ("170            "  SELECT c.name, SUM(o.total_amount) AS total_spent"171            "  FROM customers c"172            "  JOIN orders o ON c.id = o.customer_id"173            "  WHERE o.status = 'completed'"174            "  GROUP BY c.id, c.name"175            ") "176            "SELECT name, total_spent FROM customer_totals "177            "WHERE total_spent > (SELECT AVG(total_spent) FROM customer_totals) "178            "ORDER BY total_spent DESC"179        )180    if "best-selling" in d or "best selling" in d:181        return (182            "WITH product_sales AS ("183            "  SELECT p.category, p.name AS product_name, SUM(oi.quantity) AS total_quantity"184            "  FROM products p JOIN order_items oi ON p.id = oi.product_id"185            "  GROUP BY p.id, p.category, p.name"186            "), ranked AS ("187            "  SELECT category, product_name, total_quantity,"188            "    RANK() OVER (PARTITION BY category ORDER BY total_quantity DESC) AS rnk"189            "  FROM product_sales"190            ") "191            "SELECT category, product_name, total_quantity FROM ranked WHERE rnk = 1 "192            "ORDER BY category, product_name"193        )194    if "2022" in d and "2023" in d and "2024" in d:195        return (196            "SELECT c.name, c.email FROM customers c "197            "WHERE (SELECT COUNT(DISTINCT STRFTIME('%Y', o.order_date)) "198            "FROM orders o WHERE o.customer_id = c.id "199            "AND STRFTIME('%Y', o.order_date) IN ('2022','2023','2024')) = 3 "200            "ORDER BY c.name ASC"201        )202 203    # Medium204    if "total spending" in d or "total spent" in d:205        return (206            "SELECT c.name, SUM(o.total_amount) AS total_spent "207            "FROM customers c JOIN orders o ON c.id = o.customer_id "208            "WHERE o.status = 'completed' "209            "GROUP BY c.id, c.name ORDER BY total_spent DESC"210        )211    if "never" in d and ("order" in d or "appear" in d):212        return (213            "SELECT p.name, p.category, p.price FROM products p "214            "LEFT JOIN order_items oi ON p.id = oi.product_id WHERE oi.id IS NULL"215        )216    if "average" in d and "month" in d:217        return (218            "SELECT STRFTIME('%Y-%m', order_date) AS month, "219            "ROUND(AVG(total_amount), 2) AS avg_order_value "220            "FROM orders WHERE order_date LIKE '2023%' "221            "GROUP BY month ORDER BY month ASC"222        )223 224    return "SELECT name FROM customers LIMIT 5"225 226 227# ── Run One Task Episode ───────────────────────────────────────────────────228 229def run_task(230    task_id: str,231    difficulty: str,232    client: OpenAI,233) -> float:234    """235    Run a single task as one episode. Returns the best score in [0, 1].236 237    Emits [START], [STEP]..., [END] to stdout.238    """239    # Import env locally to keep module-level clean240    sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))241    from server.sql_environment import SQLEnvironment242    from models import SQLAction243 244    env = SQLEnvironment()245    rewards: List[float] = []246    steps_taken = 0247    best_score = 0.0248    success = False249 250    log_start(task=task_id, env=BENCHMARK, model=MODEL_NAME)251 252    try:253        obs = env.reset(difficulty=difficulty, task_id=task_id)254        previous_attempts: List[dict] = []255 256        for step in range(1, MAX_STEPS + 1):257            if obs.done:258                break259 260            # Get query from LLM or fallback261            if API_KEY and API_BASE_URL:262                query = get_sql_from_llm(263                    client,264                    schema_info=obs.schema_info,265                    task_description=obs.task_description,266                    expected_columns=obs.expected_columns,267                    previous_attempts=previous_attempts,268                )269            else:270                query = _fallback_query(obs.task_description)271 272            action = SQLAction(query=query, difficulty=difficulty, task_id=task_id)273            obs = env.step(action)274 275            reward = obs.reward276            done = obs.done277            error = obs.error if obs.error else None278 279            rewards.append(reward)280            steps_taken = step281            best_score = max(best_score, reward)282 283            log_step(284                step=step,285                action=query,286                reward=reward,287                done=done,288                error=error,289            )290 291            previous_attempts.append({292                "query": query,293                "reward": reward,294                "message": obs.message,295                "error": obs.error,296            })297            if len(previous_attempts) > 2:298                previous_attempts.pop(0)299 300            # Near-perfect score — stop early301            if reward >= 0.99:302                break303 304            if done:305                break306 307        # Final score — clamped to open interval (0.01, 0.99) per Phase 2 spec308        score = min(0.99, max(0.01, best_score))309        success = score >= 0.5310 311    except Exception as exc:312        print(f"[DEBUG] Exception during episode: {exc}", flush=True)313        score = 0.0314 315    finally:316        try:317            env.close()318        except Exception:319            pass320        log_end(321            success=success,322            steps=steps_taken,323            score=score,324            rewards=rewards,325        )326 327    return score328 329 330# ── Main ────────────────────────────────────────────────────────────────────331 332def main() -> None:333    client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY or "sk-placeholder")334 335    total_score = 0.0336    task_scores = {}337 338    for task_id, difficulty in ALL_TASKS:339        score = run_task(task_id, difficulty, client)340        task_scores[task_id] = score341        total_score += score342 343    # Summary (not part of mandatory format — informational only)344    print("\n" + "=" * 60, flush=True)345    print("INFERENCE SUMMARY", flush=True)346    print("=" * 60, flush=True)347    for task_id, score in task_scores.items():348        status = "PASS" if score >= 0.5 else "FAIL"349        print(f"  [{status}] {task_id}: score={score:.2f}", flush=True)350    avg = total_score / len(ALL_TASKS) if ALL_TASKS else 0.0351    print(f"\n  Average score: {avg:.2f} ({total_score:.2f}/{len(ALL_TASKS)})", flush=True)352 353 354if __name__ == "__main__":355    main()356