Team Ai
Apppublic

openenv-testing/android_env-pr-162

sourceHugging Faceupdated 11mo agoView on Hugging Face
0likes
interfaces.py119 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 7from abc import ABC, abstractmethod8from typing import Any, Protocol, TypedDict9 10from .types import Action, Observation, State11 12 13class Message(TypedDict):14    """A message in a conversation.15 16    Compatible with Huggingface chat template format.17    """18 19    role: str20    content: str21 22 23class ModelTokenizer(Protocol):24    """Protocol for tokenizers that support chat templates.25 26    This protocol defines the interface that tokenizers must implement27    to work with chat-based environments. It's compatible with28    Huggingface transformers tokenizers.29    """30 31    def apply_chat_template(32        self,33        conversation: list[Message],34        tokenize: bool = True,35        return_tensors: str | None = None,36        **kwargs: Any,37    ) -> Any:38        """Apply a chat template to format and optionally tokenize a conversation.39 40        Args:41            conversation: List of message dictionaries with 'role' and 'content'42            tokenize: Whether to tokenize the output43            return_tensors: Format for returned tensors ('pt' for PyTorch)44            **kwargs: Additional arguments45 46        Returns:47            Formatted and optionally tokenized conversation48        """49        ...50 51    def decode(52        self, token_ids: Any, skip_special_tokens: bool = False, **kwargs: Any53    ) -> str:54        """Decode token IDs back to text.55 56        Args:57            token_ids: Token IDs to decode58            skip_special_tokens: Whether to skip special tokens in output59            **kwargs: Additional arguments60 61        Returns:62            Decoded text string63        """64        ...65 66 67class Transform(ABC):68    """Transform observations to add rewards, metrics, or other modifications.69 70    Transforms follow the TorchRL pattern where they take an observation71    and return a (potentially modified) observation. This allows for72    flexible reward computation and observation augmentation.73    """74 75    @abstractmethod76    def __call__(self, observation: Observation) -> Observation:77        """Transform an observation.78 79        Args:80            observation: The input observation81 82        Returns:83            The transformed observation84        """85        pass86 87 88class Environment(ABC):89    """Base class for all environment servers following Gym/Gymnasium API.90 91    Args:92        transform: Optional transform to apply to observations93    """94 95    def __init__(self, transform: Transform | None = None):96        self.transform = transform97 98    @abstractmethod99    def reset(self) -> Observation:100        """Reset the environment and return initial observation."""101        pass102 103    @abstractmethod104    def step(self, action: Action) -> Observation:105        """Take a step in the environment."""106        pass107 108    @property109    @abstractmethod110    def state(self) -> State:111        """Get the current environment state."""112        pass113 114    def _apply_transform(self, observation: Observation) -> Observation:115        """Apply transform if one is provided."""116        if self.transform is not None:117            return self.transform(observation)118        return observation119