Traffic-Control-Env / inference.py
dootisaha25's picture
Initial OpenEnv traffic signal environment
852e969
Raw
History Blame Contribute Delete
8.47 kB
"""Inference Script - Indian Traffic Signal OpenEnv
================================================
MANDATORY environment variables (injected by the validator):
API_BASE_URL The LiteLLM proxy endpoint.
API_KEY Your API key for the proxy.
MODEL_NAME The model identifier to use for inference.
STDOUT FORMAT (exact - do not deviate):
[START] task=<task_name> env=<benchmark> model=<model_name>
[STEP] step=<n> action=<action_str> reward=<0.00> done=<true|false> error=<msg|null>
[END] success=<true|false> steps=<n> score=<0.000> rewards=<r1,r2,...,rn>
"""
import json
import os
import sys
from typing import List, Optional
from openai import OpenAI
from env import IndianTrafficEnv
from grader import grade_rollout
from models import TrafficAction, TrafficState
# -------------------------------------------------------------------
# MANDATORY: read from injected environment variables — no hardcoding.
# The validator checks that all LLM calls flow through API_BASE_URL.
# -------------------------------------------------------------------
API_BASE_URL: str = os.environ["API_BASE_URL"] # must be set by validator
API_KEY: str = os.environ.get("API_KEY") or os.environ.get("HF_TOKEN", "")
MODEL_NAME: str = os.environ.get("MODEL_NAME", "Qwen/Qwen2.5-7B-Instruct")
BENCHMARK: str = os.environ.get("BENCHMARK", "indian-traffic-signal-openenv")
SUCCESS_SCORE_THRESHOLD = 0.5
TASKS = ["single_intersection", "rush_hour", "emergency_priority"]
VALID_ACTIONS = [a.value for a in TrafficAction]
# Single shared client — always routed through the injected proxy URL.
_client = OpenAI(
base_url=API_BASE_URL,
api_key=API_KEY,
timeout=30.0, # generous timeout for proxy round-trips
max_retries=1,
)
# ---------------------------------------------------------------------------
# Logging helpers — exact format required by the validator
# ---------------------------------------------------------------------------
def log_start(task: str, env: str, model: str) -> None:
print(f"[START] task={task} env={env} model={model}", flush=True)
def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:
error_val = error if error else "null"
print(
f"[STEP] step={step} action={action} reward={reward:.2f} "
f"done={str(done).lower()} error={error_val}",
flush=True,
)
def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:
rewards_str = ",".join(f"{r:.2f}" for r in rewards)
print(
f"[END] success={str(success).lower()} steps={steps} score={score:.3f} rewards={rewards_str}",
flush=True,
)
# ---------------------------------------------------------------------------
# Fallback policy — used only if the LLM call itself raises an exception
# ---------------------------------------------------------------------------
def _fallback_action(task_name: str, state: TrafficState) -> str:
"""Deterministic fallback — mirrors the baseline policy."""
if state.emergency_vehicle.present:
return TrafficAction.EMERGENCY_OVERRIDE.value
if state.pedestrian_count >= 16 and state.pedestrian_wait_time > 18:
return TrafficAction.PEDESTRIAN_CROSS.value
if state.time_since_last_phase_switch < 3 and state.current_signal_phase in (
TrafficAction.NS_GREEN,
TrafficAction.EW_GREEN,
TrafficAction.LEFT_PRIORITY,
):
return TrafficAction.EXTEND_GREEN.value
ns = state.lane_queues["N"].total + state.lane_queues["S"].total
ew = state.lane_queues["E"].total + state.lane_queues["W"].total
if task_name == "emergency_priority":
if abs(ns - ew) >= 10:
return TrafficAction.NS_GREEN.value if ns > ew else TrafficAction.EW_GREEN.value
return TrafficAction.NS_GREEN.value
if abs(ns - ew) >= 12:
return TrafficAction.NS_GREEN.value if ns > ew else TrafficAction.EW_GREEN.value
cycle = (state.tick // 8) % 4
return [
TrafficAction.NS_GREEN.value,
TrafficAction.EW_GREEN.value,
TrafficAction.LEFT_PRIORITY.value,
TrafficAction.PEDESTRIAN_CROSS.value,
][cycle]
# ---------------------------------------------------------------------------
# LLM call — ALWAYS goes through the injected proxy (API_BASE_URL / _client)
# ---------------------------------------------------------------------------
def get_action_from_llm(state: TrafficState, task_name: str) -> str:
"""Call the LLM via the injected proxy to choose a signal action."""
preferred = _fallback_action(task_name, state)
state_summary = {
"tick": state.tick,
"current_phase": state.current_signal_phase.value,
"time_since_switch": state.time_since_last_phase_switch,
"emergency": state.emergency_vehicle.model_dump(),
"pedestrian_count": state.pedestrian_count,
"pedestrian_wait": round(state.pedestrian_wait_time, 2),
"rain_level": round(state.rain_level, 3),
"lane_queues": {
lane: {"total": q.total, **q.model_dump()}
for lane, q in state.lane_queues.items()
},
}
system_prompt = (
"You are an AI traffic controller managing an Indian urban intersection. "
f"Task: {task_name}. "
f"Choose exactly one action from: {', '.join(VALID_ACTIONS)}. "
"Analyse the state and pick the best signal phase. "
f"Suggested action: {preferred}. "
"Reply with only the action name — no explanation, no punctuation."
)
user_prompt = f"Intersection state: {json.dumps(state_summary)}"
# This call MUST reach the proxy — do not wrap in a silent broad except.
response = _client.chat.completions.create(
model=MODEL_NAME,
messages=[
{"role": "system", "content": system_prompt},
{"role": "user", "content": user_prompt},
],
temperature=0.0,
max_tokens=16,
stream=False,
)
action = (response.choices[0].message.content or "").strip().upper()
return action if action in VALID_ACTIONS else preferred
# ---------------------------------------------------------------------------
# Main inference loop
# ---------------------------------------------------------------------------
def run_inference() -> None:
if not API_KEY:
print("Warning: API_KEY / HF_TOKEN not set.", file=sys.stderr, flush=True)
for task_name in TASKS:
env = IndianTrafficEnv(task_id=task_name)
env.reset(seed=42, task_id=task_name)
rewards: List[float] = []
steps_taken = 0
success = False
score = 0.001
done = False
log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)
try:
step = 1
while not done:
state = env.get_state()
error: Optional[str] = None
try:
action_str = get_action_from_llm(state, task_name)
except Exception as exc:
# LLM call failed — log it, use fallback, keep running
action_str = _fallback_action(task_name, state)
error = f"llm_error:{type(exc).__name__}"
try:
traffic_action = TrafficAction(action_str)
except ValueError:
traffic_action = TrafficAction.ALL_RED
action_str = TrafficAction.ALL_RED.value
try:
_, reward, done, _ = env.step(traffic_action)
except Exception as exc:
reward = 0.0
done = True
error = str(exc)
rewards.append(reward)
steps_taken = step
log_step(step=step, action=action_str, reward=reward, done=done, error=error)
step += 1
grader_result = grade_rollout(task_id=task_name, seed=42)
score = float(grader_result.score)
success = score >= SUCCESS_SCORE_THRESHOLD
except Exception as exc:
print(f"Fatal error in task {task_name}: {exc}", file=sys.stderr, flush=True)
success = False
score = 0.001
finally:
log_end(success=success, steps=steps_taken, score=score, rewards=rewards)
if __name__ == "__main__":
run_inference()