From 6f5dccb8172d7706512e59cc1f6aaa11b0470bbe Mon Sep 17 00:00:00 2001 From: eligotts Date: Thu, 20 Aug 2026 23:22:15 +0000 Subject: [PATCH] feat: retain inline multimodal rollout messages --- tests/v1/test_graph.py | 46 ++++++++ verifiers/v1/clients/train.py | 40 +++---- verifiers/v1/graph.py | 216 +++++++++++----------------------- verifiers/v1/trace.py | 18 --- verifiers/v1/types.py | 5 +- 5 files changed, 136 insertions(+), 189 deletions(-) diff --git a/tests/v1/test_graph.py b/tests/v1/test_graph.py index 3341b749d6..10defdb83d 100644 --- a/tests/v1/test_graph.py +++ b/tests/v1/test_graph.py @@ -1,6 +1,7 @@ import base64 import numpy as np +import pytest import verifiers.v1 as vf from verifiers.v1 import graph @@ -232,6 +233,51 @@ def second_turn(trace, prompt_ids): ] +def test_expanded_prompt_is_canonical_while_bridge_uses_logical_tokens(): + trace = vf.Trace( + agent=vf.AgentInfo(config=vf.AgentConfig()), + task=vf.TraceTask(type="Task", data=vf.TaskData(idx=0, prompt="x")), + ) + user = vf.UserMessage(content="image") + assistant = vf.AssistantMessage(content="looked") + graph.prepare_turn(trace, [user]).commit( + vf.Response( + id="a", + created=0, + model="t", + message=assistant, + finish_reason="stop", + tokens=TurnTokens( + prompt_ids=[1, 9, 9, 3], + renderer_prompt_ids=[1, 2, 3], + completion_ids=[4], + message_spans=[(0, 2)], + ), + ) + ) + + assert trace.branches[0].token_ids == [1, 9, 9, 3, 4] + turn = graph.prepare_turn(trace, [user, assistant, vf.UserMessage(content="next")]) + assert turn.previous_token_ids() == ([1, 2, 3], [4]) + + with pytest.raises(ValueError, match="exactly extend"): + turn.commit( + vf.Response( + id="b", + created=0, + model="t", + message=vf.AssistantMessage(content="bad"), + finish_reason="stop", + tokens=TurnTokens( + prompt_ids=[1, 9, 8, 3, 4, 5], + renderer_prompt_ids=[1, 2, 3, 4, 5], + completion_ids=[6], + message_spans=[None, None, (4, 5)], + ), + ) + ) + + def test_prompt_supplied_assistant_messages_are_not_sampled_turns(): task = vf.TaskData(idx=0, prompt="few-shot") trace = vf.Trace( diff --git a/verifiers/v1/clients/train.py b/verifiers/v1/clients/train.py index 2670b1fd37..298f80ac8c 100644 --- a/verifiers/v1/clients/train.py +++ b/verifiers/v1/clients/train.py @@ -149,11 +149,11 @@ def response_from_generate( # million-token contexts synchronously on the event loop. tokens=TurnTokens.model_construct( prompt_ids=prompt_ids, + renderer_prompt_ids=result.get("renderer_prompt_ids"), completion_ids=completion_ids, completion_logprobs=result.get("completion_logprobs") or [], message_spans=message_spans, is_content=attribution.is_content if attribution is not None else None, - multi_modal_data=result.get("multi_modal_data"), routed_experts=result.get("routed_experts"), kept_tokens=KeptTokens(**kept) if (kept := result.get("kept_tokens")) @@ -348,11 +348,9 @@ async def get_response( from renderers.client import generate wire_tools = [tool_to_wire(t) for t in tools] if tools else None - wire_messages = ( - [message_to_wire(m) for m in turn.tail] if turn is not None else [] - ) + wire_messages = [message_to_wire(m) for m in prompt] + wire_tail = [message_to_wire(m) for m in turn.tail] if turn is not None else [] prompt_ids: list[int] | None = None - multi_modal_data = None prompt_attribution: RenderedTokens | None = None model = body["model"] raw_sampling = sampling.model_dump(exclude_none=True) @@ -371,13 +369,15 @@ async def get_response( async with pool.acquire() as slot: renderer = slot.renderer - # 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) - ) + raw_multimodal = _has_multimodal_content(prompt) + if raw_multimodal and not getattr( + renderer, "supports_raw_multimodal", False + ): + raise NotImplementedError( + f"{type(renderer).__name__} does not support raw multimodal inference" + ) + # Only build the O(context) previous token stream for a bridgeable tail. + can_bridge = turn is not None and _is_valid_incremental_tail(wire_tail) 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 @@ -386,33 +386,33 @@ def bridge(): return renderer.bridge_to_next_turn( previous_prompt_ids, previous_completion_ids, - wire_messages, + wire_tail, tools=wire_tools, + **({"process_multimodal": False} if raw_multimodal else {}), ) bridged = await slot.run(bridge) if bridged is not None: prompt_ids = bridged.token_ids - multi_modal_data = bridged.multi_modal_data prompt_attribution = bridged bridged_turn = turn sampling_params["routed_experts_prompt_start"] = max( - len(previous_prompt_ids) + len(previous_completion_ids) - 1, - 0, + turn.path_len - 1, 0 ) # Render here (encode-side, so through the slot) rather than inside `generate`: # handed prebuilt prompt_ids, generate's own renderer touches are decode-side # and stop-id reads, safe on a bare renderer without lock or thread hop. if prompt_ids is None: - wire_messages = [message_to_wire(m) for m in prompt] rendered = await slot.run( lambda: renderer.render( - wire_messages, tools=wire_tools, add_generation_prompt=True + wire_messages, + tools=wire_tools, + add_generation_prompt=True, + **({"process_multimodal": False} if raw_multimodal else {}), ) ) prompt_ids = rendered.token_ids - multi_modal_data = rendered.multi_modal_data prompt_attribution = rendered try: @@ -422,10 +422,10 @@ def bridge(): messages=wire_messages, model=model, prompt_ids=prompt_ids, - multi_modal_data=multi_modal_data, prompt_attribution=prompt_attribution, tools=wire_tools, sampling_params=sampling_params, + raw_multimodal=raw_multimodal, extra_headers={SESSION_ID_HEADER: session_id} if session_id else None, diff --git a/verifiers/v1/graph.py b/verifiers/v1/graph.py index bb8f6b6eee..2debea0f2f 100644 --- a/verifiers/v1/graph.py +++ b/verifiers/v1/graph.py @@ -27,7 +27,7 @@ import numpy as np from pydantic import BaseModel, ConfigDict, Field, field_serializer, field_validator from pydantic.json_schema import SkipJsonSchema -from renderers.base import MultiModalData, PlaceholderRange, RenderedTokens +from renderers.base import RenderedTokens from verifiers.v1.types import ( AssistantMessage, @@ -85,6 +85,8 @@ class MessageNode(BaseModel): template scaffold + its own tokens — for an assistant, the generation-prompt scaffold followed by the sampled completion. Concatenated along a path, these reproduce the exact `prompt_ids + completion_ids` the model saw.""" + renderer_token_ids: list[int] = Field(default_factory=list, exclude=True) + """Logical renderer tokens retained only while extending a live rollout.""" mask: list[bool] = Field(default_factory=list) """Per-token, parallel to `token_ids`: True for trainable, model-sampled tokens (only an assistant node's completion span); False for template scaffold and every input-message @@ -111,12 +113,6 @@ class MessageNode(BaseModel): `logprobs`. None means no reference model scored this node.""" loss_weights: dict[str, list[float]] | None = None """Named loss-weight streams aligned to `token_ids`, consumer-stamped.""" - 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).""" 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 @@ -132,50 +128,6 @@ class MessageNode(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) - @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"`).""" - if mmd is None: - return None - return { - "mm_hashes": {k: list(v) for k, v in mmd.mm_hashes.items()}, - "mm_placeholders": { - modality: [{"offset": p.offset, "length": p.length} for p in ranges] - 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 - ] - for modality, items in mmd.mm_items.items() - }, - } - - @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): - return 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() - }, - ) - @field_serializer("routed_experts") def serialize_ndarray_field(self, arr: np.ndarray | None) -> dict | None: """Integer array -> raw-bytes `__nd__` dict so it rides the wire (numpy can't JSON).""" @@ -333,27 +285,24 @@ def previous_token_ids(self) -> tuple[list[int], list[int]] | None: """Return `(previous_prompt_ids, previous_completion_ids)` for a bridge anchor. The anchor must end at a sampled assistant node. That node stores generation-prompt - scaffold followed by sampled completion tokens, so split at the first sampled token. + scaffold followed by sampled completion tokens, so split off the sampled suffix. """ if not self.prefix_node_ids: return None last = self.trace.nodes[self.prefix_node_ids[-1]] if not last.sampled: return None - first_sampled = next( - (i for i, sampled in enumerate(last.mask) if sampled), None - ) - if first_sampled is None: - return None - if any(not sampled for sampled in last.mask[first_sampled:]): + num_sampled = sum(last.mask) + if not num_sampled: return None prompt_ids: list[int] = [] for nid in self.prefix_node_ids[:-1]: - prompt_ids.extend(self.trace.nodes[nid].token_ids) - prompt_ids.extend(last.token_ids[:first_sampled]) - # Slicing already returns an independent list; avoid a second completion-sized copy. - completion_ids = last.token_ids[first_sampled:] + node = self.trace.nodes[nid] + prompt_ids.extend(node.renderer_token_ids or node.token_ids) + last_ids = last.renderer_token_ids or last.token_ids + prompt_ids.extend(last_ids[:-num_sampled]) + completion_ids = last_ids[-num_sampled:] if not prompt_ids or not completion_ids: return None return prompt_ids, completion_ids @@ -364,15 +313,27 @@ def prompt_message_spans( """Convert bridge-tail attribution into full-prompt message spans.""" # Reused bridge tokens are unattributed, so scan only the newly rendered tail. tail_spans = RenderedTokens( - message_indices=tail_attribution.message_indices[self.path_len :], + message_indices=tail_attribution.message_indices[self.renderer_path_len :], message_roles=tail_attribution.message_roles, ).message_token_spans() # Tail spans are slice-relative; restore their full-prompt token offsets. return [None] * self.tail_start + [ - None if span is None else (span[0] + self.path_len, span[1] + self.path_len) + None + if span is None + else (span[0] + self.renderer_path_len, span[1] + self.renderer_path_len) for span in tail_spans ] + @property + def renderer_path_len(self) -> int: + return sum( + len( + self.trace.nodes[nid].renderer_token_ids + or self.trace.nodes[nid].token_ids + ) + for nid in self.prefix_node_ids + ) + 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) @@ -430,53 +391,6 @@ def prepare_turn(trace: Trace, prompt: list[Message]) -> PendingTurn: ) -def _part_modality(part) -> str | None: - """The multimodal modality a content part introduces (currently only images), or None.""" - return "image" if getattr(part, "type", None) == "image_url" else None - - -def _attribute_mm( - trace: Trace, - path: list[tuple[int, Message]], - num_reused: int, - mmd: MultiModalData | None, -) -> None: - """Attach each new image's renderer item to the node whose message introduced it. The - 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.""" - if mmd is None or mmd.is_empty(): - return - cursors: dict[str, int] = {} - for pos, (node_id, msg) in enumerate(path): - content = msg.content - if not isinstance(content, list): - continue - node_items: dict[str, list] = {} - node_hashes: dict[str, list] = {} - for part in content: - modality = _part_modality(part) - if modality is None: - continue - k = cursors.get(modality, 0) - cursors[modality] = k + 1 - # Reused prefix: advance the cursor over its media, don't re-attribute. - if pos < num_reused: - continue - items = mmd.mm_items.get(modality) or [] - hashes = mmd.mm_hashes.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 - ) - - def _attribute_routed_experts( trace: Trace, new_node_ids: list[int], @@ -534,62 +448,72 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int: trace = turn.trace prompt = turn.prompt tokens = response.tokens - multi_modal_data = tokens.multi_modal_data if tokens else None prompt_ids = tokens.prompt_ids if tokens else [] + renderer_prompt_ids = ( + tokens.renderer_prompt_ids + if tokens and tokens.renderer_prompt_ids is not None + else prompt_ids + ) spans = tokens.message_spans if tokens else None is_content = tokens.is_content if tokens else None - has_is_content = is_content is not None and len(is_content) == len(prompt_ids) + attribution_matches = renderer_prompt_ids == prompt_ids + has_is_content = ( + attribution_matches + and is_content is not None + and len(is_content) == len(prompt_ids) + ) idx = _head_index(trace) - # Token-based prefix reuse. `prepare_turn` matched the prefix by message hash (content); when - # this turn carries token ids, tighten that to token identity — the stored prefix must be an - # exact token prefix of what the model saw this turn (`prompt_ids`). Reuse whole nodes within - # the longest common token prefix and fork at the first divergence, so a retokenized prior - # (BPE drift, dropped ``, rewritten tool calls) branches off with this turn's real - # tokens instead of silently inheriting stale ones. Comparing the *concatenated* prefix (not - # per-message spans) is what makes this correct: a prior assistant's stored generation form - # and its re-rendered input form place the turn-close scaffold in different nodes but at the - # same position, so only a genuine content/token change shifts the common prefix. The bridge - # keeps the prior verbatim so it matches fully (stays linear); the eval relay carries no token - # ids and keeps the message-hash prefix. prefix = turn.prefix_node_ids - path_len = turn.path_len # cumulative stored token length of the reused prefix + path_len = turn.path_len + renderer_path_len = turn.renderer_path_len if tokens is not None and prefix: - # Compare node by node against the prompt_ids slice at the running offset (C-level list - # ==, short-circuits at the first divergent node) — no full concatenation materialized. keep = 0 off = 0 + renderer_off = 0 for nid in prefix: - node_tokens = trace.nodes[nid].token_ids - if prompt_ids[off : off + len(node_tokens)] != node_tokens: + node = trace.nodes[nid] + node_renderer_ids = node.renderer_token_ids or node.token_ids + if ( + prompt_ids[off : off + len(node.token_ids)] != node.token_ids + or renderer_prompt_ids[ + renderer_off : renderer_off + len(node_renderer_ids) + ] + != node_renderer_ids + ): break - off += len(node_tokens) + off += len(node.token_ids) + renderer_off += len(node_renderer_ids) keep += 1 + if tokens.renderer_prompt_ids is not None and keep != len(prefix): + raise ValueError( + "vLLM prompt tokens do not exactly extend the stored rollout prefix" + ) prefix = prefix[:keep] path_len = off + renderer_path_len = renderer_off num_reused = len(prefix) parent = prefix[-1] if prefix else None - # cursor: in prompt_ids, the end of the previous *new* message's tokens cursor: int | None = None - # Track new nodes separately so routed-expert attribution does not need this full path. + renderer_cursor: int | None = None new_node_ids: list[int] = [] - # Materialize the reused message path only for multimodal cursor attribution. - mm_path: list[tuple[int, Message]] | None = None - if multi_modal_data is not None: - mm_path = [(nid, prompt[i]) for i, nid in enumerate(prefix)] for i, msg in enumerate(prompt[num_reused:], start=num_reused): key = (parent, message_hash(msg)) start = path_len if cursor is None else cursor - span = spans[i] if spans and i < len(spans) else None + span = spans[i] if attribution_matches and spans and i < len(spans) else None end = span[1] if span else start node_tokens = prompt_ids[start:end] + renderer_start = ( + renderer_path_len if renderer_cursor is None else renderer_cursor + ) + renderer_span = spans[i] if spans and i < len(spans) else None + renderer_end = renderer_span[1] if renderer_span else renderer_start trace.nodes.append( - # Every value is already typed framework data; avoid revalidating and copying - # potentially huge token slices a second time. MessageNode.model_construct( parent=parent, message=msg, token_ids=node_tokens, + renderer_token_ids=renderer_prompt_ids[renderer_start:renderer_end], mask=[False] * len(node_tokens), is_content=is_content[start:end] if has_is_content else [], ) @@ -597,20 +521,23 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int: parent = len(trace.nodes) - 1 idx[key] = parent new_node_ids.append(parent) - if mm_path is not None: - mm_path.append((parent, msg)) cursor = end + renderer_cursor = renderer_end - # Assistant node: trailing scaffold (the generation prompt) + the sampled completion. comp_ids = tokens.completion_ids if tokens else [] gen_start = path_len if cursor is None else cursor gen_prompt = prompt_ids[gen_start:] + renderer_gen_start = ( + renderer_path_len if renderer_cursor is None else renderer_cursor + ) + renderer_gen_prompt = renderer_prompt_ids[renderer_gen_start:] trace.nodes.append( MessageNode.model_construct( parent=parent, message=response.message, sampled=True, token_ids=[*gen_prompt, *comp_ids], + renderer_token_ids=[*renderer_gen_prompt, *comp_ids], mask=[False] * len(gen_prompt) + [True] * len(comp_ids), is_content=([False] * len(gen_prompt) + [True] * len(comp_ids)) if has_is_content @@ -619,15 +546,10 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int: logprobs=tokens.completion_logprobs if tokens else [], ) ) - # Register the assistant so the next turn's prompt (which restates it) reuses this node. assistant_id = len(trace.nodes) - 1 idx[(parent, message_hash(response.message))] = assistant_id new_node_ids.append(assistant_id) - # Attribute this turn's images onto the input nodes that introduced them (by content part). - if mm_path is not None: - _attribute_mm(trace, mm_path, num_reused, multi_modal_data) - # Attribute this turn's expert-routing array onto the nodes created this turn (new input # nodes in creation order, then the assistant node), each getting the routing for its tokens. _attribute_routed_experts( diff --git a/verifiers/v1/trace.py b/verifiers/v1/trace.py index 8d6b0f9776..709752565c 100644 --- a/verifiers/v1/trace.py +++ b/verifiers/v1/trace.py @@ -8,7 +8,6 @@ import numpy as np from pydantic import BaseModel, Field, PrivateAttr -from renderers.base import MultiModalData from typing_extensions import TypeVar if TYPE_CHECKING: @@ -41,7 +40,6 @@ EXCLUDE_FIELDS: dict = { "nodes": { "__all__": { - "multi_modal_data", "routed_experts", "kept_tokens", } @@ -283,22 +281,6 @@ def loss_weights(self, name: str) -> list[float] | None: weights.extend(node_weights) return weights - @property - def multi_modal_data(self) -> MultiModalData | None: - """Node image data concatenated in token order for training; never persisted.""" - merged = MultiModalData() - found = False - for node in self.nodes: - mmd = node.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) - return merged if found else None - @property def routed_experts(self) -> np.ndarray | None: """uint8 `[tokens, layers, top_k]` routing; partial data returns None.""" diff --git a/verifiers/v1/types.py b/verifiers/v1/types.py index 703a2c2263..79997b7a0d 100644 --- a/verifiers/v1/types.py +++ b/verifiers/v1/types.py @@ -3,7 +3,6 @@ from typing import Annotated, Any, Literal from pydantic import AliasChoices, BaseModel, ConfigDict, Field -from renderers.base import MultiModalData from typing_extensions import TypedDict @@ -200,6 +199,7 @@ class TurnTokens(BaseModel): model_config = ConfigDict(arbitrary_types_allowed=True) prompt_ids: list[int] = Field(default_factory=list) + renderer_prompt_ids: list[int] | None = Field(default=None, exclude=True) completion_ids: list[int] = Field(default_factory=list) completion_logprobs: list[float] = Field(default_factory=list) @@ -209,9 +209,6 @@ 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. - 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`, # then dropped. None unless the engine ran with `enable_return_routed_experts`.