Spaces:
Sleeping
Sleeping
| """ | |
| inference.py β Space Manufacturing RL submission entry point. | |
| Default policy: OpenAI (falls back to heuristic if the client cannot be built). | |
| Environment variables: | |
| API_BASE_URL β OpenAI-compatible endpoint base URL (required) | |
| API_KEY β API key (required) | |
| MODEL_NAME β Model to use (required) | |
| BASELINE_POLICY β Force policy: "openai" (default) or "heuristic" | |
| TEMPERATURE β Sampling temperature (default: 0.0) | |
| MAX_TOKENS β Max tokens per response (default: 300) | |
| REQUEST_DELAY β Seconds to sleep between steps (default: 0.0) | |
| REQUEST_TIMEOUT β HTTP timeout in seconds (default: 30.0) | |
| STEP_TIMEOUT β Per-step inference wall-clock timeout in seconds (default: 45.0) | |
| TASK_TIMEOUT β Per-task wall-clock timeout in seconds, 0 = no limit (default: 0.0) | |
| DEBUG β Print raw model responses when "true" | |
| Usage: | |
| API_BASE_URL=https://... API_KEY=hf_... MODEL_NAME=mistralai/... python inference.py | |
| # Force heuristic baseline: | |
| BASELINE_POLICY=heuristic python inference.py | |
| # Timeouts: | |
| STEP_TIMEOUT=20 TASK_TIMEOUT=300 python inference.py | |
| """ | |
| from __future__ import annotations | |
| import asyncio | |
| import json | |
| import os | |
| import re | |
| import sys | |
| import textwrap | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from typing import Any, Dict, List, Optional | |
| from urllib.parse import urlparse | |
| try: | |
| from openai import OpenAI | |
| except ImportError: # pragma: no cover | |
| OpenAI = None # type: ignore[assignment,misc] | |
| try: | |
| from dotenv import load_dotenv | |
| except Exception: # pragma: no cover | |
| load_dotenv = None # type: ignore[assignment] | |
| # ββ package imports ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| # inference.py is a script inside SpaceFactory/. Insert the parent directory so | |
| # the whole folder is importable as the 'SpaceFactory' package, which keeps all | |
| # relative imports inside the package working correctly. | |
| _pkg_parent = str(Path(__file__).resolve().parent.parent) | |
| if _pkg_parent not in sys.path: | |
| sys.path.insert(0, _pkg_parent) | |
| from SpaceFactory.env import ManufacturingTaskEnv | |
| from SpaceFactory.graders import ManufacturingTaskGrader | |
| from SpaceFactory.models import ManufacturingAction, ManufacturingObservation | |
| from SpaceFactory.tasks import EasyTask, HardTask, MediumTask | |
| if load_dotenv is not None: | |
| load_dotenv() | |
| # ββ warn_once ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| WARNINGS_EMITTED: set[str] = set() | |
| def warn_once(key: str, message: str) -> None: | |
| if key in WARNINGS_EMITTED: | |
| return | |
| WARNINGS_EMITTED.add(key) | |
| print(f"[warn] {message}", file=sys.stderr) | |
| # ββ env helpers ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def read_float_env(name: str, default: float) -> float: | |
| raw = os.getenv(name) | |
| if raw is None: | |
| return default | |
| try: | |
| return float(raw) | |
| except (TypeError, ValueError): | |
| warn_once(f"env:{name}", f"Invalid {name}={raw!r}; using default {default}.") | |
| return default | |
| def read_int_env(name: str, default: int) -> int: | |
| raw = os.getenv(name) | |
| if raw is None: | |
| return default | |
| try: | |
| return int(raw) | |
| except (TypeError, ValueError): | |
| warn_once(f"env:{name}", f"Invalid {name}={raw!r}; using default {default}.") | |
| return default | |
| # ββ configuration ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| API_BASE_URL = os.getenv("API_BASE_URL","https://router.huggingface.co/v1") | |
| API_KEY = os.getenv("API_KEY") | |
| MODEL_NAME = os.getenv("MODEL_NAME","meta-llama/Llama-3.1-8B-Instruct:novita") | |
| BASELINE_POLICY = os.getenv("BASELINE_POLICY", "openai").lower() | |
| TEMPERATURE = read_float_env("TEMPERATURE", 0.0) | |
| MAX_TOKENS = read_int_env("MAX_TOKENS", 300) | |
| REQUEST_DELAY = read_float_env("REQUEST_DELAY", 0.0) | |
| REQUEST_TIMEOUT = read_float_env("REQUEST_TIMEOUT", 30.0) | |
| STEP_TIMEOUT = read_float_env("STEP_TIMEOUT", 45.0) # per-step wall-clock limit | |
| TASK_TIMEOUT = read_float_env("TASK_TIMEOUT", 0.0) # per-task limit; 0 = no limit | |
| DEBUG = os.getenv("DEBUG", "false").lower() == "true" | |
| FALLBACK_ACTION = "recharge" | |
| TASK_ORDER = ["easy", "medium", "hard"] | |
| TASK_TYPES = {"easy": EasyTask, "medium": MediumTask, "hard": HardTask} | |
| VALID_ACTIONS = {"produce", "assemble", "deliver", "recharge"} | |
| ACTION_PATTERN = re.compile(r"(produce|assemble|deliver|recharge)", re.IGNORECASE) | |
| _SCORE_EPS = 1e-9 # keeps every score strictly inside (0, 1) | |
| def _clamp_score(value: float) -> float: | |
| """Clamp *value* to the open interval (0, 1) exclusive.""" | |
| return max(_SCORE_EPS, min(1.0 - _SCORE_EPS, float(value))) | |
| # ββ system prompt ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| SYSTEM_PROMPT = textwrap.dedent(""" | |
| You are controlling orbital manufacturing platforms. | |
| Each step, output ONLY a JSON object mapping platform IDs (as strings) to one of: | |
| "produce", "assemble", "deliver", "recharge" | |
| Decision guidance: | |
| - recharge immediately if energy < 15 | |
| - deliver when product_stock > 0 and a delivery window is open | |
| - assemble when component_stock >= 10 and product_stock < 5 | |
| - produce when material_stock >= 15 and component_stock < 30 | |
| - recharge when energy < 40 and no urgent action is available | |
| - avoid invalid actions (e.g. assemble with no components) | |
| - keep all platforms energy-healthy across the full episode | |
| Output format (no explanation, no markdown): | |
| {"0": "produce", "1": "assemble", "2": "deliver"} | |
| """).strip() | |
| # ββ result dataclass ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| class TaskRunResult: | |
| task_name: str | |
| score: float | |
| total_reward: float | |
| steps: int | |
| done: bool | |
| metrics: Dict[str, float] | |
| # ββ observation helpers ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def observation_to_dict(obs: ManufacturingObservation) -> Dict[str, Any]: | |
| return { | |
| "platforms": [ | |
| { | |
| "id": p.id, | |
| "energy": p.energy, | |
| "material_stock": p.material_stock, | |
| "component_stock": p.component_stock, | |
| "product_stock": p.product_stock, | |
| "last_action": p.last_action, | |
| } | |
| for p in obs.platforms | |
| ], | |
| "time_step": obs.time_step, | |
| "delivery_windows": [ | |
| {"order_id": w.order_id, "product_type": w.product_type, "deadline": w.deadline} | |
| for w in obs.delivery_windows | |
| ], | |
| "solar_conditions": obs.solar_conditions, | |
| "pending_orders": [ | |
| {"order_id": o.order_id, "product_type": o.product_type, | |
| "requires_assembly": o.requires_assembly} | |
| for o in obs.pending_orders | |
| ], | |
| "total_reward": obs.total_reward, | |
| "done": obs.done, | |
| "reward": obs.reward, | |
| "metadata": obs.metadata, | |
| } | |
| def build_idle_actions(obs_dict: Dict[str, Any]) -> Dict[int, str]: | |
| return {int(p["id"]): FALLBACK_ACTION for p in obs_dict.get("platforms", [])} | |
| # ββ heuristic policy βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def heuristic_action(obs_dict: Dict[str, Any]) -> Dict[int, str]: | |
| actions: Dict[int, str] = {} | |
| has_open_window = len(obs_dict.get("delivery_windows", [])) > 0 | |
| for p in obs_dict.get("platforms", []): | |
| pid = int(p["id"]) | |
| energy = float(p["energy"]) | |
| mat = float(p["material_stock"]) | |
| comp = float(p["component_stock"]) | |
| prod = int(p["product_stock"]) | |
| if energy < 15.0: | |
| action = "recharge" | |
| elif prod > 0 and has_open_window: | |
| action = "deliver" | |
| elif comp >= 10.0 and prod < 5: | |
| action = "assemble" | |
| elif mat >= 15.0 and comp < 30.0: | |
| action = "produce" | |
| elif energy < 40.0: | |
| action = "recharge" | |
| elif mat >= 15.0: | |
| action = "produce" | |
| else: | |
| action = "recharge" | |
| actions[pid] = action | |
| return actions | |
| def safe_heuristic_action(obs_dict: Dict[str, Any], reason: str) -> Dict[int, str]: | |
| try: | |
| return heuristic_action(obs_dict) | |
| except Exception as exc: # noqa: BLE001 | |
| warn_once( | |
| f"heuristic:{reason}", | |
| f"Heuristic fallback failed after {reason}: {exc}. Returning all-{FALLBACK_ACTION}.", | |
| ) | |
| return build_idle_actions(obs_dict) | |
| # ββ prompt formatting ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def build_history_lines(history: List[str]) -> str: | |
| return "\n".join(history[-6:]) if history else "None" | |
| def format_observation(task_name: str, obs_dict: Dict[str, Any]) -> str: | |
| platforms_lines = [] | |
| for p in obs_dict.get("platforms", []): | |
| platforms_lines.append( | |
| f" [{p['id']}] energy={p['energy']:.1f} mat={p['material_stock']:.1f}" | |
| f" comp={p['component_stock']:.1f} prod={p['product_stock']}" | |
| f" last={p['last_action']}" | |
| ) | |
| windows_lines = [] | |
| for w in obs_dict.get("delivery_windows", []): | |
| step_now = obs_dict.get("time_step", 0) | |
| urgency = w["deadline"] - step_now | |
| windows_lines.append( | |
| f" order={w['order_id']} type={w['product_type']}" | |
| f" deadline={w['deadline']} ({urgency} steps left)" | |
| ) | |
| solar_str = ", ".join( | |
| f"{z}={round(v * 100)}%" | |
| for z, v in obs_dict.get("solar_conditions", {}).items() | |
| ) | |
| return textwrap.dedent(f""" | |
| Task: {task_name} | |
| Time Step: {obs_dict.get('time_step', 0)} | |
| Total Reward: {obs_dict.get('total_reward', 0.0):.2f} | |
| Solar: {solar_str or 'n/a'} | |
| Platforms: | |
| {chr(10).join(platforms_lines) or ' None'} | |
| Open Delivery Windows: | |
| {chr(10).join(windows_lines) if windows_lines else ' None'} | |
| Pending orders: {len(obs_dict.get('pending_orders', []))} | |
| """).strip() | |
| def build_user_prompt( | |
| task_name: str, | |
| step: int, | |
| obs_dict: Dict[str, Any], | |
| history: List[str], | |
| total_reward: float, | |
| ) -> str: | |
| return textwrap.dedent(f""" | |
| Step: {step} | |
| Aggregate reward so far: {total_reward:+.2f} | |
| Current state: | |
| {format_observation(task_name, obs_dict)} | |
| Previous steps: | |
| {build_history_lines(history)} | |
| Reply with exactly one JSON object. | |
| """).strip() | |
| # ββ model response parsing ββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def extract_response_text(completion: Any) -> str: | |
| choices = getattr(completion, "choices", None) | |
| if not choices: | |
| return "" | |
| message = getattr(choices[0], "message", None) | |
| if message is None: | |
| return "" | |
| content = getattr(message, "content", "") | |
| if isinstance(content, str): | |
| return content | |
| if isinstance(content, list): | |
| parts: List[str] = [] | |
| for item in content: | |
| text = item.get("text") if isinstance(item, dict) else getattr(item, "text", None) | |
| if text: | |
| parts.append(str(text)) | |
| return "\n".join(parts) | |
| return str(content or "") | |
| def parse_model_action( | |
| response_text: str, obs_dict: Dict[str, Any] | |
| ) -> Dict[int, str]: | |
| if not response_text: | |
| return safe_heuristic_action(obs_dict, "empty model response") | |
| try: | |
| json_match = re.search(r"\{.*\}", response_text.strip(), re.DOTALL) | |
| if json_match: | |
| parsed = json.loads(json_match.group(0)) | |
| valid_ids = {int(p["id"]) for p in obs_dict.get("platforms", [])} | |
| actions: Dict[int, str] = {} | |
| for key, value in parsed.items(): | |
| pid = int(key) | |
| if pid not in valid_ids: | |
| continue | |
| action = str(value).strip().lower() | |
| if action not in VALID_ACTIONS: | |
| action = FALLBACK_ACTION | |
| actions[pid] = action | |
| if actions: | |
| fallback = heuristic_action(obs_dict) | |
| for pid in valid_ids: | |
| actions.setdefault(pid, fallback.get(pid, FALLBACK_ACTION)) | |
| return actions | |
| except (json.JSONDecodeError, TypeError, ValueError): | |
| pass | |
| warn_once("parse:model-response", "Model response was not valid JSON; using heuristic.") | |
| return safe_heuristic_action(obs_dict, "invalid model response") | |
| # ββ client construction ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def validate_api_base_url(base_url: Optional[str]) -> Optional[str]: | |
| if not base_url: | |
| return None | |
| cleaned = base_url.strip().rstrip("/") | |
| parsed = urlparse(cleaned) | |
| if parsed.scheme not in {"http", "https"} or not parsed.netloc: | |
| warn_once( | |
| "config:api-base-url", | |
| f"Invalid API_BASE_URL={base_url!r}; falling back to heuristic policy.", | |
| ) | |
| return None | |
| return cleaned | |
| def build_client() -> Optional[Any]: | |
| if BASELINE_POLICY == "heuristic": | |
| return None | |
| if OpenAI is None: | |
| warn_once("client:import", "openai package not installed; falling back to heuristic policy.") | |
| return None | |
| if not API_KEY: | |
| warn_once( | |
| "config:missing", | |
| "OPENAI_API_KEY is not set. Falling back to heuristic policy.", | |
| ) | |
| return None | |
| validated_base = validate_api_base_url(API_BASE_URL) # None = use OpenAI default endpoint | |
| try: | |
| kwargs: Dict[str, Any] = {"api_key": API_KEY, "timeout": REQUEST_TIMEOUT} | |
| if validated_base: | |
| kwargs["base_url"] = validated_base | |
| return OpenAI(**kwargs) | |
| except Exception as exc: # noqa: BLE001 | |
| warn_once("client:init", f"Failed to build OpenAI client: {exc}. Using heuristic.") | |
| return None | |
| # ββ action chooser βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| async def choose_actions( | |
| client: Optional[Any], | |
| task_name: str, | |
| step: int, | |
| obs_dict: Dict[str, Any], | |
| history: List[str], | |
| total_reward: float, | |
| ) -> Dict[int, str]: | |
| if client is None: | |
| return safe_heuristic_action(obs_dict, "heuristic mode") | |
| user_prompt = build_user_prompt(task_name, step, obs_dict, history, total_reward) | |
| def _call() -> str: | |
| completion = client.chat.completions.create( | |
| model=MODEL_NAME, | |
| messages=[ | |
| {"role": "system", "content": SYSTEM_PROMPT}, | |
| {"role": "user", "content": user_prompt}, | |
| ], | |
| temperature=TEMPERATURE, | |
| max_tokens=MAX_TOKENS, | |
| ) | |
| return extract_response_text(completion) | |
| try: | |
| loop = asyncio.get_event_loop() | |
| response_text = await asyncio.wait_for( | |
| loop.run_in_executor(None, _call), | |
| timeout=STEP_TIMEOUT, | |
| ) | |
| except asyncio.TimeoutError: | |
| warn_once( | |
| f"timeout:{task_name}", | |
| f"[{task_name}] Step {step} timed out after {STEP_TIMEOUT}s. Using heuristic.", | |
| ) | |
| return safe_heuristic_action(obs_dict, f"step timeout on {task_name} step {step}") | |
| except Exception as exc: # noqa: BLE001 | |
| warn_once( | |
| f"model:{task_name}", | |
| f"[{task_name}] Model request failed at step {step}: {exc}. Using heuristic.", | |
| ) | |
| return safe_heuristic_action(obs_dict, f"model failure on {task_name} step {step}") | |
| if DEBUG: | |
| print(f"[DEBUG] [{task_name}] step={step} model_response={response_text[:300]!r}") | |
| try: | |
| return parse_model_action(response_text, obs_dict) | |
| except Exception as exc: # noqa: BLE001 | |
| warn_once( | |
| f"parse:{task_name}", | |
| f"[{task_name}] Parse failed at step {step}: {exc}. Using heuristic.", | |
| ) | |
| return safe_heuristic_action(obs_dict, f"parse failure on {task_name} step {step}") | |
| # ββ episode runner βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| async def run_task(task_name: str, client: Optional[Any]) -> TaskRunResult: | |
| env = ManufacturingTaskEnv(task_name=task_name) | |
| grader = ManufacturingTaskGrader(task_name=task_name) | |
| history: List[str] = [] | |
| obs = env.reset() | |
| state = env.state() | |
| step_limit = state.max_steps | |
| print(f"[START] task={task_name} max_steps={step_limit}", flush=True) | |
| for step in range(1, step_limit + 1): | |
| obs_dict = observation_to_dict(obs) | |
| actions = await choose_actions(client, task_name, step, obs_dict, history, obs.total_reward) | |
| # Wrap dict back into ManufacturingAction | |
| action_obj = ManufacturingAction(platform_actions=actions) | |
| obs, reward, done, info = env.step(action_obj) | |
| reward_value = float(reward.value) | |
| history.append(f"step {step}: {actions} -> reward {reward_value:+.2f}") | |
| print( | |
| f"[STEP] task={task_name} step={step}/{step_limit}" | |
| f" reward={reward_value:.4f} total={obs.total_reward:.4f}" | |
| f" done={done}", | |
| flush=True, | |
| ) | |
| if REQUEST_DELAY > 0 and not done: | |
| await asyncio.sleep(REQUEST_DELAY) | |
| if done: | |
| break | |
| final_state = env.state() | |
| metrics = {k: float(v) for k, v in final_state.metrics.items()} | |
| score = _clamp_score(grader.grade(metrics, final_state.step_count, final_state.platforms)) | |
| print( | |
| f"[END] task={task_name} score={score:.4f}" | |
| f" steps={final_state.step_count} total_reward={final_state.total_reward:.4f}" | |
| f" done={final_state.done}", | |
| flush=True, | |
| ) | |
| return TaskRunResult( | |
| task_name=task_name, | |
| score=score, | |
| total_reward=final_state.total_reward, | |
| steps=final_state.step_count, | |
| done=final_state.done, | |
| metrics=metrics, | |
| ) | |
| # ββ summary printer ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| def print_summary(results: List[TaskRunResult]) -> None: | |
| aggregate = sum(r.score for r in results) / len(results) | |
| print("\nInference Summary") | |
| print("=" * 60) | |
| for r in results: | |
| print( | |
| f"{r.task_name:<8} score={r.score:.4f}" | |
| f" reward={r.total_reward:.2f}" | |
| f" steps={r.steps}" | |
| f" done={r.done}" | |
| ) | |
| print("-" * 60) | |
| print(f"aggregate_score={aggregate:.4f}") | |
| print("=" * 60) | |
| # ββ async main βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ | |
| async def async_main() -> None: | |
| client = build_client() | |
| results = [] | |
| for task_name in TASK_ORDER: | |
| if TASK_TIMEOUT > 0: | |
| try: | |
| result = await asyncio.wait_for(run_task(task_name, client), timeout=TASK_TIMEOUT) | |
| except asyncio.TimeoutError: | |
| print( | |
| f"[TIMEOUT] task={task_name} exceeded {TASK_TIMEOUT}s; skipping.", | |
| file=sys.stderr, | |
| flush=True, | |
| ) | |
| continue | |
| else: | |
| result = await run_task(task_name, client) | |
| results.append(result) | |
| if results: | |
| print_summary(results) | |
| def main() -> None: | |
| try: | |
| asyncio.run(async_main()) | |
| except KeyboardInterrupt: | |
| print("\nInference interrupted.", file=sys.stderr) | |
| raise SystemExit(130) from None | |
| except Exception as exc: # noqa: BLE001 | |
| print(f"\nInference failed: {exc}", file=sys.stderr) | |
| raise SystemExit(1) from None | |
| if __name__ == "__main__": | |
| main() | |