Team Ai
Apppublic

Codexzzz/sql-env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py312 linesDownload Raw Back to root
1# """2# Inference Script — SQL Query Grader Environment3# Mandatory stdout format: [START], [STEP], [END]4# """5# import asyncio6# import os7# from typing import List, Optional8 9# from openai import OpenAI10# from sql_env import SqlAction, SqlEnv11 12# # ── Required variables (checklist compliant) ──────────────────────────────────13# API_BASE_URL     = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")14# MODEL_NAME       = os.getenv("MODEL_NAME",   "Qwen/Qwen2.5-72B-Instruct")15# HF_TOKEN         = os.getenv("HF_TOKEN")                                   # no default16# LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME") or os.getenv("IMAGE_NAME")17 18# # ── Config ────────────────────────────────────────────────────────────────────19# TASKS = ["select_basics", "aggregate_filter", "multi_join", "data_anomalies"]20# TASK_MAX_STEPS = {21#     "select_basics":    5,22#     "aggregate_filter": 5,23#     "multi_join":       7,24#     "data_anomalies":   7,25# }26# BENCHMARK         = "sql_env"27# SUCCESS_THRESHOLD = 0.728 29# # CRITICAL: ALL numeric values printed to stdout must be strictly in (0, 1).30# # Using 0.001 / 0.999 as bounds ensures :.3f formatting NEVER rounds to 0.000 or 1.000.31# _VAL_MIN = 0.00132# _VAL_MAX = 0.99933 34# SYSTEM_PROMPT = (35#     "You are an expert SQL writer. You will be given a database schema and a task. "36#     "Write a correct SQL query to solve the task. "37#     "Reply with ONLY the raw SQL — no markdown, no backticks, no explanation."38# )39 40 41# # ── Logging helpers ───────────────────────────────────────────────────────────42 43# def log_start(task: str, env: str, model: str) -> None:44#     print(f"[START] task={task} env={env} model={model}", flush=True)45 46 47# def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:48#     error_val = error if error else "null"49#     # :.3f so 0.999 prints as 0.999, never rounds to 1.00050#     print(51#         f"[STEP] step={step} action={action} "52#         f"reward={reward:.3f} done={str(done).lower()} error={error_val}",53#         flush=True,54#     )55 56 57# def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:58#     # :.3f for all values — 0.999 → "0.999", 0.001 → "0.001"59#     rewards_str = ",".join(f"{r:.3f}" for r in rewards)60#     print(61#         f"[END] success={str(success).lower()} steps={steps} "62#         f"score={score:.3f} rewards={rewards_str}",63#         flush=True,64#     )65 66 67# def _clamp(value: float) -> float:68#     """Clamp any reward/score to strictly (0, 1) using bounds safe for :.3f formatting."""69#     return min(max(float(value), _VAL_MIN), _VAL_MAX)70 71 72# # ── Task runner ───────────────────────────────────────────────────────────────73 74# async def run_task(task_name: str) -> float:75#     client    = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)76#     env       = await SqlEnv.from_docker_image(LOCAL_IMAGE_NAME)77#     max_steps = TASK_MAX_STEPS.get(task_name, 5)78 79#     rewards:     List[float] = []80#     steps_taken: int         = 081#     score:       float       = _VAL_MIN82#     success:     bool        = False83 84#     log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)85 86#     try:87#         result = await env.reset(task=task_name)88#         obs    = result.observation89 90#         for step in range(1, max_steps + 1):91#             if result.done:92#                 break93 94#             completion = client.chat.completions.create(95#                 model=MODEL_NAME,96#                 messages=[97#                     {"role": "system", "content": SYSTEM_PROMPT},98#                     {99#                         "role": "user",100#                         "content": (101#                             f"Schema:\n{obs.schema_info}\n\n"102#                             f"Task:\n{obs.task_description}\n\n"103#                             f"Previous feedback:\n{obs.feedback}\n\n"104#                             "Write the SQL query:"105#                         ),106#                     },107#                 ],108#                 max_tokens=400,109#                 temperature=0.3,110#             )111 112#             sql        = (completion.choices[0].message.content or "").strip()113#             result     = await env.step(SqlAction(sql_query=sql))114#             obs        = result.observation115#             raw_reward = result.reward if result.reward is not None else 0.0116#             reward     = _clamp(raw_reward)   # clamp BEFORE logging117#             error      = obs.error_message if obs.error_message else None118 119#             rewards.append(reward)120#             steps_taken = step121 122#             log_step(123#                 step   = step,124#                 action = sql[:100].replace("\n", " "),125#                 reward = reward,126#                 done   = result.done,127#                 error  = error,128#             )129 130#             if result.done:131#                 break132 133#         # Score = best clamped reward across all steps134#         score   = max(rewards) if rewards else _VAL_MIN135#         # score is already clamped because rewards list contains only clamped values136#         success = score >= SUCCESS_THRESHOLD137 138#     finally:139#         try:140#             await env.close()141#         except Exception:142#             pass143#         log_end(success=success, steps=steps_taken, score=score, rewards=rewards)144 145#     return score146 147 148# async def main() -> None:149#     for task in TASKS:150#         await run_task(task)151 152 153# if __name__ == "__main__":154#     asyncio.run(main())155 156 157 158"""159Inference Script — SQL Query Grader Environment160Mandatory stdout format: [START], [STEP], [END]161"""162import asyncio163import os164from typing import List, Optional165 166from openai import OpenAI167from sql_env import SqlAction, SqlEnv168 169# ── Required variables (checklist compliant) ──────────────────────────────────170API_BASE_URL     = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")171MODEL_NAME       = os.getenv("MODEL_NAME",   "Qwen/Qwen2.5-72B-Instruct")172HF_TOKEN         = os.getenv("HF_TOKEN")                                   # no default173LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME") or os.getenv("IMAGE_NAME")174 175# ── Config ────────────────────────────────────────────────────────────────────176TASKS = ["select_basics", "aggregate_filter", "multi_join", "data_anomalies", "window_functions"]177TASK_MAX_STEPS = {178    "select_basics":    5,179    "aggregate_filter": 5,180    "multi_join":       7,181    "data_anomalies":   7,182    "window_functions": 8,183}184BENCHMARK         = "sql_env"185SUCCESS_THRESHOLD = 0.7186 187# CRITICAL: ALL numeric values printed to stdout must be strictly in (0, 1).188# Using 0.001 / 0.999 as bounds ensures :.3f formatting NEVER rounds to 0.000 or 1.000.189_VAL_MIN = 0.001190_VAL_MAX = 0.999191 192SYSTEM_PROMPT = (193    "You are an expert SQL writer. You will be given a database schema and a task. "194    "Write a correct SQL query to solve the task. "195    "Reply with ONLY the raw SQL — no markdown, no backticks, no explanation."196)197 198 199# ── Logging helpers ───────────────────────────────────────────────────────────200 201def log_start(task: str, env: str, model: str) -> None:202    print(f"[START] task={task} env={env} model={model}", flush=True)203 204 205def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:206    error_val = error if error else "null"207    # :.3f so 0.999 prints as 0.999, never rounds to 1.000208    print(209        f"[STEP] step={step} action={action} "210        f"reward={reward:.3f} done={str(done).lower()} error={error_val}",211        flush=True,212    )213 214 215def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:216    # :.3f for all values — 0.999 → "0.999", 0.001 → "0.001"217    rewards_str = ",".join(f"{r:.3f}" for r in rewards)218    print(219        f"[END] success={str(success).lower()} steps={steps} "220        f"score={score:.3f} rewards={rewards_str}",221        flush=True,222    )223 224 225def _clamp(value: float) -> float:226    """Clamp any reward/score to strictly (0, 1) using bounds safe for :.3f formatting."""227    return min(max(float(value), _VAL_MIN), _VAL_MAX)228 229 230# ── Task runner ───────────────────────────────────────────────────────────────231 232async def run_task(task_name: str) -> float:233    client    = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN)234    env       = await SqlEnv.from_docker_image(LOCAL_IMAGE_NAME)235    max_steps = TASK_MAX_STEPS.get(task_name, 5)236 237    rewards:     List[float] = []238    steps_taken: int         = 0239    score:       float       = _VAL_MIN240    success:     bool        = False241 242    log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)243 244    try:245        result = await env.reset(task=task_name)246        obs    = result.observation247 248        for step in range(1, max_steps + 1):249            if result.done:250                break251 252            completion = client.chat.completions.create(253                model=MODEL_NAME,254                messages=[255                    {"role": "system", "content": SYSTEM_PROMPT},256                    {257                        "role": "user",258                        "content": (259                            f"Schema:\n{obs.schema_info}\n\n"260                            f"Task:\n{obs.task_description}\n\n"261                            f"Previous feedback:\n{obs.feedback}\n\n"262                            "Write the SQL query:"263                        ),264                    },265                ],266                max_tokens=400,267                temperature=0.3,268            )269 270            sql        = (completion.choices[0].message.content or "").strip()271            result     = await env.step(SqlAction(sql_query=sql))272            obs        = result.observation273            raw_reward = result.reward if result.reward is not None else 0.0274            reward     = _clamp(raw_reward)   # clamp BEFORE logging275            error      = obs.error_message if obs.error_message else None276 277            rewards.append(reward)278            steps_taken = step279 280            log_step(281                step   = step,282                action = sql[:100].replace("\n", " "),283                reward = reward,284                done   = result.done,285                error  = error,286            )287 288            if result.done:289                break290 291        # score = best clamped reward — already strictly in (0, 1) since rewards list292        # contains only clamped values293        score   = max(rewards) if rewards else _VAL_MIN294        success = score >= SUCCESS_THRESHOLD295 296    finally:297        try:298            await env.close()299        except Exception:300            pass301        log_end(success=success, steps=steps_taken, score=score, rewards=rewards)302 303    return score304 305 306async def main() -> None:307    for task in TASKS:308        await run_task(task)309 310 311if __name__ == "__main__":312    asyncio.run(main())