Spaces:
Sleeping
Sleeping
Upload folder using huggingface_hub
Browse files- 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:
|
| 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 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 254 |
-
error_msg = _sanitize(str(exc))
|
| 255 |
|
| 256 |
-
|
| 257 |
-
|
| 258 |
-
|
| 259 |
|
| 260 |
-
|
| 261 |
-
|
| 262 |
-
|
| 263 |
-
|
| 264 |
-
)
|
|
|
|
|
|
|
|
|
|
| 265 |
|
| 266 |
|
| 267 |
# ---------------------------------------------------------------------------
|
|
@@ -269,18 +273,20 @@ def run_task(task: str, client: OpenAI) -> None:
|
|
| 269 |
# ---------------------------------------------------------------------------
|
| 270 |
|
| 271 |
def main() -> None:
|
| 272 |
-
|
|
|
|
|
|
|
|
|
|
| 273 |
|
| 274 |
import time
|
| 275 |
-
#
|
| 276 |
-
for _ in range(
|
| 277 |
try:
|
| 278 |
-
r = _http.get(f"{SERVER_URL.rstrip('/')}/health", timeout=
|
| 279 |
if r.status_code == 200:
|
| 280 |
break
|
| 281 |
except Exception:
|
| 282 |
-
|
| 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)
|