openenv-testing/android_env-pr-162
0
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 