amanmurari commited on
Commit
f58d781
·
verified ·
1 Parent(s): 48e7c77

Upload folder using huggingface_hub

Browse files
Files changed (1) hide show
  1. inference.py +31 -25
inference.py CHANGED
@@ -53,7 +53,7 @@ MODEL_NAME = os.getenv("MODEL_NAME", "gpt-4o-mini")
53
  HF_TOKEN = os.getenv("HF_TOKEN")
54
  LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")
55
 
56
- SERVER_URL = os.getenv("SERVER_URL", "http://localhost:8000")
57
 
58
  SEED = 42
59
  MAX_TOKENS = 64
@@ -214,16 +214,18 @@ def get_llm_action(client: OpenAI, obs: TrafficObservation, step: int) -> Traffi
214
 
215
  def run_task(task: str, client: OpenAI) -> None:
216
  """Run a single task episode."""
217
- with TrafficControlEnv(base_url=SERVER_URL).sync() as env:
218
- # Note: openenv-core's reset takes task_id, so passing task_id=task
219
- obs = env.reset(task_id=task, seed=SEED)
220
- rewards: List[float] = []
221
- step = 0
222
- error_msg: Optional[str] = None
223
 
224
- print(f'[START] task="{task}"', flush=True)
 
 
 
 
 
 
 
 
225
 
226
- try:
227
  while not obs.done:
228
  step += 1
229
  action = get_llm_action(client, obs, step)
@@ -250,18 +252,20 @@ def run_task(task: str, client: OpenAI) -> None:
250
  flush=True,
251
  )
252
 
253
- except Exception as exc:
254
- error_msg = _sanitize(str(exc))
255
 
256
- success = not error_msg and obs.done
257
- total_reward = sum(rewards)
258
- rewards_str = ",".join(f"{r:.2f}" for r in rewards[-10:]) # last 10 for brevity
259
 
260
- print(
261
- f'[END] success={str(success).lower()} steps={step} '
262
- f'total_reward={total_reward:.2f} rewards=[{rewards_str}]',
263
- flush=True,
264
- )
 
 
 
265
 
266
 
267
  # ---------------------------------------------------------------------------
@@ -269,18 +273,20 @@ def run_task(task: str, client: OpenAI) -> None:
269
  # ---------------------------------------------------------------------------
270
 
271
  def main() -> None:
272
- client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN or "dummy-key")
 
 
 
273
 
274
  import time
275
- # Prevent Phase 2 unhandled connection exception by waiting for the server
276
- for _ in range(15):
277
  try:
278
- r = _http.get(f"{SERVER_URL.rstrip('/')}/health", timeout=2)
279
  if r.status_code == 200:
280
  break
281
  except Exception:
282
- pass
283
- time.sleep(2)
284
 
285
  for task in ["basic_flow", "emergency_priority", "dynamic_scenarios"]:
286
  run_task(task, client)
 
53
  HF_TOKEN = os.getenv("HF_TOKEN")
54
  LOCAL_IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")
55
 
56
+ SERVER_URL = os.getenv("SERVER_URL", "http://localhost:7860")
57
 
58
  SEED = 42
59
  MAX_TOKENS = 64
 
214
 
215
  def run_task(task: str, client: OpenAI) -> None:
216
  """Run a single task episode."""
217
+ print(f'[START] task="{task}"', flush=True)
 
 
 
 
 
218
 
219
+ rewards: List[float] = []
220
+ step = 0
221
+ error_msg: Optional[str] = None
222
+ success = False
223
+
224
+ try:
225
+ with TrafficControlEnv(base_url=SERVER_URL).sync() as env:
226
+ # Note: openenv-core's reset takes task_id, so passing task_id=task
227
+ obs = env.reset(task_id=task, seed=SEED)
228
 
 
229
  while not obs.done:
230
  step += 1
231
  action = get_llm_action(client, obs, step)
 
252
  flush=True,
253
  )
254
 
255
+ success = obs.done and not error_msg
 
256
 
257
+ except Exception as exc:
258
+ error_msg = _sanitize(str(exc))
259
+ success = False
260
 
261
+ total_reward = sum(rewards)
262
+ rewards_str = ",".join(f"{r:.2f}" for r in rewards[-10:]) # last 10 for brevity
263
+
264
+ print(
265
+ f'[END] success={str(success).lower()} steps={step} '
266
+ f'total_reward={total_reward:.2f} rewards=[{rewards_str}]',
267
+ flush=True,
268
+ )
269
 
270
 
271
  # ---------------------------------------------------------------------------
 
273
  # ---------------------------------------------------------------------------
274
 
275
  def main() -> None:
276
+ try:
277
+ client = OpenAI(base_url=API_BASE_URL, api_key=HF_TOKEN or "dummy-key")
278
+ except Exception:
279
+ client = None
280
 
281
  import time
282
+ # Fast reconnect logic so we don't trigger external Phase 2 timeouts
283
+ for _ in range(5):
284
  try:
285
+ r = _http.get(f"{SERVER_URL.rstrip('/')}/health", timeout=1)
286
  if r.status_code == 200:
287
  break
288
  except Exception:
289
+ time.sleep(1)
 
290
 
291
  for task in ["basic_flow", "emergency_priority", "dynamic_scenarios"]:
292
  run_task(task, client)