Team Ai
Apppublic

openenv/chat_env

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
chat_environment.py212 linesDownload Raw Back to 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 7"""8Chat Environment Implementation.9 10A chat-based environment for LLMs, designed as a blank canvas for conversation and RL.11"""12 13from openenv.core.env_server.interfaces import (14    Environment,15    Message,16    ModelTokenizer,17    Transform,18)19 20# Support both in-repo and standalone imports21try:22    # In-repo imports (when running from OpenEnv repository)23    from ..models import ChatAction, ChatObservation, ChatState24except ImportError as e:25    if "relative import" not in str(e) and "no known parent package" not in str(e):26        raise27    # Standalone imports (when running via uvicorn server.app:app)28    from models import ChatAction, ChatObservation, ChatState29 30 31class ChatEnvironment(Environment):32    """A chat-based environment for LLMs, designed as a blank canvas for conversation and RL.33 34    This environment is designed to work with language models. It provides the fundamental structure35    for managing conversation state but is intentionally minimal to allow maximum flexibility.36 37    The environment owns the tokenizer and is responsible for managing both message history and tokens.38    Actions contain only tokens that interface directly with models.39 40    Args:41        tokenizer: A tokenizer that will be used to tokenize the conversation42        system_prompt: An optional system prompt string to use during reset calls (optional)43        system_role: The role of the system (at reset time). Defaults to "system"44        transform: Optional transform to apply to observations45    """46 47    def __init__(48        self,49        tokenizer: ModelTokenizer,50        system_prompt: str | None = None,51        system_role: str = "system",52        transform: Transform | None = None,53    ):54        super().__init__(transform=transform)55 56        if not hasattr(tokenizer, "apply_chat_template") and not hasattr(57            tokenizer, "encode"58        ):59            raise ValueError(60                "Tokenizer must have 'apply_chat_template' or 'encode' method"61            )62        self.tokenizer = tokenizer63        self.system_prompt = system_prompt64        self.system_role = system_role65 66        self._state = ChatState()67 68        if system_prompt:69            system_message: Message = {"role": system_role, "content": system_prompt}70            self._state.history_messages.append(system_message)71            system_tokens = self._tokenize_conversation([system_message])72            self._state.history_tokens.append(system_tokens)73 74    def _coerce_tokens(self, tokens) -> list[int]:75        """Normalize tokenizer outputs into a flat list of ints."""76        if hasattr(tokens, "tolist") and callable(tokens.tolist):77            tokens = tokens.tolist()78 79        if isinstance(tokens, tuple):80            tokens = list(tokens)81 82        if isinstance(tokens, list):83            flattened: list[int] = []84            for token in tokens:85                flattened.extend(self._coerce_tokens(token))86            return flattened87 88        return [int(tokens)]89 90    def _tokenize_conversation(self, conversation: list[Message]) -> list[int]:91        """Tokenize a conversation with a chat-template fallback for base tokenizers."""92        try:93            tokens = self.tokenizer.apply_chat_template(conversation=conversation, tokenize=True)94        except Exception:95            # Some tokenizers (e.g. gpt2) do not define `chat_template`.96            fallback_text = "".join(97                f"{m['role']}: {m['content']}\n" for m in conversation98            )99            if hasattr(self.tokenizer, "encode"):100                tokens = self.tokenizer.encode(fallback_text)  # type: ignore[attr-defined]101            else:102                raise ValueError("Tokenizer must support apply_chat_template or encode")103 104        return self._coerce_tokens(tokens)105 106    def reset(self) -> ChatObservation:107        """Reset the environment to initial state.108 109        Returns:110            ChatObservation: Initial observation with system prompt (if any)111        """112        self._state.history_messages = []113        self._state.history_tokens = []114        if self.system_prompt:115            system_message: Message = {116                "role": self.system_role,117                "content": self.system_prompt,118            }119            self._state.history_messages = [system_message]120            system_tokens = self._tokenize_conversation([system_message])121            self._state.history_tokens = [system_tokens]122 123        return self._create_observation()124 125    def step(self, action: ChatAction) -> ChatObservation:  # type: ignore[override]126        """Take a step in the environment by adding tokens to the chat history.127 128        Args:129            action: A ChatAction object containing tokens.130 131        Returns:132            ChatObservation: The updated observation with the new tokens added.133        """134        action_tokens = [int(token) for token in action.tokens]135 136        # Store the tokens directly from the action137        self._state.history_tokens.append(action_tokens)138 139        # Decode tokens to text and add as a message to history140        decoded_text = self.tokenizer.decode(action_tokens, skip_special_tokens=True)141        assistant_message: Message = {"role": "assistant", "content": decoded_text}142        self._state.history_messages.append(assistant_message)143 144        return self._create_observation()145 146    def _create_observation(self) -> ChatObservation:147        """Create a ChatObservation from the current state.148 149        Returns both the message history and the tokens flattened as a single tensor150        ready to be used by models.151 152        Returns:153            ChatObservation: Observation with messages and flattened tokens154        """155        if self._state.history_tokens:156            flattened_tokens = [157                token158                for token_list in self._state.history_tokens159                for token in token_list160            ]161        else:162            flattened_tokens = []163 164        observation = ChatObservation(165            messages=self._state.history_messages.copy(),  # Copy to prevent external mutation166            tokens=flattened_tokens,167        )168 169        transformed = self._apply_transform(observation)170        if isinstance(transformed, ChatObservation):171            return transformed172        else:173            # If transform returns base Observation, convert back to ChatObservation174            return ChatObservation(175                messages=getattr(transformed, "messages", []),176                tokens=self._coerce_tokens(getattr(transformed, "tokens", [])),177                done=transformed.done,178                reward=transformed.reward,179            )180 181    @property182    def state(self) -> ChatState:183        """Get the current state of the environment.184 185        Returns:186            ChatState: The current state.187        """188        return self._state189 190    def message_to_action(self, message: Message) -> ChatAction:191        """Convert a message dictionary to a ChatAction with tokens.192 193        Args:194            message: Dictionary with 'role' and 'content' keys195 196        Returns:197            ChatAction: A new ChatAction instance with tokenized content198 199        Raises:200            ValueError: If required keys are missing201        """202        if "role" not in message:203            raise ValueError("Message must contain a 'role' key")204        if "content" not in message:205            raise ValueError("Message must contain a 'content' key")206        if message["content"] is None:207            raise ValueError("Message content cannot be None")208 209        tokens = self._tokenize_conversation([message])210 211        return ChatAction(tokens=tokens)212