Skip to content
Merged
2 changes: 1 addition & 1 deletion environments/alphabet_sort/alphabet_sort.py
Original file line number Diff line number Diff line change
Expand Up @@ -224,7 +224,7 @@ async def env_response(
assistant_count = len([m for m in messages if m["role"] == "assistant"])
follow_ups = state["info"]["follow_ups"]
follow_up_idx = assistant_count - 1
return [{"role": "user", "content": follow_ups[follow_up_idx]}]
return [vf.UserMessage(content=follow_ups[follow_up_idx])]

def score_response(
predicted: List[str], expected: List[str], apply_power: bool = True
Expand Down
2 changes: 1 addition & 1 deletion environments/doublecheck/doublecheck.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ async def env_response(
self, messages: Messages, state: State, **kwargs
) -> Messages:
"""Generate a response from the environment."""
return [{"role": "user", "content": "Are you sure?"}]
return [vf.UserMessage(content="Are you sure?")]


def load_environment(
Expand Down
7 changes: 1 addition & 6 deletions environments/sentence_repeater/sentence_repeater.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,12 +83,7 @@ async def env_response(
self, messages: Messages, state: State, **kwargs
) -> Messages:
num_turn = len(state["trajectory"])
return [
{
"role": "user",
"content": state["info"]["questions"][num_turn],
}
]
return [vf.UserMessage(content=state["info"]["questions"][num_turn])]


def load_environment(**kwargs) -> vf.Environment:
Expand Down
6 changes: 2 additions & 4 deletions verifiers/envs/environment.py
Original file line number Diff line number Diff line change
Expand Up @@ -494,7 +494,7 @@ def get_state_usage(self, state: State) -> TokenUsage | None:
async def get_model_response(
self,
state: State,
prompt: Messages | str,
prompt: Messages,
client: Client | None = None,
model: str | None = None,
tool_defs: list[Tool] | None = None,
Expand Down Expand Up @@ -539,12 +539,10 @@ def resolve_optional_args(
client, model, tool_defs, sampling_args
)

normalized_prompt = normalize_messages(prompt, field_name="prompt")
Comment thread
mikasenghaas marked this conversation as resolved.

self._get_usage_tracker(state, create_if_missing=True)

response = await client.get_response(
prompt=normalized_prompt,
prompt=prompt,
model=model,
tools=tool_defs,
sampling_args=sampling_args,
Expand Down
2 changes: 1 addition & 1 deletion verifiers/envs/experimental/cli_agent_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -436,7 +436,7 @@ async def get_prompt_messages(self, state: State) -> Messages:
async def get_model_response(
self,
state: State,
prompt: Messages | str,
prompt: Messages,
client: Client | None = None,
model: str | None = None,
tool_defs: list[Tool] | None = None,
Expand Down
6 changes: 3 additions & 3 deletions verifiers/envs/experimental/gym_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,12 +142,12 @@ def obs_to_text(self, obs: Any) -> str:
return self.obs_to_text_fn(obs)
return str(obs)

def wrap_response(self, text: str) -> vf.Messages | str:
return cast(vf.Messages, [{"role": "user", "content": text}])
def wrap_response(self, text: str) -> vf.Messages:
return [vf.UserMessage(content=text)]

async def env_response(
self, messages: vf.Messages, state: State, **kwargs: Any
) -> vf.Messages | str:
) -> vf.Messages:
if "gym_env" not in state:
env = self.env_cls(**self.env_kwargs)
seed = int(state["answer"])
Expand Down
47 changes: 20 additions & 27 deletions verifiers/envs/multiturn_env.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,10 @@
State,
TrajectoryStep,
)
from verifiers.utils.message_utils import concat_messages, normalize_messages
from verifiers.utils.message_utils import (
concat_messages,
maybe_normalize_messages,
)
from verifiers.utils.response_utils import (
parse_response_message,
parse_response_tokens,
Expand Down Expand Up @@ -41,7 +44,7 @@ def __init__(self, max_turns: int = -1, **kwargs):
@abstractmethod
async def env_response(
self, messages: Messages, state: State, **kwargs
) -> Messages | str:
) -> Messages:
"""
Generate a response from the environment.
"""
Expand Down Expand Up @@ -71,42 +74,29 @@ async def setup_state(self, state: State) -> State:
async def get_prompt_messages(self, state: State) -> Messages:
"""Override for rollouts with non-linear message sequences."""
if len(state["trajectory"]) == 0:
return normalize_messages(state["prompt"], field_name="state.prompt")
prev_turn_prompt = normalize_messages(
state["trajectory"][-1]["prompt"], field_name="trajectory.prompt"
)
prev_turn_completion = normalize_messages(
state["trajectory"][-1]["completion"], field_name="trajectory.completion"
)
return state["prompt"]
prev_turn_prompt = state["trajectory"][-1]["prompt"]
prev_turn_completion = state["trajectory"][-1]["completion"]
messages = concat_messages([prev_turn_prompt, prev_turn_completion])
env_response = await self.env_response(messages, state)
env_response_messages = normalize_messages(
env_response, field_name="env_response"
)
return concat_messages([messages, env_response_messages])
env_response = maybe_normalize_messages(env_response, field_name="env_response")
return concat_messages([messages, env_response])

async def render_completion(self, state: State):
"""Override for rollouts with non-linear message sequences."""
if len(state["trajectory"]) == 0:
state["completion"] = []
return
last_prompt = normalize_messages(
state["trajectory"][-1]["prompt"], field_name="trajectory.prompt"
)
last_completion = normalize_messages(
state["trajectory"][-1]["completion"], field_name="trajectory.completion"
)
last_prompt = state["trajectory"][-1]["prompt"]
last_completion = state["trajectory"][-1]["completion"]
full_conversation = concat_messages([last_prompt, last_completion])
if state.get("final_env_response"):
full_conversation = concat_messages(
[
full_conversation,
normalize_messages(
state["final_env_response"], field_name="final_env_response"
),
]
final_resp = state["final_env_response"]
final_resp = maybe_normalize_messages(
final_resp, field_name="final_env_response"
)
prompt_messages = normalize_messages(state["prompt"], field_name="state.prompt")
full_conversation = concat_messages([full_conversation, final_resp])
prompt_messages = state["prompt"]
state["completion"] = full_conversation[len(prompt_messages) :]

async def add_trajectory_step(self, state: State, trajectory_step: TrajectoryStep):
Expand Down Expand Up @@ -156,6 +146,9 @@ async def rollout(
while not await self.is_completed(state):
try:
prompt_messages = await self.get_prompt_messages(state)
prompt_messages = maybe_normalize_messages(
prompt_messages, field_name="prompt_messages"
)
if state.get("final_env_response") is not None:
continue
response = await self.get_model_response(state, prompt_messages)
Expand Down
19 changes: 18 additions & 1 deletion verifiers/utils/logging_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,25 @@
from verifiers.errors import Error
from verifiers.types import ErrorInfo, Messages
from verifiers.utils.error_utils import ErrorChain
from verifiers.utils.message_utils import format_messages

LOGGER_NAME = "verifiers"

_seen_once_keys: set[tuple[str, str]] = set()


def log_once(logger: logging.Logger, level: int, msg: str) -> None:
"""Log a message only once per (logger name, message) pair for the process lifetime."""
key = (logger.name, msg)
Comment thread
cursor[bot] marked this conversation as resolved.
if key in _seen_once_keys:
return
_seen_once_keys.add(key)
logger.log(level, msg)


def warning_once(logger: logging.Logger, msg: str) -> None:
"""Shorthand for ``log_once(logger, logging.WARNING, ...)``."""
log_once(logger, logging.WARNING, msg)


class JsonFormatter(logging.Formatter):
"""JSON formatter for structured logging."""
Expand Down Expand Up @@ -144,6 +159,8 @@ def print_prompt_completions_sample(
step: int,
num_samples: int = 1,
) -> None:
from verifiers.utils.message_utils import format_messages

def format_error(error: ErrorInfo | BaseException) -> Text:
out = Text()
if isinstance(error, BaseException):
Expand Down
26 changes: 26 additions & 0 deletions verifiers/utils/message_utils.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import json
import logging
import re
from collections.abc import Mapping
from typing import Any, cast
Expand All @@ -19,6 +20,8 @@
UserMessage,
)

logger = logging.getLogger(__name__)


def from_raw_content_part(part: dict[str, Any]) -> ContentPart:
"""Convert a raw content-part dict to a typed content part when possible."""
Expand Down Expand Up @@ -139,6 +142,29 @@ def normalize_messages(
return normalized


def maybe_normalize_messages(
value: Messages | str,
*,
field_name: str = "messages",
) -> Messages:
"""Normalize messages only if needed, logging a warning on first occurrence."""
from verifiers.utils.logging_utils import warning_once

requires_normalize = not isinstance(value, list) or not all(
isinstance(m, Message) for m in value
)
if not requires_normalize:
return cast(Messages, value)
warning_once(
logger,
f"{field_name} returned raw dicts/strings instead of vf.Messages. This"
" repeatedly triggers normalize_messages(), causing unnecessary"
" Pydantic validation overhead. Return vf.Message types (e.g."
" vf.UserMessage, vf.AssistantMessage) to avoid this.",
)
return normalize_messages(value, field_name=field_name)
Comment thread
mikasenghaas marked this conversation as resolved.


def concat_messages(messages_list: list[Messages]) -> Messages:
"""Concatenate multiple Messages lists into one."""
result = []
Expand Down
14 changes: 6 additions & 8 deletions verifiers/utils/response_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,14 +9,12 @@
async def parse_response_message(response: Response) -> Messages:
"""Parse a vf.Response into a vf.Messages list (single vf.AssistantMessage)."""
response_message = response.message
message_payload = {
"role": "assistant",
"content": response_message.content,
"reasoning_content": response_message.reasoning_content,
"thinking_blocks": response_message.thinking_blocks,
"tool_calls": response_message.tool_calls,
}
message = AssistantMessage.model_validate(message_payload)
message = AssistantMessage(
content=response_message.content,
reasoning_content=response_message.reasoning_content,
thinking_blocks=response_message.thinking_blocks,
tool_calls=response_message.tool_calls,
)
return [message]


Expand Down
Loading