sravaniamere commited on
Commit
574f8e7
Β·
1 Parent(s): e5d6640

error handling in inference.py

Browse files
Files changed (1) hide show
  1. inference.py +48 -23
inference.py CHANGED
@@ -102,7 +102,6 @@ def get_model_action(client: OpenAI, obs: dict, history: List[str]) -> str:
102
  # ── Main episode loop ─────────────────────────────────────────
103
  async def run_task(task_name: str) -> None:
104
  client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)
105
- http = httpx.AsyncClient(base_url=ENV_URL, timeout=30.0)
106
 
107
  rewards: List[float] = []
108
  history: List[str] = []
@@ -112,26 +111,46 @@ async def run_task(task_name: str) -> None:
112
 
113
  log_start(task_name, BENCHMARK, MODEL_NAME)
114
 
 
115
  try:
116
- reset_resp = await http.post("/reset", json={"task_name": task_name})
117
- reset_resp.raise_for_status()
118
- obs = reset_resp.json()
119
 
120
- for step in range(1, MAX_STEPS + 1):
121
- action_str = get_model_action(client, obs, history)
 
 
 
 
 
122
 
123
- step_resp = await http.post("/step", json={"corrected_query": action_str})
124
- step_resp.raise_for_status()
125
- result = step_resp.json()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
126
 
127
- obs = result["observation"]
128
- reward = float(result["reward"])
129
- done = bool(result["done"])
130
- error = result.get("info", {}).get("error")
131
 
132
  rewards.append(reward)
133
  steps_taken = step
134
- history.append(f"Step {step}: attempt={action_str!r} reward={reward:+.2f}")
 
 
135
 
136
  log_step(step, action_str, reward, done, error)
137
 
@@ -145,17 +164,23 @@ async def run_task(task_name: str) -> None:
145
  print(f"[DEBUG] Episode error: {exc}", flush=True)
146
 
147
  finally:
148
- try:
149
- await http.aclose()
150
- except Exception as e:
151
- print(f"[DEBUG] HTTP client close error: {e}", flush=True)
 
152
  log_end(success, steps_taken, score, rewards)
153
 
154
-
155
  async def main() -> None:
156
- for difficulty in ("easy", "medium", "hard"):
157
- await run_task(difficulty)
158
- print("", flush=True)
 
 
 
159
 
160
  if __name__ == "__main__":
161
- asyncio.run(main())
 
 
 
 
102
  # ── Main episode loop ─────────────────────────────────────────
103
  async def run_task(task_name: str) -> None:
104
  client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)
 
105
 
106
  rewards: List[float] = []
107
  history: List[str] = []
 
111
 
112
  log_start(task_name, BENCHMARK, MODEL_NAME)
113
 
114
+ http = None
115
  try:
116
+ http = httpx.AsyncClient(base_url=ENV_URL, timeout=60.0)
 
 
117
 
118
+ try:
119
+ reset_resp = await http.post("/reset", json={"task_name": task_name})
120
+ reset_resp.raise_for_status()
121
+ obs = reset_resp.json()
122
+ except Exception as e:
123
+ print(f"[DEBUG] Reset failed: {e}", flush=True)
124
+ return
125
 
126
+ for step in range(1, MAX_STEPS + 1):
127
+ try:
128
+ action_str = get_model_action(client, obs, history)
129
+ except Exception as e:
130
+ print(f"[DEBUG] Model failed: {e}", flush=True)
131
+ action_str = "SELECT 1"
132
+
133
+ try:
134
+ step_resp = await http.post(
135
+ "/step",
136
+ json={"corrected_query": action_str},
137
+ )
138
+ step_resp.raise_for_status()
139
+ result = step_resp.json()
140
+ except Exception as e:
141
+ print(f"[DEBUG] Step failed: {e}", flush=True)
142
+ break
143
 
144
+ obs = result["observation"]
145
+ reward = float(result["reward"])
146
+ done = bool(result["done"])
147
+ error = result.get("info", {}).get("error")
148
 
149
  rewards.append(reward)
150
  steps_taken = step
151
+ history.append(
152
+ f"Step {step}: attempt={action_str!r} reward={reward:+.2f}"
153
+ )
154
 
155
  log_step(step, action_str, reward, done, error)
156
 
 
164
  print(f"[DEBUG] Episode error: {exc}", flush=True)
165
 
166
  finally:
167
+ if http is not None:
168
+ try:
169
+ await http.aclose()
170
+ except Exception as e:
171
+ print(f"[DEBUG] HTTP close error: {e}", flush=True)
172
  log_end(success, steps_taken, score, rewards)
173
 
 
174
  async def main() -> None:
175
+ try:
176
+ for difficulty in ("easy", "medium", "hard"):
177
+ await run_task(difficulty)
178
+ print("", flush=True)
179
+ except Exception as e:
180
+ print(f"[DEBUG] Main error: {e}", flush=True)
181
 
182
  if __name__ == "__main__":
183
+ try:
184
+ asyncio.run(main())
185
+ except Exception as e:
186
+ print(f"[DEBUG] Fatal error: {e}", flush=True)