Codexzzz/sql-env
0
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())