diff --git a/tests/test_client_multimodal_types.py b/tests/test_client_multimodal_types.py index 934cccb2f0..a8fcf80995 100644 --- a/tests/test_client_multimodal_types.py +++ b/tests/test_client_multimodal_types.py @@ -296,3 +296,54 @@ async def test_anthropic_tool_call_round_trips_thinking_blocks(): {"type": "thinking", "thinking": "hidden chain", "signature": "sig_1"}, {"type": "tool_use", "id": "call_1", "name": "lookup", "input": {"q": "x"}}, ] + + +def test_prepare_images_inplace_offloads_every_image_part_shape(tmp_path): + """Ingress offload must cover the full renderer part treaty: nested + ``image_url`` dicts, direct-string ``image_url``, direct ``image`` + strings, and typed pydantic parts — and reject non-file leftovers.""" + import base64 + + pytest.importorskip( + "renderers.mm_store", reason="needs a renderers version with raw image offload" + ) + from verifiers.utils.multimodal import prepare_images_inplace + + raw = b"png-ish bytes" + data_url = "data:image/png;base64," + base64.b64encode(raw).decode("ascii") + wire = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": data_url}}, + {"type": "image_url", "image_url": data_url}, + {"type": "image", "image": data_url}, + ], + } + ] + } + typed = UserMessage( + content=[ImageUrlContentPart(image_url=ImageUrlSource(url=data_url))] + ) + + prepare_images_inplace(wire, image_dir=tmp_path) + prepare_images_inplace(typed, image_dir=tmp_path) + + parts = wire["messages"][0]["content"] + offloaded = [ + parts[0]["image_url"]["url"], + parts[1]["image_url"], + parts[2]["image"], + typed.content[0].image_url.url, + ] + assert len(set(offloaded)) == 1 # content-addressed: same bytes, same file + assert offloaded[0].startswith("file://") + written = list(tmp_path.iterdir()) + assert len(written) == 1 and written[0].read_bytes() == raw + + with pytest.raises(RuntimeError, match="file://"): + prepare_images_inplace( + {"type": "image", "image": "https://example.com/x.png"}, + image_dir=tmp_path, + ) diff --git a/uv.lock b/uv.lock index 38298ea398..664381fc07 100644 --- a/uv.lock +++ b/uv.lock @@ -1,5 +1,5 @@ version = 1 -revision = 4 +revision = 3 requires-python = ">=3.11, <3.14" resolution-markers = [ "python_full_version >= '3.13' and sys_platform == 'win32'", @@ -163,6 +163,12 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [ + { name = "datasets" }, + { name = "verifiers" }, +] + [[package]] name = "annotated-doc" version = "0.0.4" @@ -692,6 +698,9 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [{ name = "verifiers" }] + [[package]] name = "color-codeword-v1" version = "0.1.0" @@ -701,6 +710,12 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [ + { name = "pillow", specifier = ">=11.0.0" }, + { name = "verifiers" }, +] + [[package]] name = "colorama" version = "0.4.6" @@ -718,6 +733,9 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [{ name = "verifiers" }] + [[package]] name = "connect-python" version = "0.9.0" @@ -952,6 +970,9 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [{ name = "verifiers" }] + [[package]] name = "defusedxml" version = "0.7.1" @@ -1355,6 +1376,9 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [{ name = "verifiers" }] + [[package]] name = "googleapis-common-protos" version = "1.75.0" @@ -1503,6 +1527,9 @@ dependencies = [ { name = "datasets" }, ] +[package.metadata] +requires-dist = [{ name = "datasets" }] + [[package]] name = "h11" version = "0.16.0" @@ -2061,6 +2088,9 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [{ name = "verifiers" }] + [[package]] name = "latex2sympy2-extended" version = "1.11.0" @@ -2810,6 +2840,9 @@ dependencies = [ { name = "verifiers", extra = ["openenv"] }, ] +[package.metadata] +requires-dist = [{ name = "verifiers", extras = ["openenv"] }] + [[package]] name = "opentelemetry-api" version = "1.43.0" @@ -3250,6 +3283,9 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [{ name = "verifiers" }] + [[package]] name = "protobuf" version = "6.33.6" @@ -4052,6 +4088,9 @@ dependencies = [ { name = "datasets" }, ] +[package.metadata] +requires-dist = [{ name = "datasets" }] + [[package]] name = "rich" version = "15.0.0" @@ -4225,6 +4264,9 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [{ name = "verifiers" }] + [[package]] name = "secretstorage" version = "3.5.0" @@ -4918,6 +4960,79 @@ examples = [ { name = "wordle-v1" }, ] +[package.metadata] +requires-dist = [ + { name = "aiohttp", specifier = ">=3.9.0" }, + { name = "aiolimiter", specifier = ">=1.2.1" }, + { name = "anthropic", specifier = ">=0.78.0" }, + { name = "datasets", specifier = ">=3.3.0,<6.0.0" }, + { name = "gepa", specifier = ">=0.0.6" }, + { name = "harbor", marker = "python_full_version >= '3.12' and extra == 'harbor'", specifier = "==0.20.0" }, + { name = "httpx", specifier = ">=0.27.0" }, + { name = "loguru", specifier = ">=0.7.0" }, + { name = "math-verify", specifier = ">=0.8.0" }, + { name = "mcp", specifier = ">=1.24.0,<2" }, + { name = "modal", marker = "extra == 'modal'", specifier = ">=1.4.0" }, + { name = "msgpack", specifier = ">=1.1.2" }, + { name = "nest-asyncio", marker = "extra == 'notebook'", specifier = ">=1.6.0" }, + { name = "nltk", marker = "extra == 'ta'", specifier = ">=3.9.2" }, + { name = "numpy", specifier = ">=2.1.0" }, + { name = "openai", specifier = ">=2.9.0" }, + { name = "openai-agents", specifier = ">=0.8.2" }, + { name = "openenv", marker = "extra == 'openenv'", specifier = ">=0.4.1" }, + { name = "prime-pydantic-config", extras = ["toml"], specifier = ">=0.4.2" }, + { name = "prime-sandboxes", specifier = ">=0.2.33" }, + { name = "prime-tunnel", specifier = ">=0.1.8" }, + { name = "pydantic", specifier = ">=2.12.3" }, + { name = "python-dotenv", marker = "extra == 'browser'", specifier = ">=1.0.0" }, + { name = "pyzmq", specifier = ">=27.1.0" }, + { name = "reasoning-gym", marker = "extra == 'rg'", specifier = ">=0.1.8" }, + { name = "renderers", specifier = ">=0.1.9.dev9" }, + { name = "requests" }, + { name = "rich", specifier = ">=11.0.0" }, + { name = "setproctitle", specifier = ">=1.3.0" }, + { name = "stagehand", marker = "extra == 'browser'", specifier = ">=3.4.0" }, + { name = "tenacity", specifier = ">=8.5.0" }, + { name = "textarena", marker = "extra == 'ta'", specifier = "==0.7.4" }, + { name = "tomli-w", specifier = ">=1.0.0" }, + { name = "typing-extensions", specifier = ">=4.12.2" }, + { name = "uvloop", marker = "platform_python_implementation != 'PyPy' and sys_platform != 'cygwin' and sys_platform != 'win32'", specifier = ">=0.21.0" }, +] +provides-extras = ["browser", "harbor", "modal", "notebook", "openenv", "rg", "ta"] + +[package.metadata.requires-dev] +dev = [ + { name = "nest-asyncio", specifier = ">=1.6.0" }, + { name = "nltk", specifier = ">=3.9.2" }, + { name = "pre-commit", specifier = ">=2.19.0" }, + { name = "pytest", specifier = ">=7.0.0,<9.1.0" }, + { name = "pytest-asyncio", specifier = ">=0.21.0" }, + { name = "pytest-cov", specifier = ">=4.0.0" }, + { name = "pytest-xdist", specifier = ">=3.8.0" }, + { name = "python-dotenv", specifier = ">=1.0.0" }, + { name = "reasoning-gym", specifier = ">=0.1.8" }, + { name = "ruff", specifier = ">=0.16.0" }, + { name = "stagehand", specifier = ">=3.4.0" }, + { name = "textarena", specifier = "==0.7.4" }, + { name = "ty", specifier = ">=0.0.63" }, +] +examples = [ + { name = "alphabet-sort-v1", editable = "environments/alphabet_sort_v1" }, + { name = "code-golf-v1", editable = "environments/code_golf_v1" }, + { name = "color-codeword-v1", editable = "environments/color_codeword_v1" }, + { name = "compact", editable = "environments/compact" }, + { name = "deepwiki-v1", editable = "environments/deepwiki_v1" }, + { name = "glossary-v1", editable = "environments/glossary_v1" }, + { name = "gsm8k-v1", editable = "environments/gsm8k_v1" }, + { name = "kuhn-poker-v1", editable = "environments/kuhn_poker_v1" }, + { name = "openenv-wordle-v1", editable = "environments/openenv_wordle_v1" }, + { name = "proposer-solver-v1", editable = "environments/proposer_solver_v1" }, + { name = "reverse-text-v1", editable = "environments/reverse_text_v1" }, + { name = "scratchpad-v1", editable = "environments/scratchpad_v1" }, + { name = "wiki-search-v1", editable = "environments/wiki_search_v1" }, + { name = "wordle-v1", editable = "environments/wordle_v1" }, +] + [[package]] name = "virtualenv" version = "21.6.1" @@ -5062,6 +5177,13 @@ dependencies = [ { name = "verifiers" }, ] +[package.metadata] +requires-dist = [ + { name = "chromadb", specifier = ">=0.4.13" }, + { name = "datasets" }, + { name = "verifiers" }, +] + [[package]] name = "win32-setctime" version = "1.2.0" @@ -5079,6 +5201,9 @@ dependencies = [ { name = "verifiers", extra = ["ta"] }, ] +[package.metadata] +requires-dist = [{ name = "verifiers", extras = ["ta"] }] + [[package]] name = "xxhash" version = "3.8.1" diff --git a/verifiers/clients/renderer_client.py b/verifiers/clients/renderer_client.py index bfd916f59c..855fb255e8 100644 --- a/verifiers/clients/renderer_client.py +++ b/verifiers/clients/renderer_client.py @@ -10,6 +10,7 @@ concurrent rollouts tokenize in parallel instead of blocking the event loop. """ +import asyncio import json import threading from collections.abc import Mapping @@ -58,6 +59,7 @@ UserMessage, ) from verifiers.utils.client_utils import setup_openai_client +from verifiers.utils.multimodal import prepare_images_inplace # Module-level bridge counters. Incremented by every RendererClient instance # that tries to stitch a multi-turn prompt; callers (e.g. prime-rl's @@ -474,6 +476,7 @@ def _get_renderer_or_pool( async def to_native_prompt( self, messages: Messages ) -> tuple[list[RendererMessage], dict]: + await asyncio.to_thread(prepare_images_inplace, messages) return ( _attach_tool_call_names([_to_renderer_message(m) for m in messages]), {}, diff --git a/verifiers/types.py b/verifiers/types.py index 87aa10a7b2..a0bbbc68d4 100644 --- a/verifiers/types.py +++ b/verifiers/types.py @@ -216,7 +216,7 @@ class ResponseTokens(CustomBaseModel): completion_logprobs: list[float] routed_experts: RoutedExpertsPayload | None = None # Renderer-emitted multimodal sidecar (renderers.base.MultiModalData) - # carrying processed pixel_values / placeholder ranges per modality. + # carrying raw image descriptors / placeholder ranges per modality. # Populated by the renderer client when the rollout went through a # multimodal-aware renderer; ``None`` otherwise. Stored as ``Any`` to # avoid a hard import dependency on ``renderers`` at this layer. @@ -263,7 +263,7 @@ class TrajectoryStepTokens(TypedDict): is_truncated: bool routed_experts: RoutedExpertsPayload | None # Renderer-emitted multimodal sidecar (renderers.base.MultiModalData) - # carrying processed pixel_values / placeholder ranges per modality. + # carrying raw image descriptors / placeholder ranges per modality. # ``NotRequired`` because text-only rollouts (and non-renderer client # types) never populate it. multi_modal_data: NotRequired[Any] diff --git a/verifiers/utils/multimodal.py b/verifiers/utils/multimodal.py new file mode 100644 index 0000000000..d3065ce31b --- /dev/null +++ b/verifiers/utils/multimodal.py @@ -0,0 +1,103 @@ +"""Multimodal ingress helpers for renderer-backed training.""" + +from __future__ import annotations + +from importlib import import_module +from pathlib import Path +from typing import Any + + +def _offload_image_url(url: object, image_dir: Path | None) -> str | None: + try: + offload_image_to_run_assets = import_module( + "renderers.mm_store" + ).offload_image_to_run_assets + except ( + ImportError, + AttributeError, + ) as exc: # pragma: no cover - dependency-version guard + raise RuntimeError( + "Multimodal training requires a renderers version with raw image asset offload support." + ) from exc + + return offload_image_to_run_assets(url, image_dir=image_dir) + + +def _part_image_field(part_type: object) -> str | None: + """The field carrying a content part's image source, keyed by ``type``. + + Mirrors the renderer-side part treaty (``renderers.qwen3_vl._image_source``): + ``image_url`` parts nest the URL under ``image_url`` (``{"url": ...}`` or a + direct string), ``image`` parts carry the URL string directly. + """ + if part_type in ("image_url", "image"): + return str(part_type) + return None + + +def _get_field(container: Any, name: str) -> object: + if isinstance(container, dict): + return container.get(name) + return getattr(container, name, None) + + +def _set_field(container: Any, name: str, value: str) -> None: + if isinstance(container, dict): + container[name] = value + else: + setattr(container, name, value) + + +def _prepare_image_part(part: Any, field: str, *, image_dir: Path | None) -> None: + """Offload one image part's source to run assets and require ``file://``.""" + source = _get_field(part, field) + if source is None: + return + if isinstance(source, str): + container, key = part, field + elif field == "image_url": + container, key = source, "url" + else: + raise RuntimeError( + f"multimodal training requires string image sources; got {type(source).__name__} under {field!r}" + ) + + url = _get_field(container, key) + offloaded = _offload_image_url(url, image_dir) + if offloaded is not None: + _set_field(container, key, offloaded) + url = offloaded + if not isinstance(url, str) or not url.startswith("file://"): + raise RuntimeError( + "multimodal training requires image sources offloaded to file:// " + f"run image assets; got {url!r} under {field!r}" + ) + + +def prepare_images_inplace(value: Any, *, image_dir: Path | None = None) -> None: + """Offload image URLs reachable from ``value`` to run image assets. + + Handles OpenAI wire dicts/lists and the pydantic v0/v1 message/content-part + models used by trajectories and traces. + """ + if isinstance(value, dict): + field = _part_image_field(value.get("type")) + if field is not None: + _prepare_image_part(value, field, image_dir=image_dir) + for child in value.values(): + prepare_images_inplace(child, image_dir=image_dir) + return + + if isinstance(value, (list, tuple)): + for child in value: + prepare_images_inplace(child, image_dir=image_dir) + return + + field = _part_image_field(getattr(value, "type", None)) + if field is not None: + _prepare_image_part(value, field, image_dir=image_dir) + return + + content = getattr(value, "content", None) + if isinstance(content, (list, tuple)): + prepare_images_inplace(content, image_dir=image_dir) diff --git a/verifiers/v1/clients/client.py b/verifiers/v1/clients/client.py index 5940f49f26..392679ef5f 100644 --- a/verifiers/v1/clients/client.py +++ b/verifiers/v1/clients/client.py @@ -29,6 +29,15 @@ class RelayReply: class Client(ABC): + async def prepare_request_body(self, dialect: Dialect, body: dict) -> dict: + """Normalize a provider request before the interception server parses/traces it. + + Relay clients keep the request verbatim. Training clients may rewrite heavy + in-process payloads (for example base64 images) into stable run-asset refs so the + trace, renderer, and trainer all see the same cheap message content. + """ + return body + @abstractmethod async def get_response( self, diff --git a/verifiers/v1/clients/train.py b/verifiers/v1/clients/train.py index 785f05a923..2874f56e7b 100644 --- a/verifiers/v1/clients/train.py +++ b/verifiers/v1/clients/train.py @@ -8,6 +8,7 @@ needs a running vLLM engine. """ +import asyncio import json from collections.abc import Mapping from typing import Any @@ -15,7 +16,9 @@ from openai import AsyncOpenAI, OpenAIError from renderers import OverlongPromptError as RendererOverlongPromptError from renderers import RenderedTokens, RendererConfig +from renderers.base import is_multimodal +from verifiers.utils.multimodal import prepare_images_inplace from verifiers.v1.clients.client import SESSION_ID_HEADER, Client from verifiers.v1.dialects import FINISH_REASONS, ChatDialect, Dialect, parse_tools from verifiers.v1.dialects.chat import message_to_wire @@ -169,16 +172,6 @@ def _is_valid_incremental_tail(messages: list[dict[str, Any]]) -> bool: return all(role == "tool" for role in roles) -def _has_multimodal_content(messages) -> bool: - for message in messages: - content = getattr(message, "content", None) - if not isinstance(content, list): - continue - if any(getattr(part, "type", None) == "image_url" for part in content): - return True - return False - - class TrainClient(Client): """Renders prompts to token ids and calls a vLLM `/inference/v1/generate` engine.""" @@ -215,6 +208,11 @@ def _renderer_pool( ) return self._pool + async def prepare_request_body(self, dialect: Dialect, body: dict) -> dict: + if isinstance(dialect, ChatDialect): + await asyncio.to_thread(prepare_images_inplace, body) + return body + async def get_response( self, dialect: Dialect, @@ -265,23 +263,24 @@ async def get_response( ) bridged_turn: PendingTurn | None = None - # Only build the (O(context)) previous-turn token ids once the cheap guards pass — a - # multimodal prompt or a tail that isn't a clean `[tool*, user?]` extension can't bridge. - can_bridge = ( - turn is not None - and not _has_multimodal_content(prompt) - and _is_valid_incremental_tail(wire_messages) - ) + # Only build the (O(context)) previous-turn token ids once the cheap guards pass: a + # tail that isn't a clean `[tool*, user?]` extension can't bridge. + can_bridge = turn is not None and _is_valid_incremental_tail(wire_messages) previous_ids = turn.previous_token_ids() if can_bridge else None if previous_ids is not None: previous_prompt_ids, previous_completion_ids = previous_ids def bridge(): + kwargs: dict[str, Any] = {"tools": wire_tools} + if is_multimodal(renderer): + kwargs["previous_multi_modal_data"] = ( + turn.previous_multi_modal_data() + ) return renderer.bridge_to_next_turn( previous_prompt_ids, previous_completion_ids, wire_messages, - tools=wire_tools, + **kwargs, ) bridged = await _maybe_offload(renderer, bridge) diff --git a/verifiers/v1/graph.py b/verifiers/v1/graph.py index f81480d43b..2d3811dd58 100644 --- a/verifiers/v1/graph.py +++ b/verifiers/v1/graph.py @@ -31,12 +31,14 @@ from verifiers.v1.types import ( AssistantMessage, + FinishReason, KeptTokens, Message, Response, TextContentPart, Tool, ToolMessage, + Usage, ) if TYPE_CHECKING: @@ -61,6 +63,46 @@ def _decode_ndarray(d: dict) -> np.ndarray: return np.frombuffer(d["data"], dtype=np.dtype(d["dtype"])).reshape(d["shape"]) +_PROCESSED_MM_KEYS = frozenset({"pixel_values", "image_embeds", "image_features"}) + + +def _contains_processed_mm_key(value: Any) -> bool: + if isinstance(value, dict): + return bool(_PROCESSED_MM_KEYS.intersection(value)) or any( + _contains_processed_mm_key(v) for v in value.values() + ) + if isinstance(value, (list, tuple)): + return any(_contains_processed_mm_key(v) for v in value) + return False + + +def _validate_raw_mm_item(item: Any) -> dict[str, Any]: + if not isinstance(item, dict): + raise TypeError( + "v1 multimodal sidecars must be raw image descriptor dicts, " + f"got {type(item).__name__}" + ) + if _contains_processed_mm_key(item): + raise TypeError( + "v1 multimodal sidecars must be raw image descriptors, " + "not processed multimodal payloads" + ) + if not isinstance(item.get("raw_image_uri"), str) or not item["raw_image_uri"]: + raise ValueError("v1 multimodal sidecars require raw_image_uri") + return dict(item) + + +def _validate_raw_mm_data(mmd: MultiModalData) -> MultiModalData: + return MultiModalData( + mm_hashes={k: list(v) for k, v in mmd.mm_hashes.items()}, + mm_placeholders={k: list(v) for k, v in mmd.mm_placeholders.items()}, + mm_items={ + modality: [_validate_raw_mm_item(item) for item in items] + for modality, items in mmd.mm_items.items() + }, + ) + + class MessageNode(BaseModel): """One message in the graph: a message plus the tokens it adds to the cumulative sequence. Concatenating a root→leaf path's nodes reconstructs that branch's full token @@ -101,12 +143,19 @@ class MessageNode(BaseModel): logprobs: list[float] = Field(default_factory=list) """Sampling logprobs for the sampled tokens — length equals the number of True entries in `mask`; empty for input messages.""" + finish_reason: FinishReason = None + """The response's finish reason (assistant nodes only) — kept for truncation detection.""" multi_modal_data: SkipJsonSchema[MultiModalData | None] = None - """The renderer items for the images this message's content introduces (pixel tensors, - grids, hashes, placeholders) — the only carrier of the pixels from the env server to the - trainer. `Branch.multi_modal_data` concatenates them along the path into the training - `mm_kwargs`. Rides the wire as raw bytes (msgpack `bin`) since pydantic can't JSON the numpy; - kept off disk by the dump-site `exclude` in prime-rl (the tensors bloat the rollout jsonl).""" + """The renderer items for images this message introduces. + + With the raw-image path, items are lightweight descriptors (hashes, grid metadata, and + optional run-image refs), not image processor tensors. `Branch.multi_modal_data` concatenates + them along the path for the trainer. Old processed-payload sidecars are rejected. + """ + usage: Usage | None = None + """Provider-reported token usage for this message's response (assistant nodes). Preserved on + the wire and on disk so dashboards can show token counts and cost even when the endpoint + returns no token ids.""" routed_experts: SkipJsonSchema[np.ndarray | None] = None """This node's slice of the MoE expert-routing array — uint8 `[len(token_ids), layers, top_k]`, the expert ids inference selected for exactly this node's tokens. Attributed from @@ -124,10 +173,10 @@ class MessageNode(BaseModel): @field_serializer("multi_modal_data") def serialize_multi_modal_data(self, mmd: MultiModalData | None) -> dict | None: - """`MultiModalData` -> msgpack-safe dict so the pixel tensors ride the wire; numpy - `mm_items` values become raw-bytes `__nd__` dicts (every renderer emits `return_tensors="np"`).""" + """`MultiModalData` -> msgpack-safe raw descriptor dict.""" if mmd is None: return None + mmd = _validate_raw_mm_data(mmd) return { "mm_hashes": {k: list(v) for k, v in mmd.mm_hashes.items()}, "mm_placeholders": { @@ -135,9 +184,7 @@ def serialize_multi_modal_data(self, mmd: MultiModalData | None) -> dict | None: for modality, ranges in mmd.mm_placeholders.items() }, "mm_items": { - modality: [ - {k: _encode_ndarray(v) for k, v in item.items()} for item in items - ] + modality: [dict(item) for item in items] for modality, items in mmd.mm_items.items() }, } @@ -145,25 +192,29 @@ def serialize_multi_modal_data(self, mmd: MultiModalData | None) -> dict | None: @field_validator("multi_modal_data", mode="before") @classmethod def deserialize_multi_modal_data(cls, value: Any) -> MultiModalData | None: - if value is None or isinstance(value, MultiModalData): + if value is None: return value + if isinstance(value, MultiModalData): + return _validate_raw_mm_data(value) if not isinstance(value, dict): raise TypeError(f"cannot build MultiModalData from {type(value).__name__}") - return MultiModalData( - mm_hashes={k: list(v) for k, v in (value.get("mm_hashes") or {}).items()}, - mm_placeholders={ - modality: [ - PlaceholderRange(offset=p["offset"], length=p["length"]) - for p in ranges - ] - for modality, ranges in (value.get("mm_placeholders") or {}).items() - }, - mm_items={ - modality: [ - {k: _decode_ndarray(v) for k, v in item.items()} for item in items - ] - for modality, items in (value.get("mm_items") or {}).items() - }, + return _validate_raw_mm_data( + MultiModalData( + mm_hashes={ + k: list(v) for k, v in (value.get("mm_hashes") or {}).items() + }, + mm_placeholders={ + modality: [ + PlaceholderRange(offset=p["offset"], length=p["length"]) + for p in ranges + ] + for modality, ranges in (value.get("mm_placeholders") or {}).items() + }, + mm_items={ + modality: list(items) + for modality, items in (value.get("mm_items") or {}).items() + }, + ) ) @field_serializer("routed_experts") @@ -363,6 +414,23 @@ def prompt_message_spans( for span in tail_spans ] + def previous_multi_modal_data(self) -> MultiModalData | None: + """Concatenate multimodal sidecars attached to the reusable prefix.""" + merged = MultiModalData() + found = False + for nid in self.prefix_node_ids: + mmd = self.trace.nodes[nid].multi_modal_data + if mmd is None or mmd.is_empty(): + continue + found = True + for modality, items in mmd.mm_items.items(): + merged.mm_items.setdefault(modality, []).extend(items) + for modality, hashes in mmd.mm_hashes.items(): + merged.mm_hashes.setdefault(modality, []).extend(hashes) + for modality, placeholders in mmd.mm_placeholders.items(): + merged.mm_placeholders.setdefault(modality, []).extend(placeholders) + return merged if found else None + def commit(self, response: Response, tools: list[Tool] | None = None) -> int: """Add this turn to the graph; returns the committed assistant node's id.""" assistant_id = _commit_turn(self, response) @@ -423,8 +491,9 @@ def _attribute_mm( renderer emits items per modality in prompt order (message order, then content-part order), so we walk the path advancing a per-modality cursor over every message's media but write only the nodes created this turn — `path[:num_reused]` is the reused prefix, already - attributed when first created. Item order is all training needs; placeholder offsets aren't - carried.""" + attributed when first created. Each node gets the hashes/items/placeholders for exactly the + media it introduced, preserving vLLM multimodal-list alignment when those node sidecars are + later merged for bridge or training.""" if mmd is None or mmd.is_empty(): return cursors: dict[str, int] = {} @@ -434,6 +503,7 @@ def _attribute_mm( continue node_items: dict[str, list] = {} node_hashes: dict[str, list] = {} + node_placeholders: dict[str, list[PlaceholderRange]] = {} for part in content: modality = _part_modality(part) if modality is None: @@ -445,13 +515,20 @@ def _attribute_mm( continue items = mmd.mm_items.get(modality) or [] hashes = mmd.mm_hashes.get(modality) or [] + placeholders = mmd.mm_placeholders.get(modality) or [] if k < len(items): node_items.setdefault(modality, []).append(items[k]) if k < len(hashes): node_hashes.setdefault(modality, []).append(hashes[k]) - if node_items: - trace.nodes[node_id].multi_modal_data = MultiModalData( - mm_items=node_items, mm_hashes=node_hashes + if k < len(placeholders): + node_placeholders.setdefault(modality, []).append(placeholders[k]) + if node_items or node_hashes or node_placeholders: + trace.nodes[node_id].multi_modal_data = _validate_raw_mm_data( + MultiModalData( + mm_items=node_items, + mm_hashes=node_hashes, + mm_placeholders=node_placeholders, + ) ) diff --git a/verifiers/v1/interception/server.py b/verifiers/v1/interception/server.py index b9e5574d19..bda31795c2 100644 --- a/verifiers/v1/interception/server.py +++ b/verifiers/v1/interception/server.py @@ -37,6 +37,7 @@ from verifiers.v1.dialects import DIALECTS, Dialect from verifiers.v1.dialects.base import is_sse_done_event from verifiers.v1.errors import ( + InterceptionError, OverlongPromptError, ProviderError, RolloutError, @@ -305,6 +306,18 @@ async def handle_request( # alias after parsing so the wire body does not survive model inference. request._read_bytes = None del raw + try: + body = await session.ctx.client.prepare_request_body(dialect, body) + except RolloutError as e: + return self._fail(session, dialect, e) + except Exception as e: # noqa: BLE001 - ingress boundary surfaces every prep failure + return self._fail( + session, + dialect, + InterceptionError( + f"request preparation failed: {type(e).__name__}: {e}" + ), + ) streaming = dialect.streaming(body) logger.debug( "intercept %s: id=%s stream=%s", diff --git a/verifiers/v1/legacy.py b/verifiers/v1/legacy.py index d964bae322..bde801702d 100644 --- a/verifiers/v1/legacy.py +++ b/verifiers/v1/legacy.py @@ -13,6 +13,7 @@ v1 stays importable without the v0 package present. """ +import asyncio import contextlib import logging from pathlib import Path @@ -193,6 +194,7 @@ def _to_v1_tokens(raw: Any) -> TurnTokens | None: prompt_ids=list(raw.get("prompt_ids") or []), completion_ids=list(raw.get("completion_ids") or []), completion_logprobs=list(raw.get("completion_logprobs") or []), + multi_modal_data=raw.get("multi_modal_data"), ) @@ -421,6 +423,20 @@ def _v0_client(self, client_config: ClientConfig, model: str): self._clients[key] = resolve_client(v0_config) return self._clients[key] + async def _state_output_with_live_trajectory(self, state: Any) -> dict: + """Build v0 rollout output metadata while preserving live trajectory sidecars. + + The JSON save path deltas ``tokens.multi_modal_data`` to avoid repeated + cumulative multimodal sidecars. Trace reconstruction needs the live, + cumulative sidecar for each turn so image descriptors align with the + full prompt the renderer saw. + """ + from verifiers.utils.save_utils import state_to_output + + out = await asyncio.to_thread(state_to_output, state, []) + out["trajectory"] = state.get("trajectory", []) + return out + async def _run_v0( self, task_idx: int, @@ -429,13 +445,13 @@ async def _run_v0( sampling: SamplingConfig, ) -> dict: client = self._v0_client(client_config, model) - return await self.env.run_rollout( + state = await self.env._run_rollout_state( input=dict(self.dataset[task_idx]), client=client, model=model, sampling_args=sampling.model_dump(exclude_none=True), - state_columns=["trajectory"], ) + return await self._state_output_with_live_trajectory(state) @staticmethod def _row(req: RunRequest) -> int: @@ -460,12 +476,14 @@ async def _run(self, req: RunRequest) -> RunResponse: async def _run_group(self, req: RunGroupRequest) -> RunGroupResponse: client = self._v0_client(req.client, req.model) # run_group scores the rollouts together so group/preference reward funcs apply. - outs = await self.env.run_group( + states = await self.env._run_group_states( group_inputs=[dict(self.dataset[req.task_idx]) for _ in range(req.n)], client=client, model=req.model, sampling_args=req.sampling.model_dump(exclude_none=True), - state_columns=["trajectory"], + ) + outs = await asyncio.gather( + *(self._state_output_with_live_trajectory(state) for state in states) ) traces = [ rollout_output_to_trace(out, req.task_idx).model_dump() for out in outs diff --git a/verifiers/v1/trace.py b/verifiers/v1/trace.py index 6fbf25e403..88adc8bf89 100644 --- a/verifiers/v1/trace.py +++ b/verifiers/v1/trace.py @@ -224,7 +224,12 @@ def logprobs(self) -> list[float]: @property def multi_modal_data(self) -> MultiModalData | None: - """Node image data concatenated in token order for training; never persisted.""" + """The branch's multimodal sidecar — every node's images concatenated in path order. + + None when the branch has no images. The raw-image path carries lightweight descriptors + plus placeholder ranges, so downstream vLLM/training multimodal payloads can align hashes, + placeholders, and item refs without reprocessing images in the env worker. + """ merged = MultiModalData() found = False for node in self.nodes: @@ -236,6 +241,8 @@ def multi_modal_data(self) -> MultiModalData | None: merged.mm_items.setdefault(modality, []).extend(items) for modality, hashes in mmd.mm_hashes.items(): merged.mm_hashes.setdefault(modality, []).extend(hashes) + for modality, placeholders in mmd.mm_placeholders.items(): + merged.mm_placeholders.setdefault(modality, []).extend(placeholders) return merged if found else None @property diff --git a/verifiers/v1/types.py b/verifiers/v1/types.py index dcb99f7067..00fa5e7856 100644 --- a/verifiers/v1/types.py +++ b/verifiers/v1/types.py @@ -202,8 +202,8 @@ class TurnTokens(BaseModel): default=None, exclude=True ) is_content: list[bool] | None = Field(default=None, exclude=True) - # Transient carrier (excluded): the renderer's multimodal sidecar (image tensors + offsets), - # attributed per node by the turn's `commit`, then dropped — never persisted. + # Transient carrier (excluded): the renderer's multimodal sidecar (raw-image descriptors, + # hashes, and placeholder offsets), attributed per node by the turn's `commit`, then dropped. multi_modal_data: MultiModalData | None = Field(default=None, exclude=True) # Transient carrier (excluded): the MoE expert-routing data from `generate` (expert ids # per token), attributed per node by the turn's `commit` into `MessageNode.routed_experts`,