openenv/chat_env
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 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 