Spaces:
Running
Running
Commit ·
386bf01
1
Parent(s): b4b9b2e
final one
Browse files- 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
|
| 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
|