Team Ai
Apppublic

openenv/coding_env

sourceHugging Faceupdated 3mo agoView on Hugging Face
21likes
interfaces.py298 linesDownload Raw Back to env_server
1# Copyright (c) Meta Platforms, Inc. and affiliates.2# All rights reserved.3#4# This source code is licensed under the BSD-style license found in the5# LICENSE file in the root directory of this source tree.6 7import inspect8from abc import ABC, abstractmethod9from typing import Any, Generic, Optional, Protocol, TYPE_CHECKING, TypeVar10 11from typing_extensions import TypedDict12 13from .types import Action, EnvironmentMetadata, Observation, State14 15if TYPE_CHECKING:16    from openenv.core.rubrics import Rubric17 18ActT = TypeVar("ActT", bound=Action)19ObsT = TypeVar("ObsT", bound=Observation)20StateT = TypeVar("StateT", bound=State)21 22 23class Message(TypedDict):24    """A message in a conversation.25 26    Compatible with Huggingface chat template format.27    """28 29    role: str30    content: str31 32 33class ModelTokenizer(Protocol):34    """Protocol for tokenizers that support chat templates.35 36    This protocol defines the interface that tokenizers must implement37    to work with chat-based environments. It's compatible with38    Huggingface transformers tokenizers.39    """40 41    def apply_chat_template(42        self,43        conversation: list[Message],44        tokenize: bool = True,45        return_tensors: str | None = None,46        **kwargs: Any,47    ) -> Any:48        """Apply a chat template to format and optionally tokenize a conversation.49 50        Args:51            conversation: List of message dictionaries with 'role' and 'content'52            tokenize: Whether to tokenize the output53            return_tensors: Format for returned tensors ('pt' for PyTorch)54            **kwargs: Additional arguments55 56        Returns:57            Formatted and optionally tokenized conversation58        """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: Token IDs to decode68            skip_special_tokens: Whether to skip special tokens in output69            **kwargs: Additional arguments70 71        Returns:72            Decoded text string73        """74        ...75 76 77class Transform(ABC, Generic[ObsT]):78    """Transform observations to add rewards, metrics, or other modifications.79 80    Transforms follow the TorchRL pattern where they take an observation81    and return a (potentially modified) observation. This allows for82    flexible reward computation and observation augmentation.83    """84 85    @abstractmethod86    def __call__(self, observation: ObsT) -> ObsT:87        """Transform an observation.88 89        Args:90            observation: The input observation91 92        Returns:93            The transformed observation94        """95        pass96 97 98class Environment(ABC, Generic[ActT, ObsT, StateT]):99    """Base class for all environment servers following Gym/Gymnasium API.100 101    Args:102        transform: Optional transform to apply to observations103        rubric: Optional rubric for reward computation. When provided, the104            rubric's output can be used to set the observation's reward in step().105 106    Class Attributes:107        SUPPORTS_CONCURRENT_SESSIONS: Whether this environment supports concurrent sessions.108            When True, multiple WebSocket connections can each have their own109            environment instance (up to max_concurrent_envs). When False (default),110            the environment should only be used with a single session at a time.111 112            Set this to True in your Environment subclass if:113            - The environment uses proper session isolation (e.g., unique working dirs)114            - No shared mutable state exists between instances115            - External resources (databases, APIs) can handle concurrent access116 117    Attributes:118        rubric: Optional rubric for computing rewards. Environments can set this119            in __init__ and use it in step() to compute observation rewards.120            Training infrastructure can access it for introspection:121                for name, r in env.rubric.named_rubrics():122                    print(f"{name}: {r.last_score}")123 124    See RFC 004 for rubric design: rfcs/004-rubrics.md125    """126 127    # Class-level flag indicating whether this environment supports concurrent sessions128    SUPPORTS_CONCURRENT_SESSIONS: bool = False129 130    # Optional rubric for reward computation131    rubric: Optional["Rubric"]132 133    def __init__(134        self,135        transform: Optional[Transform[ObsT]] = None,136        rubric: Optional["Rubric"] = None,137    ):138        self.transform = transform139        self.rubric = rubric140 141    @abstractmethod142    def reset(143        self,144        seed: Optional[int] = None,145        episode_id: Optional[str] = None,146        **kwargs: Any,147    ) -> ObsT:148        """Reset the environment and return initial observation."""149        pass150 151    async def reset_async(152        self,153        seed: Optional[int] = None,154        episode_id: Optional[str] = None,155        **kwargs: Any,156    ) -> ObsT:157        """Async version of reset. Default implementation calls sync reset.158 159        Override to provide true async implementation.160        """161        return self.reset(seed=seed, episode_id=episode_id, **kwargs)162 163    @abstractmethod164    def step(165        self,166        action: ActT,167        timeout_s: Optional[float] = None,168        **kwargs: Any,169    ) -> ObsT:170        """Take a step in the environment."""171        pass172 173    async def step_async(174        self,175        action: ActT,176        timeout_s: Optional[float] = None,177        **kwargs: Any,178    ) -> ObsT:179        """Async version of step. Default implementation calls sync step.180 181        Override to provide true async implementation.182        """183        return self.step(action, timeout_s=timeout_s, **kwargs)184 185    @property186    @abstractmethod187    def state(self) -> StateT:188        """Get the current environment state."""189        pass190 191    def get_metadata(self) -> EnvironmentMetadata:192        """193        Get metadata about this environment.194 195        Override this method to provide custom metadata for the environment.196        Default implementation returns basic metadata derived from class name.197 198        Returns:199            EnvironmentMetadata with environment information200        """201        return EnvironmentMetadata(202            name=self.__class__.__name__,203            description=f"{self.__class__.__name__} environment",204            version="1.0.0",205        )206 207    def _apply_transform(self, observation: ObsT) -> ObsT:208        """Apply transform if one is provided."""209        if self.transform is not None:210            return self.transform(observation)211        return observation212 213    def _apply_rubric(self, action: ActT, observation: ObsT) -> float:214        """Apply rubric if one is provided.215 216        Args:217            action: The action taken by the agent.218            observation: The resulting observation.219 220        Returns:221            Reward value from the rubric, or 0.0 if no rubric is set.222 223        Usage in step():224            def step(self, action: MyAction, ...) -> MyObservation:225                # ... execute action and create observation ...226                observation.reward = self._apply_rubric(action, observation)227                return observation228        """229        if self.rubric is not None:230            return self.rubric(action, observation)231        return 0.0232 233    async def _apply_rubric_async(self, action: ActT, observation: ObsT) -> float:234        """Apply rubric asynchronously if one is provided.235 236        Args:237            action: The action taken by the agent.238            observation: The resulting observation.239 240        Returns:241            Reward value from the rubric, or 0.0 if no rubric is set.242 243        Usage in step_async():244            async def step_async(self, action: MyAction, ...) -> MyObservation:245                # ... execute action and create observation ...246                observation.reward = await self._apply_rubric_async(action, observation)247                return observation248        """249        if self.rubric is not None:250            result = self.rubric(action, observation)251            # If rubric returns a coroutine, await it252            if inspect.iscoroutine(result):253                return await result254            return result255        return 0.0256 257    def _reset_rubric(self) -> None:258        """Reset the rubric state if one is provided.259 260        Call this in reset() to clear any trajectory state in the rubric.261 262        Usage in reset():263            def reset(self, ...) -> MyObservation:264                self._reset_rubric()265                # ... create initial observation ...266                return observation267        """268        if self.rubric is not None:269            self.rubric.reset()270 271    async def _reset_rubric_async(self) -> None:272        """Reset the rubric state asynchronously if one is provided.273 274        Call this in reset_async() to clear any trajectory state in the rubric.275 276        Usage in reset_async():277            async def reset_async(self, ...) -> MyObservation:278                await self._reset_rubric_async()279                # ... create initial observation ...280                return observation281        """282        if self.rubric is not None:283            # Check if rubric has async reset method284            if hasattr(self.rubric, "reset_async"):285                result = self.rubric.reset_async()286                if inspect.iscoroutine(result):287                    await result288            else:289                self.rubric.reset()290 291    def close(self) -> None:292        """Clean up resources used by the environment.293 294        Override this method to implement custom cleanup logic.295        Called when the environment is being destroyed or reset.296        """297        pass298