channu07/microgrid-env-v2
0
1"""2Baseline Inference Script - MicrogridEnv v23Runs all 3 tasks and emits [START] [STEP] [END] blocks per task.4"""5import asyncio6import json7import os8import re9import textwrap10from typing import List, Optional11 12from openai import OpenAI13 14from microgrid_env_v2 import MicrogridEnv, MicrogridAction15 16IMAGE_NAME = os.getenv("LOCAL_IMAGE_NAME")17API_KEY = os.getenv("HF_TOKEN") or os.getenv("API_KEY")18API_BASE_URL = os.getenv("API_BASE_URL", "https://router.huggingface.co/v1")19MODEL_NAME = os.getenv("MODEL_NAME", "Qwen/Qwen2.5-72B-Instruct")20 21TASKS = ["steady_state", "storm_day", "cascade_fault"]22BENCHMARK = "microgrid_env_v2"23MAX_STEPS = 824TEMPERATURE = 0.325MAX_TOKENS = 25026SUCCESS_SCORE_THRESHOLD = 0.327 28SYSTEM_PROMPT = textwrap.dedent("""29You are an autonomous microgrid operator. Each turn, output ONE JSON action:30{"battery_power": float, "grid_import": float, "load_shed_frac": float, "dispatch_priority": int}31 32Constraints:33- battery_power in [-5, +5] MW (positive=discharge, negative=charge)34- grid_import in [0, 10] MW (expensive)35- load_shed_frac in [0, 0.3] (shed load only if necessary)36- dispatch_priority in {0,1,2}37 38Objective: keep frequency near 50 Hz, minimize cost, avoid unserved load.39- If solar+wind > load: charge battery (negative battery_power)40- If solar+wind < load: discharge battery first, then grid41- If fault_flag==1: shed 15-20% load and reduce battery output to isolate fault42- If frequency drifts from 50 Hz: reduce power imbalance43 44Reply with ONLY the JSON object, no explanation.45""").strip()46 47 48def log_start(task: str, env: str, model: str) -> None:49 print(f"[START] task={task} env={env} model={model}", flush=True)50 51 52def log_step(step: int, action: str, reward: float, done: bool, error: Optional[str]) -> None:53 err = error if error else "null"54 print(55 f"[STEP] step={step} action={action} reward={reward:.2f} "56 f"done={str(done).lower()} error={err}",57 flush=True,58 )59 60 61def log_end(success: bool, steps: int, score: float, rewards: List[float]) -> None:62 r_str = ",".join(f"{r:.2f}" for r in rewards)63 print(64 f"[END] success={str(success).lower()} steps={steps} "65 f"score={score:.3f} rewards={r_str}",66 flush=True,67 )68 69 70def parse_action(text: str) -> MicrogridAction:71 try:72 match = re.search(r"\{.*\}", text, re.DOTALL)73 if match:74 data = json.loads(match.group())75 return MicrogridAction(76 battery_power=float(data.get("battery_power", 0.0)),77 grid_import=float(data.get("grid_import", 0.0)),78 load_shed_frac=float(data.get("load_shed_frac", 0.0)),79 dispatch_priority=int(data.get("dispatch_priority", 0)),80 )81 except Exception:82 pass83 return MicrogridAction()84 85 86def get_action(client: OpenAI, obs: dict, step: int, history: List[str]) -> MicrogridAction:87 hist = "\n".join(history[-3:]) if history else "None"88 user = f"Step {step}\nObservation: {obs}\nPrevious:\n{hist}\n\nOutput JSON action:"89 try:90 comp = client.chat.completions.create(91 model=MODEL_NAME,92 messages=[93 {"role": "system", "content": SYSTEM_PROMPT},94 {"role": "user", "content": user},95 ],96 temperature=TEMPERATURE,97 max_tokens=MAX_TOKENS,98 stream=False,99 )100 return parse_action((comp.choices[0].message.content or "").strip())101 except Exception as exc:102 print(f"[DEBUG] model error: {exc}", flush=True)103 return MicrogridAction()104 105 106async def run_task(task_name: str, client: OpenAI) -> None:107 if IMAGE_NAME:108 env = await MicrogridEnv.from_docker_image(IMAGE_NAME)109 else:110 env = MicrogridEnv(base_url=os.getenv("ENV_URL", "http://localhost:7860"))111 112 history: List[str] = []113 rewards: List[float] = []114 steps_taken = 0115 score = 0.0116 success = False117 118 log_start(task=task_name, env=BENCHMARK, model=MODEL_NAME)119 120 try:121 result = await env.reset(task=task_name, seed=42)122 obs = result.observation123 124 for step in range(1, MAX_STEPS + 1):125 if result.done:126 break127 action = get_action(client, obs.model_dump(), step, history)128 act_str = (129 f"batt={action.battery_power:.1f},grid={action.grid_import:.1f},"130 f"shed={action.load_shed_frac:.2f},pri={action.dispatch_priority}"131 )132 result = await env.step(action)133 obs = result.observation134 reward = result.reward or 0.0135 rewards.append(reward)136 steps_taken = step137 log_step(step=step, action=act_str, reward=reward, done=result.done, error=None)138 history.append(139 f"s{step}: r={reward:+.2f} freq={obs.frequency_hz:.2f} "140 f"soc={obs.battery_soc:.2f} fault={obs.fault_flag}"141 )142 if result.done:143 break144 145 if rewards:146 score = min(max((sum(rewards) + 2 * len(rewards)) / (3.5 * len(rewards)), 0.0), 1.0)147 success = score >= SUCCESS_SCORE_THRESHOLD148 149 finally:150 try:151 await env.close()152 except Exception as e:153 print(f"[DEBUG] close error: {e}", flush=True)154 log_end(success=success, steps=steps_taken, score=score, rewards=rewards)155 156 157async def main() -> None:158 client = OpenAI(base_url=API_BASE_URL, api_key=API_KEY)159 for task in TASKS:160 await run_task(task, client)161 162 163if __name__ == "__main__":164 asyncio.run(main())