CoolFace
Apppublic

channu07/microgrid-env-v2

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
inference.py164 linesDownload Raw Back to root
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())