P-Karthik-Mohan commited on
Commit
e106b6e
Β·
1 Parent(s): 4e17f50

Fixed HF_TOKEN has no default value as required :)

Browse files
Files changed (1) hide show
  1. inference.py +41 -35
inference.py CHANGED
@@ -1,10 +1,12 @@
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> score=<0.000> rewards=<r1,r2,...>
 
 
8
  """
9
 
10
  import os
@@ -17,6 +19,8 @@ from dotenv import load_dotenv
17
 
18
  load_dotenv()
19
 
 
 
20
  ENV_BASE_URL = "https://p-karthik-mohan-sql-analyst-env.hf.space"
21
  MAX_ATTEMPTS = 5
22
  TASK_IDS = [1, 2, 3, 4, 5, 6, 7, 8]
@@ -24,14 +28,14 @@ BENCHMARK = "sql-analyst-env"
24
 
25
  API_BASE_URL = os.environ.get("API_BASE_URL", "https://api.groq.com/openai/v1")
26
  MODEL_NAME = os.environ.get("MODEL_NAME", "llama-3.1-8b-instant")
27
- HF_TOKEN = os.environ.get("HF_TOKEN", "")
28
 
29
  client = OpenAI(
30
  base_url=API_BASE_URL,
31
  api_key=HF_TOKEN if HF_TOKEN else "no-key-needed",
32
  )
33
 
34
- # ── Mandatory stdout log functions ────────────────────────────────────────────
35
 
36
  def log_start(task, env, model):
37
  print(f"[START] task={task} env={env} model={model}", flush=True)
@@ -42,9 +46,14 @@ def log_step(step, action, reward, done, error=None):
42
  done_val = str(done).lower()
43
  print(f"[STEP] step={step} action={action_clean} reward={reward:.2f} done={done_val} error={error_val}", flush=True)
44
 
45
- def log_end(success, steps, score, rewards):
46
  rewards_str = ",".join(f"{r:.2f}" for r in rewards)
47
- print(f"[END] success={str(success).lower()} steps={steps} rewards={rewards_str}", flush=True)
 
 
 
 
 
48
 
49
  # ── Environment helpers ───────────────────────────────────────────────────────
50
 
@@ -59,18 +68,18 @@ def env_step(sql):
59
  return r.json()
60
 
61
  def wait_for_server(retries=10, delay=2.0):
62
- print("Waiting for environment server...", flush=True)
63
  for i in range(retries):
64
  try:
65
  r = requests.get(f"{ENV_BASE_URL}/health", timeout=3)
66
  if r.status_code == 200:
67
- print("Server is ready.\n", flush=True)
68
  return
69
  except requests.exceptions.ConnectionError:
70
  pass
71
- print(f" Not ready yet... ({i+1}/{retries})", flush=True)
72
  time.sleep(delay)
73
- print("ERROR: Server did not start in time.", flush=True)
74
  sys.exit(1)
75
 
76
  # ── LLM ───────────────────────────────────────────────────────────────────────
@@ -143,27 +152,26 @@ def solve_task(task_id):
143
  hint = obs["hint"]
144
  difficulty = obs["difficulty"]
145
 
146
- print(f"\n{'='*60}", flush=True)
147
- print(f"TASK {task_id} ({difficulty.upper()})", flush=True)
148
- print(f"Task: {task_desc}\n", flush=True)
149
 
150
  log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)
151
 
152
  previous_attempts = []
153
  all_rewards = []
154
  best_reward = 0.0
155
- final_sql = ""
156
  steps_taken = 0
157
  success = False
158
 
159
  for attempt in range(1, MAX_ATTEMPTS + 1):
160
- print(f" Attempt {attempt}/{MAX_ATTEMPTS} β€” asking LLM...", flush=True)
161
  error = None
162
  sql = ""
163
 
164
  try:
165
  sql = ask_llm(task_desc, schema, hint, attempt, previous_attempts)
166
- print(f" SQL: {sql[:120]}{'...' if len(sql) > 120 else ''}", flush=True)
167
  except Exception as e:
168
  error = str(e)
169
  log_step(step=attempt, action="", reward=0.0, done=False, error=error)
@@ -186,9 +194,8 @@ def solve_task(task_id):
186
  all_rewards.append(reward)
187
  steps_taken = attempt
188
  best_reward = max(best_reward, reward)
189
- final_sql = sql
190
 
191
- print(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})", flush=True)
192
 
193
  log_step(step=attempt, action=sql, reward=reward, done=done, error=error)
194
 
@@ -200,15 +207,15 @@ def solve_task(task_id):
200
  })
201
 
202
  if done:
203
- print(f" PERFECT SCORE on attempt {attempt}!", flush=True)
204
  success = True
205
  break
206
  elif reward >= 0.8:
207
- print(f" Score is close ({reward:.3f}). Trying to improve...", flush=True)
208
  else:
209
- print(f" Score is low ({reward:.3f}). Refining query...", flush=True)
210
 
211
- log_end(success=success, steps=steps_taken, score=best_reward, rewards=all_rewards)
212
 
213
  return {
214
  "task_id": task_id,
@@ -216,17 +223,16 @@ def solve_task(task_id):
216
  "difficulty": difficulty,
217
  "best_reward": best_reward,
218
  "attempts": steps_taken,
219
- "final_sql": final_sql,
220
  "solved": success,
221
  }
222
 
223
  # ── Main ──────────────────────────────────────────────────────────────────────
224
 
225
  def main():
226
- print("SQL Analyst OpenEnv β€” Baseline Inference Agent", flush=True)
227
- print(f"Model : {MODEL_NAME}", flush=True)
228
- print(f"API Base : {API_BASE_URL}", flush=True)
229
- print(f"Env Server : {ENV_BASE_URL}", flush=True)
230
 
231
  wait_for_server()
232
 
@@ -238,23 +244,23 @@ def main():
238
  results.append(result)
239
  total_score += result["best_reward"]
240
 
241
- print(f"\n{'='*60}", flush=True)
242
- print("FINAL RESULTS", flush=True)
243
- print('='*60, flush=True)
244
 
245
  for r in results:
246
  status = "SOLVED" if r["solved"] else f"best={r['best_reward']:.3f}"
247
- print(f" Task {r['task_id']} ({r['difficulty']:6s}) {status} in {r['attempts']} attempt(s)", flush=True)
248
 
249
  avg_score = total_score / len(results)
250
  tasks_solved = sum(1 for r in results if r["solved"])
251
 
252
- print(f"\n Average reward : {avg_score:.3f} / 1.000", flush=True)
253
- print(f" Tasks solved : {tasks_solved} / {len(results)}", flush=True)
254
 
255
  with open("results.json", "w") as f:
256
  json.dump({"results": results, "avg_score": round(avg_score, 3), "tasks_solved": tasks_solved}, f, indent=2)
257
- print(f"\n Results saved to results.json", flush=True)
258
 
259
  if __name__ == "__main__":
260
  main()
 
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
 
19
 
20
  load_dotenv()
21
 
22
+ # ── Configuration ─────────────────────────────────────────────────────────────
23
+
24
  ENV_BASE_URL = "https://p-karthik-mohan-sql-analyst-env.hf.space"
25
  MAX_ATTEMPTS = 5
26
  TASK_IDS = [1, 2, 3, 4, 5, 6, 7, 8]
 
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
  client = OpenAI(
34
  base_url=API_BASE_URL,
35
  api_key=HF_TOKEN if HF_TOKEN else "no-key-needed",
36
  )
37
 
38
+ # ── Stdout log functions (mandatory format - stdout only) ─────────────────────
39
 
40
  def log_start(task, env, model):
41
  print(f"[START] task={task} env={env} model={model}", flush=True)
 
46
  done_val = str(done).lower()
47
  print(f"[STEP] step={step} action={action_clean} reward={reward:.2f} done={done_val} error={error_val}", flush=True)
48
 
49
+ def log_end(success, steps, rewards):
50
  rewards_str = ",".join(f"{r:.2f}" for r in rewards)
51
+ print(f"[END] success={str(success).lower()} steps={steps} rewards={rewards_str}", flush=True)
52
+
53
+ # ── Debug log (stderr only - never pollutes stdout) ───────────────────────────
54
+
55
+ def debug(msg):
56
+ print(msg, file=sys.stderr, flush=True)
57
 
58
  # ── Environment helpers ───────────────────────────────────────────────────────
59
 
 
68
  return r.json()
69
 
70
  def wait_for_server(retries=10, delay=2.0):
71
+ debug("Waiting for environment server...")
72
  for i in range(retries):
73
  try:
74
  r = requests.get(f"{ENV_BASE_URL}/health", timeout=3)
75
  if r.status_code == 200:
76
+ debug("Server is ready.")
77
  return
78
  except requests.exceptions.ConnectionError:
79
  pass
80
+ debug(f" Not ready yet... ({i+1}/{retries})")
81
  time.sleep(delay)
82
+ debug("ERROR: Server did not start in time.")
83
  sys.exit(1)
84
 
85
  # ── LLM ───────────────────────────────────────────────────────────────────────
 
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
  log_step(step=attempt, action="", reward=0.0, done=False, error=error)
 
194
  all_rewards.append(reward)
195
  steps_taken = attempt
196
  best_reward = max(best_reward, reward)
 
197
 
198
+ 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})")
199
 
200
  log_step(step=attempt, action=sql, reward=reward, done=done, error=error)
201
 
 
207
  })
208
 
209
  if done:
210
+ debug(f" PERFECT SCORE on attempt {attempt}!")
211
  success = True
212
  break
213
  elif reward >= 0.8:
214
+ debug(f" Score is close ({reward:.3f}). Trying to improve...")
215
  else:
216
+ debug(f" Score is low ({reward:.3f}). Refining query...")
217
 
218
+ log_end(success=success, steps=steps_taken, rewards=all_rewards)
219
 
220
  return {
221
  "task_id": task_id,
 
223
  "difficulty": difficulty,
224
  "best_reward": best_reward,
225
  "attempts": steps_taken,
 
226
  "solved": success,
227
  }
228
 
229
  # ── Main ──────────────────────────────────────────────────────────────────────
230
 
231
  def main():
232
+ debug("SQL Analyst OpenEnv β€” Baseline Inference Agent")
233
+ debug(f"Model : {MODEL_NAME}")
234
+ debug(f"API Base : {API_BASE_URL}")
235
+ debug(f"Env Server : {ENV_BASE_URL}")
236
 
237
  wait_for_server()
238
 
 
244
  results.append(result)
245
  total_score += result["best_reward"]
246
 
247
+ debug(f"\n{'='*60}")
248
+ debug("FINAL RESULTS")
249
+ debug('='*60)
250
 
251
  for r in results:
252
  status = "SOLVED" if r["solved"] else f"best={r['best_reward']:.3f}"
253
+ debug(f" Task {r['task_id']} ({r['difficulty']:6s}) {status} in {r['attempts']} attempt(s)")
254
 
255
  avg_score = total_score / len(results)
256
  tasks_solved = sum(1 for r in results if r["solved"])
257
 
258
+ debug(f"\n Average reward : {avg_score:.3f} / 1.000")
259
+ debug(f" Tasks solved : {tasks_solved} / {len(results)}")
260
 
261
  with open("results.json", "w") as f:
262
  json.dump({"results": results, "avg_score": round(avg_score, 3), "tasks_solved": tasks_solved}, f, indent=2)
263
+ debug(f" Results saved to results.json")
264
 
265
  if __name__ == "__main__":
266
  main()