Spaces:
Sleeping
Sleeping
Kaushalraj Puwar
refactor: improve physical simulation stability with RK2 integration and update documentation across the environment and task modules
ba6f178 | from __future__ import annotations | |
| import threading | |
| from abc import ABC, abstractmethod | |
| from typing import Any, Dict, Tuple | |
| from env.core import ThermalPlantEnv | |
| class OpenEnvInterface(ABC): | |
| """ | |
| Abstract base class defining the contract for an OpenEnv-compliant interface. | |
| This ensures that the API layer can interact with any environment that | |
| adheres to this standard structure. | |
| """ | |
| def get_state(self) -> Dict[str, Any]: | |
| """Returns the full, unrounded internal state of the environment.""" | |
| raise NotImplementedError | |
| def reset(self, task_id: str, episode_id: int) -> Dict[str, float]: | |
| """Resets the environment to a new initial state for a given task and episode.""" | |
| raise NotImplementedError | |
| def step(self, action: Dict[str, float]) -> Tuple[Dict[str, float], float, bool, Dict[str, Any]]: | |
| """ | |
| Executes one time step in the environment. | |
| Args: | |
| action: A dictionary containing the agent's action. | |
| Returns: | |
| A tuple containing: | |
| - observation (Dict[str, float]): The agent's observation of the current environment state. | |
| - reward (float): The amount of reward returned after the previous action. | |
| - done (bool): Whether the episode has ended. | |
| - info (Dict[str, Any]): Contains auxiliary diagnostic information. | |
| """ | |
| raise NotImplementedError | |
| class ConcreteOpenEnvInterface(OpenEnvInterface): | |
| """ | |
| Concrete implementation of the OpenEnvInterface. | |
| This class provides a thread-safe singleton wrapper for the core thermal | |
| plant environment. It decouples the API layer from the simulation | |
| internals, ensuring consistent state management across sequential | |
| evaluation tasks. | |
| """ | |
| _instance: "ConcreteOpenEnvInterface" | None = None | |
| _env: ThermalPlantEnv | None = None | |
| _lock: threading.RLock = threading.RLock() | |
| def __new__(cls, max_steps: int | None = None) -> "ConcreteOpenEnvInterface": | |
| with cls._lock: | |
| if cls._instance is None: | |
| cls._instance = super().__new__(cls) | |
| if max_steps is None: | |
| cls._env = ThermalPlantEnv() | |
| else: | |
| cls._env = ThermalPlantEnv(max_steps=int(max_steps)) | |
| elif cls._env is not None and max_steps is not None: | |
| cls._env.max_steps = int(max_steps) | |
| return cls._instance | |
| def __init__(self, max_steps: int | None = None) -> None: | |
| """Allow non-API callers to configure the shared env without bypassing the interface.""" | |
| with self._lock: | |
| if self._env is not None and max_steps is not None: | |
| self._env.max_steps = int(max_steps) | |
| def get_state(self) -> Dict[str, Any]: | |
| """Returns the full, unrounded internal state from the core environment.""" | |
| with self._lock: | |
| if self._env is None: | |
| raise RuntimeError("Environment not initialized.") | |
| return self._env.state() | |
| def reset(self, task_id: str, episode_id: int) -> Dict[str, float]: | |
| """Resets the core environment and returns the initial observation.""" | |
| with self._lock: | |
| if self._env is None: | |
| raise RuntimeError("Environment not initialized.") | |
| return self._env.reset(task_id=task_id, episode_id=episode_id) | |
| def step(self, action: Dict[str, float]) -> Tuple[Dict[str, float], float, bool, Dict[str, Any]]: | |
| """Performs a step in the core environment.""" | |
| with self._lock: | |
| if self._env is None: | |
| raise RuntimeError("Environment not initialized.") | |
| return self._env.step(action) | |