Team Ai
Apppublic

openenv/echo_env

sourceHugging Faceupdated 2d agoView on Hugging Face
6likes
interfaces.py431 linesDownload Raw Back to env_server
1# SPDX-License-Identifier: BSD-3-Clause2 3import inspect4from abc import ABC, abstractmethod5from typing import Any, Generic, Optional, Protocol, TYPE_CHECKING, TypeVar6 7from typing_extensions import TypedDict8 9from .types import Action, EnvironmentMetadata, Observation, State10 11if TYPE_CHECKING:12    from openenv.core.rubrics import Rubric13 14ActT = TypeVar("ActT", bound=Action)15ObsT = TypeVar("ObsT", bound=Observation)16StateT = TypeVar("StateT", bound=State)17 18 19class Message(TypedDict):20    """A message in a conversation.21 22    Compatible with Huggingface chat template format.23    """24 25    role: str26    content: str27 28 29class ModelTokenizer(Protocol):30    """Protocol for tokenizers that support chat templates.31 32    This protocol defines the interface that tokenizers must implement33    to work with chat-based environments. It's compatible with34    Huggingface transformers tokenizers.35    """36 37    def apply_chat_template(38        self,39        conversation: list[Message],40        tokenize: bool = True,41        return_tensors: str | None = None,42        **kwargs: Any,43    ) -> Any:44        """Apply a chat template to format and optionally tokenize a conversation.45 46        Args:47            conversation (`list[Message]`):48                List of message dictionaries with 'role' and 'content'.49            tokenize (`bool`, *optional*, defaults to `True`):50                Whether to tokenize the output.51            return_tensors (`str`, *optional*):52                Format for returned tensors ('pt' for PyTorch).53            **kwargs:54                Additional arguments.55 56        Returns:57            Formatted and optionally tokenized conversation.58        """59        ...60 61    def decode(62        self, token_ids: Any, skip_special_tokens: bool = False, **kwargs: Any63    ) -> str:64        """Decode token IDs back to text.65 66        Args:67            token_ids (`Any`):68                Token IDs to decode.69            skip_special_tokens (`bool`, *optional*, defaults to `False`):70                Whether to skip special tokens in output.71            **kwargs:72                Additional arguments.73 74        Returns:75            `str`: Decoded text string.76        """77        ...78 79 80class TaskProvider(Protocol):81    """82    Optional task discovery API for dataset-backed environments.83 84    An environment implements this protocol structurally — declare the methods on85    an [`~openenv.core.env_server.interfaces.Environment`] subclass, without86    inheriting from `TaskProvider`. When the methods are present,87    [`~openenv.core.env_server.http_server.HTTPEnvServer`] exposes them as HTTP88    routes under `/{env_name}/…`; when they are absent, those routes return89    `501 Not Implemented`. Each method may be sync or async.90 91    Task provider methods are for metadata/discovery only and should be92    side-effect-free. They must be callable on a freshly constructed93    environment instance because HTTP compatibility routes may create a94    short-lived instance solely for task discovery.95 96    Selecting a task is not part of this protocol — pass the chosen split and97    index to `reset()` instead. See the98    [Task API guide](https://huggingface.co/docs/openenv/guides/task-api).99 100    Examples:101 102    ```python103    env.list_splits()          # ["train", "test"]104    env.num_tasks("test")      # 7595105    env.get_task("test", 12)   # {"id": "test-12", "index": 12}106    env.reset(split="test", index=12)107    ```108    """109 110    def list_splits(self) -> list[Any]:111        """112        Return task split descriptors supported by this environment.113 114        Returns:115            `list[Any]`: Split descriptors. Plain strings, dicts, and Pydantic116                models are all accepted; the server normalizes each entry to117                `{"name": ..., "type": ...}`.118        """119        ...120 121    def list_tasks(self, split: str) -> list[Any]:122        """123        Return all task specs for a split.124 125        Args:126            split (`str`):127                Task split name.128 129        Returns:130            `list[Any]`: Task specs for the split. Environments backed by very131                large or streamed splits may return a bounded preview, but132                `num_tasks` should still report the true total.133        """134        ...135 136    def num_tasks(self, split: str) -> int:137        """138        Return the number of task specs in a split.139 140        Args:141            split (`str`):142                Task split name.143 144        Returns:145            `int`: Number of task specs available in the split.146        """147        ...148 149    def get_task(self, split: str, index: int) -> Any:150        """151        Return one task spec by split and index.152 153        Args:154            split (`str`):155                Task split name.156            index (`int`):157                Task index within the split.158 159        Returns:160            `Any`: The task spec at that position.161 162        Raises:163            `IndexError`: If `index` is out of range for the split. The HTTP164                route converts this to a `400 Bad Request`.165        """166        ...167 168    def get_task_range(169        self,170        split: str,171        start: Optional[int] = None,172        stop: Optional[int] = None,173    ) -> list[Any]:174        """175        Return task specs for Python slice-style range bounds.176 177        Args:178            split (`str`):179                Task split name.180            start (`int`, *optional*):181                Inclusive start index. Defaults to the beginning of the split.182            stop (`int`, *optional*):183                Exclusive stop index. Defaults to the end of the split.184 185        Returns:186            `list[Any]`: Task specs in `[start, stop)`.187        """188        ...189 190 191class Transform(ABC, Generic[ObsT]):192    """Transform observations to add rewards, metrics, or other modifications.193 194    Transforms follow the TorchRL pattern where they take an observation195    and return a (potentially modified) observation. This allows for196    flexible reward computation and observation augmentation.197    """198 199    @abstractmethod200    def __call__(self, observation: ObsT) -> ObsT:201        """Transform an observation.202 203        Args:204            observation (`ObsT`):205                The input observation.206 207        Returns:208            `ObsT`: The transformed observation.209        """210        pass211 212 213class Environment(ABC, Generic[ActT, ObsT, StateT]):214    """Base class for all environment servers following Gym/Gymnasium API.215 216    See [rfcs/004-rubrics.md](https://github.com/huggingface/OpenEnv/blob/main/rfcs/004-rubrics.md) for rubric design details.217 218    Args:219        transform (`Transform`, *optional*):220            Optional transform to apply to observations.221        rubric (`Rubric`, *optional*):222            Optional rubric for reward computation. When provided, the223            rubric's output can be used to set the observation's reward in step().224 225    Attributes:226        SUPPORTS_CONCURRENT_SESSIONS (`bool`):227            Whether this environment supports concurrent sessions. When ``True``,228            multiple WebSocket connections can each have their own environment229            instance (up to ``max_concurrent_envs``). When ``False`` (default),230            the environment should only be used with a single session at a time.231 232            Set this to ``True`` in your subclass if the environment uses proper233            session isolation (unique working dirs, no shared mutable state, and234            external resources that can handle concurrent access).235        rubric (`Rubric`, *optional*):236            Optional rubric for computing rewards. Set in ``__init__`` and use in237            ``step()`` to compute observation rewards. Training infrastructure can238            access it for introspection:239 240            ```python241            for name, r in env.rubric.named_rubrics():242                print(f"{name}: {r.last_score}")243            ```244    """245 246    # Class-level flag indicating whether this environment supports concurrent sessions247    SUPPORTS_CONCURRENT_SESSIONS: bool = False248 249    REQUIRES_SINGLE_THREAD_EXECUTOR: bool = False250 251    # Optional rubric for reward computation252    rubric: Optional["Rubric"]253 254    def __init__(255        self,256        transform: Optional[Transform[ObsT]] = None,257        rubric: Optional["Rubric"] = None,258    ):259        self.transform = transform260        self.rubric = rubric261 262    @abstractmethod263    def reset(264        self,265        seed: Optional[int] = None,266        episode_id: Optional[str] = None,267        **kwargs: Any,268    ) -> ObsT:269        """Reset the environment and return initial observation."""270        pass271 272    async def reset_async(273        self,274        seed: Optional[int] = None,275        episode_id: Optional[str] = None,276        **kwargs: Any,277    ) -> ObsT:278        """Async version of reset. Default implementation calls sync reset.279 280        Override to provide true async implementation.281        """282        return self.reset(seed=seed, episode_id=episode_id, **kwargs)283 284    @abstractmethod285    def step(286        self,287        action: ActT,288        timeout_s: Optional[float] = None,289        **kwargs: Any,290    ) -> ObsT:291        """Take a step in the environment."""292        pass293 294    async def step_async(295        self,296        action: ActT,297        timeout_s: Optional[float] = None,298        **kwargs: Any,299    ) -> ObsT:300        """Async version of step. Default implementation calls sync step.301 302        Override to provide true async implementation.303        """304        return self.step(action, timeout_s=timeout_s, **kwargs)305 306    @property307    @abstractmethod308    def state(self) -> StateT:309        """Get the current environment state."""310        pass311 312    def get_metadata(self) -> EnvironmentMetadata:313        """314        Get metadata about this environment.315 316        Override this method to provide custom metadata for the environment.317        Default implementation returns basic metadata derived from class name.318 319        Returns:320            [`EnvironmentMetadata`] with environment information.321        """322        return EnvironmentMetadata(323            name=self.__class__.__name__,324            description=f"{self.__class__.__name__} environment",325            version="1.0.0",326        )327 328    def _apply_transform(self, observation: ObsT) -> ObsT:329        """Apply transform if one is provided."""330        if self.transform is not None:331            return self.transform(observation)332        return observation333 334    def _apply_rubric(self, action: ActT, observation: ObsT) -> float:335        """Apply rubric if one is provided.336 337        Args:338            action (`ActT`):339                The action taken by the agent.340            observation (`ObsT`):341                The resulting observation.342 343        Returns:344            `float`: Reward value from the rubric, or 0.0 if no rubric is set.345 346        Call this in `step()` to compute and assign the reward:347 348        ```python349        def step(self, action: MyAction, ...) -> MyObservation:350            # ... execute action and create observation ...351            observation.reward = self._apply_rubric(action, observation)352            return observation353        ```354        """355        if self.rubric is not None:356            return self.rubric(action, observation)357        return 0.0358 359    async def _apply_rubric_async(self, action: ActT, observation: ObsT) -> float:360        """Apply rubric asynchronously if one is provided.361 362        Args:363            action (`ActT`):364                The action taken by the agent.365            observation (`ObsT`):366                The resulting observation.367 368        Returns:369            `float`: Reward value from the rubric, or 0.0 if no rubric is set.370 371        Call this in `step_async()` to compute and assign the reward:372 373        ```python374        async def step_async(self, action: MyAction, ...) -> MyObservation:375            # ... execute action and create observation ...376            observation.reward = await self._apply_rubric_async(action, observation)377            return observation378        ```379        """380        if self.rubric is not None:381            result = self.rubric(action, observation)382            # If rubric returns a coroutine, await it383            if inspect.iscoroutine(result):384                return await result385            return result386        return 0.0387 388    def _reset_rubric(self) -> None:389        """Reset the rubric state if one is provided.390 391        Call this in `reset()` to clear any trajectory state in the rubric:392 393        ```python394        def reset(self, ...) -> MyObservation:395            self._reset_rubric()396            # ... create initial observation ...397            return observation398        ```399        """400        if self.rubric is not None:401            self.rubric.reset()402 403    async def _reset_rubric_async(self) -> None:404        """Reset the rubric state asynchronously if one is provided.405 406        Call this in `reset_async()` to clear any trajectory state in the rubric:407 408        ```python409        async def reset_async(self, ...) -> MyObservation:410            await self._reset_rubric_async()411            # ... create initial observation ...412            return observation413        ```414        """415        if self.rubric is not None:416            # Check if rubric has async reset method417            if hasattr(self.rubric, "reset_async"):418                result = self.rubric.reset_async()419                if inspect.iscoroutine(result):420                    await result421            else:422                self.rubric.reset()423 424    def close(self) -> None:425        """Clean up resources used by the environment.426 427        Override this method to implement custom cleanup logic.428        Called when the environment is being destroyed or reset.429        """430        pass431