P-Karthik-Mohan commited on
Commit
aad228d
Β·
1 Parent(s): c4c20c0

Match passing submission pattern - score clamped, top level client

Browse files
Files changed (1) hide show
  1. inference.py +111 -135
inference.py CHANGED
@@ -1,12 +1,15 @@
1
  """
2
  inference.py β€” Baseline AI agent for SQL Analyst OpenEnv
3
 
4
- STDOUT FORMAT (mandatory):
5
- [START] task=<task_name> env=<benchmark> model=<model_name>
6
- [STEP] step=<n> action=<sql> reward=<0.00> done=<true|false> error=<msg|null>
7
- [END] success=<true|false> steps=<n> rewards=<r1,r2,...,rn>
8
-
9
- All debug output goes to stderr. Stdout has only [START], [STEP], [END] lines.
 
 
 
10
  """
11
 
12
  import os
@@ -14,50 +17,59 @@ import sys
14
  import json
15
  import time
16
  import requests
 
 
17
 
18
  # ── Configuration ─────────────────────────────────────────────────────────────
19
 
20
  ENV_BASE_URL = "https://p-karthik-mohan-sql-analyst-env.hf.space"
21
- MAX_ATTEMPTS = 1
22
  BENCHMARK = "sql-analyst-env"
23
 
24
  API_BASE_URL = os.environ.get("API_BASE_URL", "https://api.groq.com/openai/v1")
25
  MODEL_NAME = os.environ.get("MODEL_NAME", "llama-3.1-8b-instant")
26
  HF_TOKEN = os.environ.get("HF_TOKEN")
27
 
28
- client = None
 
 
 
 
 
 
 
29
 
30
  # ── Stdout log functions (mandatory format) ───────────────────────────────────
31
 
32
- def log_start(task, env, model):
33
  print(f"[START] task={task} env={env} model={model}", flush=True)
34
 
35
- def log_step(step, action, reward, done, error=None):
36
  action_clean = str(action).replace("\n", " ").strip()[:120]
37
  error_val = error if error else "null"
38
  done_val = str(done).lower()
39
  print(f"[STEP] step={step} action={action_clean} reward={reward:.2f} done={done_val} error={error_val}", flush=True)
40
 
41
- def log_end(success, steps, rewards):
42
  rewards_str = ",".join(f"{r:.2f}" for r in rewards)
43
- print(f"[END] success={str(success).lower()} steps={steps} rewards={rewards_str}", flush=True)
44
 
45
- def debug(msg):
46
  print(msg, file=sys.stderr, flush=True)
47
 
48
  # ── Environment helpers ───────────────────────────────────────────────────────
49
 
50
- def env_reset(task_id):
51
  r = requests.post(f"{ENV_BASE_URL}/reset", json={"task_id": task_id}, timeout=30)
52
  r.raise_for_status()
53
  return r.json()
54
 
55
- def env_step(sql):
56
  r = requests.post(f"{ENV_BASE_URL}/step", json={"action": sql}, timeout=30)
57
  r.raise_for_status()
58
  return r.json()
59
 
60
- def wait_for_server(retries=10, delay=3.0):
61
  debug("Waiting for environment server...")
62
  for i in range(retries):
63
  try:
@@ -74,7 +86,7 @@ def wait_for_server(retries=10, delay=3.0):
74
 
75
  # ── LLM ───────────────────────────────────────────────────────────────────────
76
 
77
- def build_system_prompt():
78
  return """You are an expert SQL analyst. Your job is to write correct SQLite queries.
79
 
80
  Rules:
@@ -114,7 +126,7 @@ Attempt number: {attempt}
114
  prompt += "\nWrite the corrected SQL query now:"
115
  return prompt
116
 
117
- def ask_llm(task_description, schema, hint, attempt, previous_attempts):
118
  messages = [
119
  {"role": "system", "content": build_system_prompt()},
120
  {"role": "user", "content": build_user_prompt(task_description, schema, hint, attempt, previous_attempts)},
@@ -131,142 +143,106 @@ def ask_llm(task_description, schema, hint, attempt, previous_attempts):
131
  sql = "\n".join(line for line in lines if not line.strip().startswith("```")).strip()
132
  return sql
133
 
134
- # ── Task solver ───────────────────────────────────────────────────────────────
135
 
136
- def solve_task(task_id):
137
- task_name = f"sql-task-{task_id}"
 
 
 
 
 
 
 
138
 
139
  try:
140
  reset_resp = env_reset(task_id)
141
- except Exception as e:
142
- debug(f"ERROR resetting task {task_id}: {e}")
143
- log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)
144
- log_step(step=1, action="", reward=0.0, done=False, error=str(e))
145
- log_end(success=False, steps=1, rewards=[0.0])
146
- return {"task_id": task_id, "task_name": task_name, "difficulty": "unknown",
147
- "best_reward": 0.0, "attempts": 1, "solved": False}
148
-
149
- obs = reset_resp["observation"]
150
- task_desc = obs["task_description"]
151
- schema = obs["schema"]
152
- hint = obs["hint"]
153
- difficulty = obs["difficulty"]
154
-
155
- debug(f"\n{'='*60}")
156
- debug(f"TASK {task_id} ({difficulty.upper()})")
157
- debug(f"Task: {task_desc}\n")
158
-
159
- log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
160
 
161
- previous_attempts = []
162
- all_rewards = []
163
- best_reward = 0.0
164
- steps_taken = 0
165
- success = False
166
 
167
- for attempt in range(1, MAX_ATTEMPTS + 1):
168
- debug(f" Attempt {attempt}/{MAX_ATTEMPTS} β€” asking LLM...")
169
- error = None
170
- sql = ""
171
 
172
- try:
173
- sql = ask_llm(task_desc, schema, hint, attempt, previous_attempts)
174
- debug(f" SQL: {sql[:120]}{'...' if len(sql) > 120 else ''}")
175
- except Exception as e:
176
- error = str(e)
177
- debug(f" LLM error: {error}")
178
- log_step(step=attempt, action="", reward=0.0, done=False, error=error)
179
- all_rewards.append(0.0)
180
- steps_taken = attempt
181
- continue
182
 
183
- try:
184
- step_resp = env_step(sql)
185
- reward = step_resp["reward"]
186
- done = step_resp["done"]
187
- details = step_resp["observation"].get("reward_breakdown", {})
188
- except Exception as e:
189
- error = str(e)
190
- debug(f" Step error: {error}")
191
- log_step(step=attempt, action=sql, reward=0.0, done=False, error=error)
192
- all_rewards.append(0.0)
193
- steps_taken = attempt
194
- continue
195
-
196
- all_rewards.append(reward)
197
- steps_taken = attempt
198
- best_reward = max(best_reward, reward)
199
-
200
- debug(f" Reward: {reward:.3f} (cols={details.get('column_score',0):.2f} rows={details.get('row_score',0):.2f} vals={details.get('value_score',0):.2f})")
201
-
202
- log_step(step=attempt, action=sql, reward=reward, done=done, error=error)
203
-
204
- previous_attempts.append({
205
- "attempt": attempt,
206
- "sql": sql,
207
- "reward": reward,
208
- "details": details,
209
- })
210
-
211
- if done:
212
- debug(f" PERFECT SCORE on attempt {attempt}!")
213
- success = True
214
- break
215
- elif reward >= 0.8:
216
- debug(f" Score is close ({reward:.3f}). Trying to improve...")
217
- else:
218
- debug(f" Score is low ({reward:.3f}). Refining query...")
219
-
220
- log_end(success=success, steps=steps_taken, rewards=all_rewards)
221
-
222
- return {
223
- "task_id": task_id,
224
- "task_name": task_name,
225
- "difficulty": difficulty,
226
- "best_reward": best_reward,
227
- "attempts": steps_taken,
228
- "solved": success,
229
- }
230
 
231
- # ── Main ──────────────────────────────────────────────────────────────────────
 
 
232
 
233
- def main():
234
- global client
 
 
 
 
235
 
236
- try:
237
- from openai import OpenAI
238
- api_key = os.environ.get("HF_TOKEN") or os.environ.get("API_KEY") or "no-key-needed"
239
- client = OpenAI(
240
- base_url=API_BASE_URL,
241
- api_key=api_key,
242
- )
243
- except Exception as e:
244
- debug(f"ERROR initializing OpenAI client: {e}")
245
- sys.exit(1)
246
 
 
247
  debug("SQL Analyst OpenEnv β€” Baseline Inference Agent")
248
  debug(f"Model : {MODEL_NAME}")
249
  debug(f"API Base : {API_BASE_URL}")
250
  debug(f"Env Server : {ENV_BASE_URL}")
251
 
252
- if len(sys.argv) > 1:
253
- task_id = int(sys.argv[1])
254
- else:
255
- task_id = int(os.environ.get("TASK_ID", "1"))
256
-
257
  wait_for_server()
258
 
259
- result = solve_task(task_id)
260
-
261
- debug(f"\nRESULT: Task {result['task_id']} ({result['difficulty']}) β€” {'SOLVED' if result['solved'] else 'best=' + str(round(result['best_reward'], 3))}")
262
-
263
- with open("results.json", "w") as f:
264
- json.dump({
265
- "results": [result],
266
- "avg_score": round(result["best_reward"], 3),
267
- "tasks_solved": 1 if result["solved"] else 0,
268
- }, f, indent=2)
269
- debug("Results saved to results.json")
270
 
271
  if __name__ == "__main__":
272
  main()
 
1
  """
2
  inference.py β€” Baseline AI agent for SQL Analyst OpenEnv
3
 
4
+ Required environment variables:
5
+ API_BASE_URL LLM API endpoint
6
+ MODEL_NAME Model identifier
7
+ HF_TOKEN HuggingFace / API key
8
+
9
+ Stdout format:
10
+ [START] task=<task> env=<benchmark> model=<model>
11
+ [STEP] step=<n> action=<action> reward=<0.00> done=<true|false> error=<msg|null>
12
+ [END] success=<true|false> steps=<n> score=<0.000> rewards=<r1,r2,...>
13
  """
14
 
15
  import os
 
17
  import json
18
  import time
19
  import requests
20
+ from typing import List, Optional
21
+ from openai import OpenAI
22
 
23
  # ── Configuration ─────────────────────────────────────────────────────────────
24
 
25
  ENV_BASE_URL = "https://p-karthik-mohan-sql-analyst-env.hf.space"
26
+ MAX_ATTEMPTS = 5
27
  BENCHMARK = "sql-analyst-env"
28
 
29
  API_BASE_URL = os.environ.get("API_BASE_URL", "https://api.groq.com/openai/v1")
30
  MODEL_NAME = os.environ.get("MODEL_NAME", "llama-3.1-8b-instant")
31
  HF_TOKEN = os.environ.get("HF_TOKEN")
32
 
33
+ TASKS = [1, 2, 3, 4, 5, 6, 7, 8]
34
+ SUCCESS_SCORE_THRESHOLD = 0.5
35
+
36
+ # Initialize client at top level like the sample
37
+ client = OpenAI(
38
+ base_url=API_BASE_URL,
39
+ api_key=HF_TOKEN if HF_TOKEN else "no-key-needed",
40
+ )
41
 
42
  # ── Stdout log functions (mandatory format) ───────────────────────────────────
43
 
44
+ def log_start(task: str, env: str, model: str) -> None:
45
  print(f"[START] task={task} env={env} model={model}", flush=True)
46
 
47
+ def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:
48
  action_clean = str(action).replace("\n", " ").strip()[:120]
49
  error_val = error if error else "null"
50
  done_val = str(done).lower()
51
  print(f"[STEP] step={step} action={action_clean} reward={reward:.2f} done={done_val} error={error_val}", flush=True)
52
 
53
+ def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:
54
  rewards_str = ",".join(f"{r:.2f}" for r in rewards)
55
+ print(f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}", flush=True)
56
 
57
+ def debug(msg: str) -> None:
58
  print(msg, file=sys.stderr, flush=True)
59
 
60
  # ── Environment helpers ───────────────────────────────────────────────────────
61
 
62
+ def env_reset(task_id: int) -> dict:
63
  r = requests.post(f"{ENV_BASE_URL}/reset", json={"task_id": task_id}, timeout=30)
64
  r.raise_for_status()
65
  return r.json()
66
 
67
+ def env_step(sql: str) -> dict:
68
  r = requests.post(f"{ENV_BASE_URL}/step", json={"action": sql}, timeout=30)
69
  r.raise_for_status()
70
  return r.json()
71
 
72
+ def wait_for_server(retries: int = 10, delay: float = 3.0) -> None:
73
  debug("Waiting for environment server...")
74
  for i in range(retries):
75
  try:
 
86
 
87
  # ── LLM ───────────────────────────────────────────────────────────────────────
88
 
89
+ def build_system_prompt() -> str:
90
  return """You are an expert SQL analyst. Your job is to write correct SQLite queries.
91
 
92
  Rules:
 
126
  prompt += "\nWrite the corrected SQL query now:"
127
  return prompt
128
 
129
+ def ask_llm(task_description, schema, hint, attempt, previous_attempts) -> str:
130
  messages = [
131
  {"role": "system", "content": build_system_prompt()},
132
  {"role": "user", "content": build_user_prompt(task_description, schema, hint, attempt, previous_attempts)},
 
143
  sql = "\n".join(line for line in lines if not line.strip().startswith("```")).strip()
144
  return sql
145
 
146
+ # ── Task runner ───────────────────────────────────────────────────────────────
147
 
148
+ def run_task(task_id: int) -> None:
149
+ task_name = f"sql-task-{task_id}"
150
+ rewards: List[float] = []
151
+ steps_taken = 0
152
+ score = 0.0
153
+ success = False
154
+ last_error: Optional[str] = None
155
+
156
+ log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)
157
 
158
  try:
159
  reset_resp = env_reset(task_id)
160
+ obs = reset_resp["observation"]
161
+ task_desc = obs["task_description"]
162
+ schema = obs["schema"]
163
+ hint = obs["hint"]
164
+ difficulty = obs["difficulty"]
165
+
166
+ debug(f"\n{'='*60}")
167
+ debug(f"TASK {task_id} ({difficulty.upper()})")
168
+ debug(f"Task: {task_desc}\n")
169
+
170
+ previous_attempts = []
171
+
172
+ for attempt in range(1, MAX_ATTEMPTS + 1):
173
+ debug(f" Attempt {attempt}/{MAX_ATTEMPTS} β€” asking LLM...")
174
+ last_error = None
175
+ sql = ""
176
+
177
+ try:
178
+ sql = ask_llm(task_desc, schema, hint, attempt, previous_attempts)
179
+ debug(f" SQL: {sql[:120]}{'...' if len(sql) > 120 else ''}")
180
+ except Exception as e:
181
+ last_error = str(e)
182
+ debug(f" LLM error: {last_error}")
183
+ log_step(step=attempt, action="", reward=0.0, done=False, error=last_error)
184
+ rewards.append(0.0)
185
+ steps_taken = attempt
186
+ continue
187
+
188
+ try:
189
+ step_resp = env_step(sql)
190
+ reward = step_resp["reward"]
191
+ done = step_resp["done"]
192
+ details = step_resp["observation"].get("reward_breakdown", {})
193
+ except Exception as e:
194
+ last_error = str(e)
195
+ debug(f" Step error: {last_error}")
196
+ log_step(step=attempt, action=sql, reward=0.0, done=False, error=last_error)
197
+ rewards.append(0.0)
198
+ steps_taken = attempt
199
+ continue
200
+
201
+ rewards.append(reward)
202
+ steps_taken = attempt
203
 
204
+ debug(f" Reward: {reward:.3f} (cols={details.get('column_score',0):.2f} rows={details.get('row_score',0):.2f} vals={details.get('value_score',0):.2f})")
 
 
 
 
205
 
206
+ log_step(step=attempt, action=sql, reward=reward, done=done, error=last_error)
 
 
 
207
 
208
+ previous_attempts.append({
209
+ "attempt": attempt,
210
+ "sql": sql,
211
+ "reward": reward,
212
+ "details": details,
213
+ })
 
 
 
 
214
 
215
+ if done:
216
+ debug(f" PERFECT SCORE on attempt {attempt}!")
217
+ break
218
+ elif reward >= 0.8:
219
+ debug(f" Score is close ({reward:.3f}). Trying to improve...")
220
+ else:
221
+ debug(f" Score is low ({reward:.3f}). Refining query...")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
222
 
223
+ except Exception as e:
224
+ last_error = str(e)
225
+ debug(f"ERROR in task {task_id}: {last_error}")
226
 
227
+ finally:
228
+ # Clamp score strictly between 0 and 1 β€” required by OpenEnv spec
229
+ score = sum(rewards) / len(rewards) if rewards else 0.0
230
+ score = max(1e-6, min(score, 1 - 1e-6))
231
+ success = score >= SUCCESS_SCORE_THRESHOLD
232
+ log_end(success=success, steps=steps_taken, score=score, rewards=rewards)
233
 
234
+ # ── Main ──────────────────────────────────────────────────────────────────────
 
 
 
 
 
 
 
 
 
235
 
236
+ def main() -> None:
237
  debug("SQL Analyst OpenEnv β€” Baseline Inference Agent")
238
  debug(f"Model : {MODEL_NAME}")
239
  debug(f"API Base : {API_BASE_URL}")
240
  debug(f"Env Server : {ENV_BASE_URL}")
241
 
 
 
 
 
 
242
  wait_for_server()
243
 
244
+ for task_id in TASKS:
245
+ run_task(task_id)
 
 
 
 
 
 
 
 
 
246
 
247
  if __name__ == "__main__":
248
  main()