P-Karthik-Mohan commited on
Commit
6025d0d
Β·
1 Parent(s): f95e9d1

FIX 2 files

Browse files
Files changed (2) hide show
  1. inference.py +46 -23
  2. openenv.yaml +2 -1
inference.py CHANGED
@@ -16,36 +16,28 @@ import time
16
  import requests
17
  from openai import OpenAI
18
  from dotenv import load_dotenv
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]
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
- def main():
34
- global client
35
- client = OpenAI(
36
- base_url=os.environ["API_BASE_URL"],
37
- api_key=os.environ["API_KEY"],
38
- )
39
- debug("SQL Analyst OpenEnv β€” Baseline Inference Agent")
40
-
41
 
42
- # ── Stdout log functions (mandatory format - stdout only) ─────────────────────
43
 
44
  def log_start(task, env, model):
45
  print(f"[START] task={task} env={env} model={model}", flush=True)
46
 
47
  def log_step(step, action, reward, done, error=None):
48
- action_clean = 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)
@@ -54,7 +46,7 @@ def log_end(success, steps, rewards):
54
  rewards_str = ",".join(f"{r:.2f}" for r in rewards)
55
  print(f"[END] success={str(success).lower()} steps={steps} rewards={rewards_str}", flush=True)
56
 
57
- # ── Debug log (stderr only - never pollutes stdout) ───────────────────────────
58
 
59
  def debug(msg):
60
  print(msg, file=sys.stderr, flush=True)
@@ -62,24 +54,24 @@ def debug(msg):
62
  # ── Environment helpers ───────────────────────────────────────────────────────
63
 
64
  def env_reset(task_id):
65
- r = requests.post(f"{ENV_BASE_URL}/reset", json={"task_id": task_id})
66
  r.raise_for_status()
67
  return r.json()
68
 
69
  def env_step(sql):
70
- r = requests.post(f"{ENV_BASE_URL}/step", json={"action": sql})
71
  r.raise_for_status()
72
  return r.json()
73
 
74
- def wait_for_server(retries=10, delay=2.0):
75
  debug("Waiting for environment server...")
76
  for i in range(retries):
77
  try:
78
- r = requests.get(f"{ENV_BASE_URL}/health", timeout=3)
79
  if r.status_code == 200:
80
  debug("Server is ready.")
81
  return
82
- except requests.exceptions.ConnectionError:
83
  pass
84
  debug(f" Not ready yet... ({i+1}/{retries})")
85
  time.sleep(delay)
@@ -149,7 +141,17 @@ def ask_llm(task_description, schema, hint, attempt, previous_attempts):
149
 
150
  def solve_task(task_id):
151
  task_name = f"sql-task-{task_id}"
152
- reset_resp = env_reset(task_id)
 
 
 
 
 
 
 
 
 
 
153
  obs = reset_resp["observation"]
154
  task_desc = obs["task_description"]
155
  schema = obs["schema"]
@@ -178,6 +180,7 @@ def solve_task(task_id):
178
  debug(f" SQL: {sql[:120]}{'...' if len(sql) > 120 else ''}")
179
  except Exception as e:
180
  error = str(e)
 
181
  log_step(step=attempt, action="", reward=0.0, done=False, error=error)
182
  all_rewards.append(0.0)
183
  steps_taken = attempt
@@ -190,6 +193,7 @@ def solve_task(task_id):
190
  details = step_resp["observation"].get("reward_breakdown", {})
191
  except Exception as e:
192
  error = str(e)
 
193
  log_step(step=attempt, action=sql, reward=0.0, done=False, error=error)
194
  all_rewards.append(0.0)
195
  steps_taken = attempt
@@ -233,13 +237,25 @@ def solve_task(task_id):
233
  # ── Main ──────────────────────────────────────────────────────────────────────
234
 
235
  def main():
 
 
 
 
 
 
 
 
 
 
 
 
 
236
  debug("SQL Analyst OpenEnv β€” Baseline Inference Agent")
237
  debug(f"Model : {MODEL_NAME}")
238
  debug(f"API Base : {API_BASE_URL}")
239
  debug(f"Env Server : {ENV_BASE_URL}")
240
 
241
- # Get task_id from command line argument or environment variable
242
- # Usage: python inference.py 1
243
  if len(sys.argv) > 1:
244
  task_id = int(sys.argv[1])
245
  else:
@@ -249,9 +265,16 @@ def main():
249
 
250
  result = solve_task(task_id)
251
 
 
 
 
252
  with open("results.json", "w") as f:
253
- json.dump({"results": [result], "avg_score": round(result["best_reward"], 3), "tasks_solved": 1 if result["solved"] else 0}, f, indent=2)
254
- debug(f"Results saved to results.json")
 
 
 
 
255
 
256
  if __name__ == "__main__":
257
  main()
 
16
  import requests
17
  from openai import OpenAI
18
  from dotenv import load_dotenv
 
19
  load_dotenv()
20
 
21
  # ── Configuration ─────────────────────────────────────────────────────────────
22
 
23
  ENV_BASE_URL = "https://p-karthik-mohan-sql-analyst-env.hf.space"
24
  MAX_ATTEMPTS = 5
 
25
  BENCHMARK = "sql-analyst-env"
26
 
27
  API_BASE_URL = os.environ.get("API_BASE_URL", "https://api.groq.com/openai/v1")
28
  MODEL_NAME = os.environ.get("MODEL_NAME", "llama-3.1-8b-instant")
29
  HF_TOKEN = os.environ.get("HF_TOKEN")
30
 
31
+ # client is initialized inside main() to avoid startup crashes
32
+ client = None
 
 
 
 
 
 
33
 
34
+ # ── Stdout log functions (mandatory format) ───────────────────────────────────
35
 
36
  def log_start(task, env, model):
37
  print(f"[START] task={task} env={env} model={model}", flush=True)
38
 
39
  def log_step(step, action, reward, done, error=None):
40
+ action_clean = str(action).replace("\n", " ").strip()[:120]
41
  error_val = error if error else "null"
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)
 
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
+ # ── Debug log (stderr only) ───────────────────────────────────────────────────
50
 
51
  def debug(msg):
52
  print(msg, file=sys.stderr, flush=True)
 
54
  # ── Environment helpers ───────────────────────────────────────────────────────
55
 
56
  def env_reset(task_id):
57
+ r = requests.post(f"{ENV_BASE_URL}/reset", json={"task_id": task_id}, timeout=30)
58
  r.raise_for_status()
59
  return r.json()
60
 
61
  def env_step(sql):
62
+ r = requests.post(f"{ENV_BASE_URL}/step", json={"action": sql}, timeout=30)
63
  r.raise_for_status()
64
  return r.json()
65
 
66
+ def wait_for_server(retries=10, delay=3.0):
67
  debug("Waiting for environment server...")
68
  for i in range(retries):
69
  try:
70
+ r = requests.get(f"{ENV_BASE_URL}/health", timeout=5)
71
  if r.status_code == 200:
72
  debug("Server is ready.")
73
  return
74
+ except Exception:
75
  pass
76
  debug(f" Not ready yet... ({i+1}/{retries})")
77
  time.sleep(delay)
 
141
 
142
  def solve_task(task_id):
143
  task_name = f"sql-task-{task_id}"
144
+
145
+ try:
146
+ reset_resp = env_reset(task_id)
147
+ except Exception as e:
148
+ debug(f"ERROR resetting task {task_id}: {e}")
149
+ log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)
150
+ log_step(step=1, action="", reward=0.0, done=False, error=str(e))
151
+ log_end(success=False, steps=1, rewards=[0.0])
152
+ return {"task_id": task_id, "task_name": task_name, "difficulty": "unknown",
153
+ "best_reward": 0.0, "attempts": 1, "solved": False}
154
+
155
  obs = reset_resp["observation"]
156
  task_desc = obs["task_description"]
157
  schema = obs["schema"]
 
180
  debug(f" SQL: {sql[:120]}{'...' if len(sql) > 120 else ''}")
181
  except Exception as e:
182
  error = str(e)
183
+ debug(f" LLM error: {error}")
184
  log_step(step=attempt, action="", reward=0.0, done=False, error=error)
185
  all_rewards.append(0.0)
186
  steps_taken = attempt
 
193
  details = step_resp["observation"].get("reward_breakdown", {})
194
  except Exception as e:
195
  error = str(e)
196
+ debug(f" Step error: {error}")
197
  log_step(step=attempt, action=sql, reward=0.0, done=False, error=error)
198
  all_rewards.append(0.0)
199
  steps_taken = attempt
 
237
  # ── Main ──────────────────────────────────────────────────────────────────────
238
 
239
  def main():
240
+ global client
241
+
242
+ # Initialize client inside main() to avoid import-time crashes
243
+ try:
244
+ api_key = os.environ.get("API_KEY") or os.environ.get("HF_TOKEN") or "no-key-needed"
245
+ client = OpenAI(
246
+ base_url=os.environ.get("API_BASE_URL", API_BASE_URL),
247
+ api_key=api_key,
248
+ )
249
+ except Exception as e:
250
+ debug(f"ERROR initializing OpenAI client: {e}")
251
+ sys.exit(1)
252
+
253
  debug("SQL Analyst OpenEnv β€” Baseline Inference Agent")
254
  debug(f"Model : {MODEL_NAME}")
255
  debug(f"API Base : {API_BASE_URL}")
256
  debug(f"Env Server : {ENV_BASE_URL}")
257
 
258
+ # Get task_id from command line or environment variable
 
259
  if len(sys.argv) > 1:
260
  task_id = int(sys.argv[1])
261
  else:
 
265
 
266
  result = solve_task(task_id)
267
 
268
+ debug(f"\n{'='*60}")
269
+ debug(f"RESULT: Task {result['task_id']} ({result['difficulty']}) β€” {'SOLVED' if result['solved'] else f'best={result[chr(98)+chr(101)+chr(115)+chr(116)+chr(95)+chr(114)+chr(101)+chr(119)+chr(97)+chr(114)+chr(100)]:.3f}'}")
270
+
271
  with open("results.json", "w") as f:
272
+ json.dump({
273
+ "results": [result],
274
+ "avg_score": round(result["best_reward"], 3),
275
+ "tasks_solved": 1 if result["solved"] else 0,
276
+ }, f, indent=2)
277
+ debug("Results saved to results.json")
278
 
279
  if __name__ == "__main__":
280
  main()
openenv.yaml CHANGED
@@ -160,4 +160,5 @@ inference:
160
  llm_env_vars:
161
  - API_BASE_URL
162
  - MODEL_NAME
163
- - API_KEY
 
 
160
  llm_env_vars:
161
  - API_BASE_URL
162
  - MODEL_NAME
163
+ - API_KEY
164
+ - HF_TOKEN