Rakshithn123 commited on
Commit
386bf01
·
1 Parent(s): b4b9b2e

final one

Browse files
Files changed (1) hide show
  1. inference.py +6 -6
inference.py CHANGED
@@ -111,7 +111,7 @@ def log_start(task: str, env: str, model: str) -> None:
111
  def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:
112
  error_val = error if error else "null"
113
  done_val = str(done).lower()
114
- action_clean = action.replace("\n", " ").replace("\r", "")[:120]
115
  print(
116
  f"[STEP] step={step} action={action_clean} reward={reward:.2f} "
117
  f"done={done_val} error={error_val}",
@@ -262,7 +262,7 @@ def get_agent_action(
262
  context_parts.append(f"Remaining questions: {obs.get('remaining_items', 0)}")
263
 
264
  if history:
265
- context_parts.append("Recent actions:\n" + "\n".join(history[-2:]))
266
 
267
  if force_synthesize:
268
  context_parts.append(
@@ -299,7 +299,7 @@ def get_agent_action(
299
  parts = text.split("```")
300
  text = parts[1] if len(parts) > 1 else parts[0]
301
  if text.startswith("json"):
302
- text = text[4:]
303
  text = text.strip()
304
  # Strip control characters that break json.loads
305
  text = re.sub(r'[\x00-\x1f\x7f]', ' ', text)
@@ -372,7 +372,7 @@ async def run_task(client: OpenAI, task_name: str) -> float:
372
 
373
  try:
374
  result = await env.reset()
375
- obs = result.observation.model_dump() if hasattr(result.observation, "model_dump") else dict(result.observation)
376
 
377
  print(
378
  f"[DEBUG] Server task: {obs.get('task_name')} "
@@ -401,7 +401,7 @@ async def run_task(client: OpenAI, task_name: str) -> float:
401
 
402
  reward = float(getattr(result, "reward", 0.0))
403
 
404
- obs = result.observation.model_dump() if hasattr(result.observation, "model_dump") else dict(result.observation)
405
  done = bool(result.done)
406
  error = None
407
 
@@ -418,7 +418,7 @@ async def run_task(client: OpenAI, task_name: str) -> float:
418
  except Exception as e:
419
  reward = 0.0
420
  done = True
421
- error = str(e)[:80]
422
 
423
  rewards.append(reward)
424
  steps_taken = step
 
111
  def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:
112
  error_val = error if error else "null"
113
  done_val = str(done).lower()
114
+ action_clean = action.replace("\n", " ").replace("\r", "")[:120] # type: ignore
115
  print(
116
  f"[STEP] step={step} action={action_clean} reward={reward:.2f} "
117
  f"done={done_val} error={error_val}",
 
262
  context_parts.append(f"Remaining questions: {obs.get('remaining_items', 0)}")
263
 
264
  if history:
265
+ context_parts.append("Recent actions:\n" + "\n".join(history[-2:])) # type: ignore
266
 
267
  if force_synthesize:
268
  context_parts.append(
 
299
  parts = text.split("```")
300
  text = parts[1] if len(parts) > 1 else parts[0]
301
  if text.startswith("json"):
302
+ text = text[4:] # type: ignore
303
  text = text.strip()
304
  # Strip control characters that break json.loads
305
  text = re.sub(r'[\x00-\x1f\x7f]', ' ', text)
 
372
 
373
  try:
374
  result = await env.reset()
375
+ obs: dict = result.observation.model_dump() if hasattr(result.observation, "model_dump") else dict(result.observation) # type: ignore
376
 
377
  print(
378
  f"[DEBUG] Server task: {obs.get('task_name')} "
 
401
 
402
  reward = float(getattr(result, "reward", 0.0))
403
 
404
+ obs: dict = result.observation.model_dump() if hasattr(result.observation, "model_dump") else dict(result.observation) # type: ignore
405
  done = bool(result.done)
406
  error = None
407
 
 
418
  except Exception as e:
419
  reward = 0.0
420
  done = True
421
+ error = str(e)[:80] # type: ignore
422
 
423
  rewards.append(reward)
424
  steps_taken = step