verifiers 0.3.2.dev149__py3-none-any.whl → 0.3.2.dev151__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -140,6 +140,16 @@ def response_from_generate(
140
140
  message_spans = bridged_turn.prompt_message_spans(attribution)
141
141
  else:
142
142
  message_spans = attribution.message_token_spans()
143
+ raw_mm_placeholders = result.get("mm_placeholders")
144
+ mm_placeholders = (
145
+ sorted(
146
+ (placeholder["offset"], placeholder["length"])
147
+ for ranges in raw_mm_placeholders.values()
148
+ for placeholder in ranges
149
+ )
150
+ if raw_mm_placeholders is not None
151
+ else None
152
+ )
143
153
  return Response(
144
154
  id=result.get("request_id", ""),
145
155
  created=0,
@@ -159,11 +169,13 @@ def response_from_generate(
159
169
  # million-token contexts synchronously on the event loop.
160
170
  tokens=TurnTokens.model_construct(
161
171
  prompt_ids=prompt_ids,
172
+ renderer_prompt_ids=result.get("renderer_prompt_ids"),
173
+ bridged=bridged_turn is not None,
162
174
  completion_ids=completion_ids,
163
175
  completion_logprobs=result.get("completion_logprobs") or [],
164
176
  message_spans=message_spans,
165
177
  is_content=attribution.is_content if attribution is not None else None,
166
- multi_modal_data=result.get("multi_modal_data"),
178
+ mm_placeholders=mm_placeholders,
167
179
  mm_token_type_id_map=mm_token_type_id_map,
168
180
  routed_experts=result.get("routed_experts"),
169
181
  sampling_mask=SamplingMask.from_sampling_mask(mask)
@@ -358,11 +370,9 @@ class TrainClient(Client):
358
370
  from renderers.client import generate
359
371
 
360
372
  wire_tools = [tool_to_wire(t) for t in tools] if tools else None
361
- wire_messages = (
362
- [message_to_wire(m) for m in turn.tail] if turn is not None else []
363
- )
373
+ wire_messages = [message_to_wire(m) for m in prompt]
374
+ wire_tail = [message_to_wire(m) for m in turn.tail] if turn is not None else []
364
375
  prompt_ids: list[int] | None = None
365
- multi_modal_data = None
366
376
  prompt_attribution: RenderedTokens | None = None
367
377
  model = body["model"]
368
378
  sampling_params = sampling.wire_args()
@@ -381,14 +391,18 @@ class TrainClient(Client):
381
391
  mm_token_type_id_map = (
382
392
  renderer.mm_token_type_id_map if is_multimodal(renderer) else None
383
393
  )
384
- # Only build the (O(context)) previous-turn token ids once the cheap guards pass — a
385
- # multimodal prompt or a tail that isn't a clean `[tool*, user?]` extension can't bridge.
386
- can_bridge = (
387
- turn is not None
388
- and not _has_multimodal_content(prompt)
389
- and _is_valid_incremental_tail(wire_messages)
390
- )
391
- previous_ids = turn.previous_token_ids() if can_bridge else None
394
+ has_images = _has_multimodal_content(prompt)
395
+ process_multimodal = not has_images
396
+ if has_images and not getattr(
397
+ renderer, "supports_process_multimodal", False
398
+ ):
399
+ raise NotImplementedError(
400
+ f"{type(renderer).__name__} does not support process_multimodal=False"
401
+ )
402
+ render_kwargs = {} if process_multimodal else {"process_multimodal": False}
403
+ # Only build the O(context) previous token stream for a bridgeable tail.
404
+ can_bridge = turn is not None and _is_valid_incremental_tail(wire_tail)
405
+ previous_ids = turn.previous_renderer_token_ids() if can_bridge else None
392
406
  if previous_ids is not None:
393
407
  previous_prompt_ids, previous_completion_ids = previous_ids
394
408
 
@@ -396,33 +410,33 @@ class TrainClient(Client):
396
410
  return renderer.bridge_to_next_turn(
397
411
  previous_prompt_ids,
398
412
  previous_completion_ids,
399
- wire_messages,
413
+ wire_tail,
400
414
  tools=wire_tools,
415
+ **render_kwargs,
401
416
  )
402
417
 
403
418
  bridged = await slot.run(bridge)
404
419
  if bridged is not None:
405
420
  prompt_ids = bridged.token_ids
406
- multi_modal_data = bridged.multi_modal_data
407
421
  prompt_attribution = bridged
408
422
  bridged_turn = turn
409
423
  sampling_params["routed_experts_prompt_start"] = max(
410
- len(previous_prompt_ids) + len(previous_completion_ids) - 1,
411
- 0,
424
+ turn.path_len - 1, 0
412
425
  )
413
426
 
414
427
  # Render here (encode-side, so through the slot) rather than inside `generate`:
415
428
  # handed prebuilt prompt_ids, generate's own renderer touches are decode-side
416
429
  # and stop-id reads, safe on a bare renderer without lock or thread hop.
417
430
  if prompt_ids is None:
418
- wire_messages = [message_to_wire(m) for m in prompt]
419
431
  rendered = await slot.run(
420
432
  lambda: renderer.render(
421
- wire_messages, tools=wire_tools, add_generation_prompt=True
433
+ wire_messages,
434
+ tools=wire_tools,
435
+ add_generation_prompt=True,
436
+ **render_kwargs,
422
437
  )
423
438
  )
424
439
  prompt_ids = rendered.token_ids
425
- multi_modal_data = rendered.multi_modal_data
426
440
  prompt_attribution = rendered
427
441
 
428
442
  try:
@@ -432,10 +446,10 @@ class TrainClient(Client):
432
446
  messages=wire_messages,
433
447
  model=model,
434
448
  prompt_ids=prompt_ids,
435
- multi_modal_data=multi_modal_data,
436
449
  prompt_attribution=prompt_attribution,
437
450
  tools=wire_tools,
438
451
  sampling_params=sampling_params,
452
+ process_multimodal=process_multimodal,
439
453
  cache_salt=cache_salt,
440
454
  extra_headers={SESSION_ID_HEADER: session_id}
441
455
  if session_id
verifiers/v1/graph.py CHANGED
@@ -37,7 +37,7 @@ from pydantic import (
37
37
  field_validator,
38
38
  )
39
39
  from pydantic.json_schema import SkipJsonSchema
40
- from renderers.base import MultiModalData, PlaceholderRange, RenderedTokens
40
+ from renderers.base import RenderedTokens
41
41
 
42
42
  from verifiers.v1.semantic import ParentLink
43
43
  from verifiers.v1.types import (
@@ -114,6 +114,10 @@ class MessageNode(BaseModel):
114
114
  template scaffold + its own tokens — for an assistant, the generation-prompt scaffold
115
115
  followed by the sampled completion. Concatenated along a path, these reproduce the exact
116
116
  `prompt_ids + completion_ids` the model saw."""
117
+ renderer_token_ids: list[int] | None = Field(default=None, exclude=True)
118
+ """Logical renderer tokens retained only while extending a live rollout.
119
+ None means they are identical to `token_ids`; an empty list is a real empty slice."""
120
+
117
121
  mask: list[bool] = Field(default_factory=list)
118
122
  """Per-token, parallel to `token_ids`: True for trainable, model-sampled tokens (only an
119
123
  assistant node's completion span); False for template scaffold and every input-message
@@ -146,12 +150,6 @@ class MessageNode(BaseModel):
146
150
  `logprobs`. None means no trainer forward annotated this node."""
147
151
  loss_weights: dict[str, list[float]] | None = None
148
152
  """Named loss-weight streams aligned to `token_ids`, consumer-stamped."""
149
- multi_modal_data: SkipJsonSchema[MultiModalData | None] = None
150
- """The renderer items for the images this message's content introduces (pixel tensors,
151
- grids, hashes, placeholders) — the only carrier of the pixels from the env server to the
152
- trainer. `Branch.multi_modal_data` concatenates them along the path into the training
153
- `mm_kwargs`. Rides the wire as raw bytes (msgpack `bin`) since pydantic can't JSON the numpy;
154
- kept off disk by the dump-site `exclude` in prime-rl (the tensors bloat the rollout jsonl)."""
155
153
  routed_experts: SkipJsonSchema[np.ndarray | None] = None
156
154
  """This node's slice of the MoE expert-routing array — uint8 `[len(token_ids), layers,
157
155
  top_k]`, the expert ids inference selected for exactly this node's tokens. Attributed from
@@ -167,6 +165,14 @@ class MessageNode(BaseModel):
167
165
 
168
166
  model_config = ConfigDict(arbitrary_types_allowed=True)
169
167
 
168
+ @property
169
+ def logical_ids(self) -> list[int]:
170
+ return (
171
+ self.token_ids
172
+ if self.renderer_token_ids is None
173
+ else self.renderer_token_ids
174
+ )
175
+
170
176
  @field_serializer(
171
177
  "logprobs",
172
178
  "advantages",
@@ -183,50 +189,6 @@ class MessageNode(BaseModel):
183
189
  return values
184
190
  return [round(value, decimals) for value in values]
185
191
 
186
- @field_serializer("multi_modal_data")
187
- def serialize_multi_modal_data(self, mmd: MultiModalData | None) -> dict | None:
188
- """`MultiModalData` -> msgpack-safe dict so the pixel tensors ride the wire; numpy
189
- `mm_items` values become raw-bytes `__nd__` dicts (every renderer emits `return_tensors="np"`)."""
190
- if mmd is None:
191
- return None
192
- return {
193
- "mm_hashes": {k: list(v) for k, v in mmd.mm_hashes.items()},
194
- "mm_placeholders": {
195
- modality: [{"offset": p.offset, "length": p.length} for p in ranges]
196
- for modality, ranges in mmd.mm_placeholders.items()
197
- },
198
- "mm_items": {
199
- modality: [
200
- {k: _encode_ndarray(v) for k, v in item.items()} for item in items
201
- ]
202
- for modality, items in mmd.mm_items.items()
203
- },
204
- }
205
-
206
- @field_validator("multi_modal_data", mode="before")
207
- @classmethod
208
- def deserialize_multi_modal_data(cls, value: Any) -> MultiModalData | None:
209
- if value is None or isinstance(value, MultiModalData):
210
- return value
211
- if not isinstance(value, dict):
212
- raise TypeError(f"cannot build MultiModalData from {type(value).__name__}")
213
- return MultiModalData(
214
- mm_hashes={k: list(v) for k, v in (value.get("mm_hashes") or {}).items()},
215
- mm_placeholders={
216
- modality: [
217
- PlaceholderRange(offset=p["offset"], length=p["length"])
218
- for p in ranges
219
- ]
220
- for modality, ranges in (value.get("mm_placeholders") or {}).items()
221
- },
222
- mm_items={
223
- modality: [
224
- {k: _decode_ndarray(v) for k, v in item.items()} for item in items
225
- ]
226
- for modality, items in (value.get("mm_items") or {}).items()
227
- },
228
- )
229
-
230
192
  @field_serializer("routed_experts")
231
193
  def serialize_ndarray_field(self, arr: np.ndarray | None) -> dict | None:
232
194
  """Integer array -> raw-bytes `__nd__` dict so it rides the wire (numpy can't JSON)."""
@@ -405,9 +367,10 @@ def _matching_node(
405
367
  parent: int | None,
406
368
  message: Message,
407
369
  token_ids: list[int] | None = None,
370
+ renderer_token_ids: list[int] | None = None,
408
371
  tools: list[Tool] | None = None,
409
372
  ) -> int | None:
410
- """Find an existing child, optionally requiring its exact physical token span.
373
+ """Find an existing child, optionally requiring its exact token spans.
411
374
 
412
375
  The head index deliberately points at only the latest content-equivalent child. Token-level
413
376
  prefix breaks can leave older physical variants under the same key, so a token mismatch falls
@@ -415,19 +378,24 @@ def _matching_node(
415
378
  """
416
379
  key = _node_key(parent, message, tools)
417
380
  indexed = _head_index(trace).get(key)
418
- if indexed is not None and (
419
- token_ids is None or trace.nodes[indexed].token_ids == token_ids
381
+ if (
382
+ indexed is not None
383
+ and (token_ids is None or trace.nodes[indexed].token_ids == token_ids)
384
+ and (
385
+ renderer_token_ids is None
386
+ or trace.nodes[indexed].logical_ids == renderer_token_ids
387
+ )
420
388
  ):
421
389
  return indexed
422
- if token_ids is None:
390
+ if token_ids is None and renderer_token_ids is None:
423
391
  return indexed
424
392
  for node_id in range(len(trace.nodes) - 1, -1, -1):
425
393
  if node_id == indexed:
426
394
  continue
427
395
  node = trace.nodes[node_id]
428
396
  if (
429
- node.parent == parent
430
- and node.token_ids == token_ids
397
+ (token_ids is None or node.token_ids == token_ids)
398
+ and (renderer_token_ids is None or node.logical_ids == renderer_token_ids)
431
399
  and _node_key(node.parent, node.message, node.tools) == key
432
400
  ):
433
401
  return node_id
@@ -441,6 +409,9 @@ def _matching_prefix_node(
441
409
  prompt_ids: list[int],
442
410
  start: int,
443
411
  stop: int,
412
+ renderer_prompt_ids: list[int],
413
+ renderer_start: int,
414
+ renderer_stop: int,
444
415
  tools: list[Tool] | None = None,
445
416
  ) -> int | None:
446
417
  """Find the longest content-equivalent child matching inside `[start, stop]`.
@@ -472,6 +443,11 @@ def _matching_prefix_node(
472
443
  if start + len(trace.nodes[node_id].token_ids) <= stop
473
444
  and prompt_ids[start : start + len(trace.nodes[node_id].token_ids)]
474
445
  == trace.nodes[node_id].token_ids
446
+ and renderer_start + len(trace.nodes[node_id].logical_ids) <= renderer_stop
447
+ and renderer_prompt_ids[
448
+ renderer_start : renderer_start + len(trace.nodes[node_id].logical_ids)
449
+ ]
450
+ == trace.nodes[node_id].logical_ids
475
451
  ]
476
452
  return max(
477
453
  matches, key=lambda node_id: len(trace.nodes[node_id].token_ids), default=None
@@ -501,34 +477,31 @@ class PendingTurn:
501
477
  def tail(self) -> list[Message]:
502
478
  return self.prompt[self.tail_start :]
503
479
 
504
- def previous_token_ids(self) -> tuple[list[int], list[int]] | None:
505
- """Return `(previous_prompt_ids, previous_completion_ids)` for a bridge anchor.
480
+ def previous_renderer_token_ids(self) -> tuple[list[int], list[int]] | None:
481
+ """Return the logical renderer prompt and completion for a bridge anchor.
506
482
 
507
483
  The anchor must end at a sampled assistant node. That node stores generation-prompt
508
- scaffold followed by sampled completion tokens, so split at the first sampled token.
484
+ scaffold followed by sampled completion tokens, so split off the sampled suffix.
509
485
  """
510
486
  if not self.prefix_node_ids:
511
487
  return None
512
488
  last = self.trace.nodes[self.prefix_node_ids[-1]]
513
489
  if not last.sampled:
514
490
  return None
515
- first_sampled = next(
516
- (i for i, sampled in enumerate(last.mask) if sampled), None
517
- )
518
- if first_sampled is None:
519
- return None
520
- if any(not sampled for sampled in last.mask[first_sampled:]):
491
+ num_sampled = sum(last.mask)
492
+ if not num_sampled:
521
493
  return None
522
494
 
523
- prompt_ids: list[int] = []
495
+ renderer_prompt_ids: list[int] = []
524
496
  for nid in self.prefix_node_ids[:-1]:
525
- prompt_ids.extend(self.trace.nodes[nid].token_ids)
526
- prompt_ids.extend(last.token_ids[:first_sampled])
527
- # Slicing already returns an independent list; avoid a second completion-sized copy.
528
- completion_ids = last.token_ids[first_sampled:]
529
- if not prompt_ids or not completion_ids:
497
+ node = self.trace.nodes[nid]
498
+ renderer_prompt_ids.extend(node.logical_ids)
499
+ last_ids = last.logical_ids
500
+ renderer_prompt_ids.extend(last_ids[:-num_sampled])
501
+ completion_ids = last_ids[-num_sampled:]
502
+ if not renderer_prompt_ids or not completion_ids:
530
503
  return None
531
- return prompt_ids, completion_ids
504
+ return renderer_prompt_ids, completion_ids
532
505
 
533
506
  def prompt_message_spans(
534
507
  self, tail_attribution: RenderedTokens
@@ -536,15 +509,23 @@ class PendingTurn:
536
509
  """Convert bridge-tail attribution into full-prompt message spans."""
537
510
  # Reused bridge tokens are unattributed, so scan only the newly rendered tail.
538
511
  tail_spans = RenderedTokens(
539
- message_indices=tail_attribution.message_indices[self.path_len :],
512
+ message_indices=tail_attribution.message_indices[self.renderer_path_len :],
540
513
  message_roles=tail_attribution.message_roles,
541
514
  ).message_token_spans()
542
515
  # Tail spans are slice-relative; restore their full-prompt token offsets.
543
516
  return [None] * self.tail_start + [
544
- None if span is None else (span[0] + self.path_len, span[1] + self.path_len)
517
+ None
518
+ if span is None
519
+ else (span[0] + self.renderer_path_len, span[1] + self.renderer_path_len)
545
520
  for span in tail_spans
546
521
  ]
547
522
 
523
+ @property
524
+ def renderer_path_len(self) -> int:
525
+ return sum(
526
+ len(self.trace.nodes[nid].logical_ids) for nid in self.prefix_node_ids
527
+ )
528
+
548
529
  def commit(self, response: Response) -> int:
549
530
  """Add this turn to the graph; returns the committed assistant node's id."""
550
531
  assistant_id = _commit_turn(self, response)
@@ -639,53 +620,6 @@ def prepare_turn(
639
620
  )
640
621
 
641
622
 
642
- def _part_modality(part) -> str | None:
643
- """The multimodal modality a content part introduces (currently only images), or None."""
644
- return "image" if getattr(part, "type", None) == "image_url" else None
645
-
646
-
647
- def _attribute_mm(
648
- trace: Trace,
649
- path: list[tuple[int, Message]],
650
- num_reused: int,
651
- mmd: MultiModalData | None,
652
- ) -> None:
653
- """Attach each new image's renderer item to the node whose message introduced it. The
654
- renderer emits items per modality in prompt order (message order, then content-part order),
655
- so we walk the path advancing a per-modality cursor over every message's media but write
656
- only the nodes created this turn — `path[:num_reused]` is the reused prefix, already
657
- attributed when first created. Item order is all training needs; placeholder offsets aren't
658
- carried."""
659
- if mmd is None or mmd.is_empty():
660
- return
661
- cursors: dict[str, int] = {}
662
- for pos, (node_id, msg) in enumerate(path):
663
- content = msg.content
664
- if not isinstance(content, list):
665
- continue
666
- node_items: dict[str, list] = {}
667
- node_hashes: dict[str, list] = {}
668
- for part in content:
669
- modality = _part_modality(part)
670
- if modality is None:
671
- continue
672
- k = cursors.get(modality, 0)
673
- cursors[modality] = k + 1
674
- # Reused prefix: advance the cursor over its media, don't re-attribute.
675
- if pos < num_reused:
676
- continue
677
- items = mmd.mm_items.get(modality) or []
678
- hashes = mmd.mm_hashes.get(modality) or []
679
- if k < len(items):
680
- node_items.setdefault(modality, []).append(items[k])
681
- if k < len(hashes):
682
- node_hashes.setdefault(modality, []).append(hashes[k])
683
- if node_items:
684
- trace.nodes[node_id].multi_modal_data = MultiModalData(
685
- mm_items=node_items, mm_hashes=node_hashes
686
- )
687
-
688
-
689
623
  def _replace_placeholder_routing_row(
690
624
  trace: Trace, prefix_node_ids: list[int], arr: np.ndarray, off: int
691
625
  ) -> None:
@@ -771,46 +705,129 @@ def _attribute_sampling_mask(
771
705
  node.sampling_mask = payload
772
706
 
773
707
 
708
+ def _project_prompt_attribution(
709
+ renderer_prompt_ids: list[int],
710
+ prompt_ids: list[int],
711
+ mm_token_type_id_map: dict[int, int],
712
+ mm_placeholders: list[tuple[int, int]] | None,
713
+ message_spans: list[tuple[int, int] | None] | None,
714
+ is_content: list[bool] | None,
715
+ ) -> tuple[list[tuple[int, int] | None] | None, list[bool] | None]:
716
+ """Project logical attribution using vLLM's multimodal placeholder ranges."""
717
+ if renderer_prompt_ids == prompt_ids:
718
+ return message_spans, is_content
719
+ if not mm_token_type_id_map or mm_placeholders is None:
720
+ raise ValueError(
721
+ "cannot align renderer and vLLM prompt tokens without multimodal placeholders"
722
+ )
723
+
724
+ offsets = [0]
725
+ prompt_offset = 0
726
+ placeholder_index = 0
727
+ for token_id in renderer_prompt_ids:
728
+ if token_id in mm_token_type_id_map:
729
+ if placeholder_index >= len(mm_placeholders):
730
+ raise ValueError(
731
+ "vLLM multimodal placeholders do not align with renderer prompt"
732
+ )
733
+ offset, length = mm_placeholders[placeholder_index]
734
+ if offset != prompt_offset or length < 1:
735
+ raise ValueError(
736
+ "vLLM multimodal placeholders do not align with renderer prompt"
737
+ )
738
+ prompt_offset += length
739
+ placeholder_index += 1
740
+ else:
741
+ if (
742
+ prompt_offset >= len(prompt_ids)
743
+ or prompt_ids[prompt_offset] != token_id
744
+ ):
745
+ raise ValueError(
746
+ "renderer prompt does not align with vLLM prompt tokens"
747
+ )
748
+ prompt_offset += 1
749
+ offsets.append(prompt_offset)
750
+ if prompt_offset != len(prompt_ids) or placeholder_index != len(mm_placeholders):
751
+ raise ValueError("renderer prompt does not align with vLLM prompt tokens")
752
+
753
+ projected_spans = None
754
+ if message_spans is not None:
755
+ projected_spans = []
756
+ for span in message_spans:
757
+ if span is None:
758
+ projected_spans.append(None)
759
+ continue
760
+ start, end = span
761
+ if not 0 <= start <= end <= len(renderer_prompt_ids):
762
+ raise ValueError("message span exceeds renderer prompt tokens")
763
+ projected_spans.append((offsets[start], offsets[end]))
764
+
765
+ projected_is_content = is_content
766
+ if is_content:
767
+ if len(is_content) != len(renderer_prompt_ids):
768
+ raise ValueError(
769
+ "content attribution does not match renderer prompt tokens"
770
+ )
771
+ projected_is_content = []
772
+ for index, value in enumerate(is_content):
773
+ projected_is_content.extend([value] * (offsets[index + 1] - offsets[index]))
774
+ return projected_spans, projected_is_content
775
+
776
+
774
777
  def _commit_turn(turn: PendingTurn, response: Response) -> int:
775
778
  trace = turn.trace
776
779
  prompt = turn.prompt
777
780
  tokens = response.tokens
778
- multi_modal_data = tokens.multi_modal_data if tokens else None
779
781
  # Constant per renderer, so re-stamping every turn is idempotent.
780
782
  if tokens is not None and tokens.mm_token_type_id_map:
781
783
  trace.mm_token_type_id_map = tokens.mm_token_type_id_map
782
784
  prompt_ids = tokens.prompt_ids if tokens else []
783
- spans = tokens.message_spans if tokens else None
784
- is_content = tokens.is_content if tokens else None
785
- has_is_content = is_content is not None and len(is_content) == len(prompt_ids)
785
+ renderer_prompt_ids = (
786
+ tokens.renderer_prompt_ids
787
+ if tokens and tokens.renderer_prompt_ids is not None
788
+ else prompt_ids
789
+ )
790
+ renderer_spans = tokens.message_spans if tokens else None
791
+ renderer_is_content = tokens.is_content if tokens else None
786
792
  idx = _head_index(trace)
787
793
 
788
- # Token-based prefix reuse. `prepare_turn` matched the prefix by message hash (content); when
789
- # this turn carries token ids, tighten that to token identity — the stored prefix must be an
790
- # exact token prefix of what the model saw this turn (`prompt_ids`). Reuse whole nodes within
791
- # the longest common token prefix and fork at the first divergence, so a retokenized prior
792
- # (BPE drift, dropped `<think>`, rewritten tool calls) branches off with this turn's real
793
- # tokens instead of silently inheriting stale ones. Comparing the *concatenated* prefix (not
794
- # per-message spans) is what makes this correct: a prior assistant's stored generation form
795
- # and its re-rendered input form place the turn-close scaffold in different nodes but at the
796
- # same position, so only a genuine content/token change shifts the common prefix. The bridge
797
- # keeps the prior verbatim so it matches fully (stays linear); the eval relay carries no token
798
- # ids and keeps the message-hash prefix.
799
794
  prefix = turn.prefix_node_ids
800
- path_len = turn.path_len # cumulative stored token length of the reused prefix
795
+ path_len = turn.path_len
796
+ renderer_path_len = turn.renderer_path_len
801
797
  if tokens is not None and prefix:
802
- # Compare node by node against the prompt_ids slice at the running offset (C-level list
803
- # ==, short-circuits at the first divergent node) — no full concatenation materialized.
804
798
  keep = 0
805
799
  off = 0
800
+ renderer_off = 0
806
801
  for nid in prefix:
807
- node_tokens = trace.nodes[nid].token_ids
808
- if prompt_ids[off : off + len(node_tokens)] != node_tokens:
802
+ node = trace.nodes[nid]
803
+ node_renderer_ids = node.logical_ids
804
+ if (
805
+ prompt_ids[off : off + len(node.token_ids)] != node.token_ids
806
+ or renderer_prompt_ids[
807
+ renderer_off : renderer_off + len(node_renderer_ids)
808
+ ]
809
+ != node_renderer_ids
810
+ ):
809
811
  break
810
- off += len(node_tokens)
812
+ off += len(node.token_ids)
813
+ renderer_off += len(node_renderer_ids)
811
814
  keep += 1
815
+ if tokens.bridged and keep != len(prefix):
816
+ raise ValueError(
817
+ "vLLM prompt tokens do not exactly extend the stored rollout prefix"
818
+ )
812
819
  prefix = prefix[:keep]
813
820
  path_len = off
821
+ renderer_path_len = renderer_off
822
+ spans, is_content = _project_prompt_attribution(
823
+ renderer_prompt_ids,
824
+ prompt_ids,
825
+ trace.mm_token_type_id_map,
826
+ tokens.mm_placeholders if tokens else None,
827
+ renderer_spans,
828
+ renderer_is_content,
829
+ )
830
+ has_is_content = is_content is not None and len(is_content) == len(prompt_ids)
814
831
 
815
832
  # A parallel request may have committed more of this prompt after `prepare_turn` resolved
816
833
  # the inference prefix. Reconcile that still-uncommitted tail now, one whole message at a
@@ -822,6 +839,10 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int:
822
839
  i = len(prefix)
823
840
  span = spans[i] if spans and i < len(spans) else None
824
841
  end = span[1] if span else path_len
842
+ renderer_span = (
843
+ renderer_spans[i] if renderer_spans and i < len(renderer_spans) else None
844
+ )
845
+ renderer_end = renderer_span[1] if renderer_span else renderer_path_len
825
846
  parent = prefix[-1] if prefix else None
826
847
  if tokens is not None and span is None:
827
848
  next_start = next(
@@ -832,6 +853,16 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int:
832
853
  ),
833
854
  len(prompt_ids),
834
855
  )
856
+ next_renderer_start = next(
857
+ (
858
+ later_span[0]
859
+ for later_span in (
860
+ renderer_spans[i + 1 :] if renderer_spans else []
861
+ )
862
+ if later_span is not None
863
+ ),
864
+ len(renderer_prompt_ids),
865
+ )
835
866
  existing = _matching_prefix_node(
836
867
  trace,
837
868
  parent,
@@ -839,48 +870,57 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int:
839
870
  prompt_ids,
840
871
  path_len,
841
872
  next_start,
842
- turn.tools,
873
+ renderer_prompt_ids,
874
+ renderer_path_len,
875
+ next_renderer_start,
876
+ tools=turn.tools,
843
877
  )
844
878
  if existing is not None:
845
879
  end += len(trace.nodes[existing].token_ids)
880
+ renderer_end += len(trace.nodes[existing].logical_ids)
846
881
  else:
847
882
  node_tokens = prompt_ids[path_len:end]
883
+ renderer_node_tokens = renderer_prompt_ids[renderer_path_len:renderer_end]
848
884
  existing = _matching_node(
849
885
  trace,
850
886
  parent,
851
887
  prompt[i],
852
888
  node_tokens if tokens is not None else None,
853
- turn.tools,
889
+ renderer_node_tokens if tokens is not None else None,
890
+ tools=turn.tools,
854
891
  )
855
892
  if existing is None:
856
893
  break
857
894
  prefix.append(existing)
858
895
  path_len = end
896
+ renderer_path_len = renderer_end
859
897
 
860
898
  num_reused = len(prefix)
861
899
  parent = prefix[-1] if prefix else None
862
- # cursor: in prompt_ids, the end of the previous *new* message's tokens
863
900
  cursor: int | None = None
901
+ renderer_cursor: int | None = None
864
902
  # Track new nodes separately so routed-expert attribution needs only node ids, not this path.
865
903
  new_node_ids: list[int] = []
866
- # Materialize the reused message path only for multimodal cursor attribution.
867
- mm_path: list[tuple[int, Message]] | None = None
868
- if multi_modal_data is not None:
869
- mm_path = [(nid, prompt[i]) for i, nid in enumerate(prefix)]
870
904
  for i, msg in enumerate(prompt[num_reused:], start=num_reused):
871
905
  key = _node_key(parent, msg, turn.tools)
872
906
  start = path_len if cursor is None else cursor
873
907
  span = spans[i] if spans and i < len(spans) else None
874
908
  end = span[1] if span else start
875
909
  node_tokens = prompt_ids[start:end]
910
+ renderer_start = (
911
+ renderer_path_len if renderer_cursor is None else renderer_cursor
912
+ )
913
+ renderer_span = (
914
+ renderer_spans[i] if renderer_spans and i < len(renderer_spans) else None
915
+ )
916
+ renderer_end = renderer_span[1] if renderer_span else renderer_start
876
917
  trace.nodes.append(
877
- # Every value is already typed framework data; avoid revalidating and copying
878
- # potentially huge token slices a second time.
879
918
  MessageNode.model_construct(
880
919
  parent=parent,
881
920
  tools=turn.tools if parent is None else [],
882
921
  message=msg,
883
922
  token_ids=node_tokens,
923
+ renderer_token_ids=renderer_prompt_ids[renderer_start:renderer_end],
884
924
  mask=[False] * len(node_tokens),
885
925
  is_content=is_content[start:end] if has_is_content else [],
886
926
  )
@@ -888,14 +928,16 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int:
888
928
  parent = len(trace.nodes) - 1
889
929
  idx[key] = parent
890
930
  new_node_ids.append(parent)
891
- if mm_path is not None:
892
- mm_path.append((parent, msg))
893
931
  cursor = end
932
+ renderer_cursor = renderer_end
894
933
 
895
- # Assistant node: trailing scaffold (the generation prompt) + the sampled completion.
896
934
  comp_ids = tokens.completion_ids if tokens else []
897
935
  gen_start = path_len if cursor is None else cursor
898
936
  gen_prompt = prompt_ids[gen_start:]
937
+ renderer_gen_start = (
938
+ renderer_path_len if renderer_cursor is None else renderer_cursor
939
+ )
940
+ renderer_gen_prompt = renderer_prompt_ids[renderer_gen_start:]
899
941
  trace.nodes.append(
900
942
  MessageNode.model_construct(
901
943
  parent=parent,
@@ -903,6 +945,7 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int:
903
945
  message=response.message,
904
946
  sampled=True,
905
947
  token_ids=[*gen_prompt, *comp_ids],
948
+ renderer_token_ids=[*renderer_gen_prompt, *comp_ids],
906
949
  mask=[False] * len(gen_prompt) + [True] * len(comp_ids),
907
950
  is_content=([False] * len(gen_prompt) + [True] * len(comp_ids))
908
951
  if has_is_content
@@ -911,15 +954,10 @@ def _commit_turn(turn: PendingTurn, response: Response) -> int:
911
954
  logprobs=tokens.completion_logprobs if tokens else [],
912
955
  )
913
956
  )
914
- # Register the assistant so the next turn's prompt (which restates it) reuses this node.
915
957
  assistant_id = len(trace.nodes) - 1
916
958
  idx[_node_key(parent, response.message, turn.tools)] = assistant_id
917
959
  new_node_ids.append(assistant_id)
918
960
 
919
- # Attribute this turn's images onto the input nodes that introduced them (by content part).
920
- if mm_path is not None:
921
- _attribute_mm(trace, mm_path, num_reused, multi_modal_data)
922
-
923
961
  # Attribute this turn's expert-routing array onto the nodes created this turn (new input
924
962
  # nodes in creation order, then the assistant node), each getting the routing for its tokens.
925
963
  # The prefix goes in too, so the position the previous turn could only pad can be corrected.
verifiers/v1/trace.py CHANGED
@@ -9,7 +9,6 @@ from typing import TYPE_CHECKING, Any, Generic
9
9
 
10
10
  import numpy as np
11
11
  from pydantic import BaseModel, Field, PrivateAttr, computed_field, field_serializer
12
- from renderers.base import MultiModalData
13
12
  from typing_extensions import TypeVar
14
13
 
15
14
  if TYPE_CHECKING:
@@ -43,7 +42,6 @@ TRACE_VERSION = 1
43
42
  EXCLUDE_FIELDS: dict = {
44
43
  "nodes": {
45
44
  "__all__": {
46
- "multi_modal_data",
47
45
  "routed_experts",
48
46
  "sampling_mask",
49
47
  }
@@ -309,31 +307,14 @@ class Branch(BaseModel):
309
307
  weights.extend(node_weights)
310
308
  return weights
311
309
 
312
- @property
313
- def multi_modal_data(self) -> MultiModalData | None:
314
- """Node image data concatenated in token order for training; never persisted."""
315
- merged = MultiModalData()
316
- found = False
317
- for node in self.nodes:
318
- mmd = node.multi_modal_data
319
- if mmd is None or mmd.is_empty():
320
- continue
321
- found = True
322
- for modality, items in mmd.mm_items.items():
323
- merged.mm_items.setdefault(modality, []).extend(items)
324
- for modality, hashes in mmd.mm_hashes.items():
325
- merged.mm_hashes.setdefault(modality, []).extend(hashes)
326
- return merged if found else None
327
-
328
310
  @property
329
311
  def mm_token_type_ids(self) -> list[int] | None:
330
312
  """Per-token modality markers aligned to `token_ids` (0 = text, 1 = image
331
- placeholder, 2 = video placeholder), driving the trainer's vision-encoder
332
- slicing; None for branches carrying no multimodal data."""
333
- if self.multi_modal_data is None:
313
+ placeholder, 2 = video placeholder); None when none are present."""
314
+ if not self.mm_token_type_id_map:
334
315
  return None
335
- mapping = self.mm_token_type_id_map
336
- return [mapping.get(t, 0) for t in self.token_ids]
316
+ token_types = [self.mm_token_type_id_map.get(t, 0) for t in self.token_ids]
317
+ return token_types if any(token_types) else None
337
318
 
338
319
  @property
339
320
  def routed_experts(self) -> np.ndarray | None:
verifiers/v1/types.py CHANGED
@@ -4,7 +4,6 @@ from typing import Annotated, Any, Literal
4
4
 
5
5
  import numpy as np
6
6
  from pydantic import AliasChoices, BaseModel, ConfigDict, Field
7
- from renderers.base import MultiModalData
8
7
  from typing_extensions import TypedDict
9
8
 
10
9
 
@@ -223,18 +222,24 @@ class TurnTokens(BaseModel):
223
222
  model_config = ConfigDict(arbitrary_types_allowed=True)
224
223
 
225
224
  prompt_ids: list[int] = Field(default_factory=list)
225
+ """Effective prompt IDs evaluated by the model, after multimodal expansion."""
226
226
  completion_ids: list[int] = Field(default_factory=list)
227
227
  completion_logprobs: list[float] = Field(default_factory=list)
228
228
 
229
+ # Transient graph-construction metadata, consumed by the turn's commit and excluded
230
+ # from serialized responses.
231
+ renderer_prompt_ids: list[int] | None = Field(default=None, exclude=True)
232
+ """Logical renderer IDs before multimodal expansion, retained for bridge extension."""
233
+ bridged: bool = Field(default=False, exclude=True)
234
+ """Whether the renderer constructed this prompt by extending a stored prefix."""
229
235
  # Transient carrier (excluded): per-message token spans into `prompt_ids` from the renderer,
230
236
  # consumed by the turn's `commit` to attribute tokens per message, then dropped.
231
237
  message_spans: list[tuple[int, int] | None] | None = Field(
232
238
  default=None, exclude=True
233
239
  )
234
240
  is_content: list[bool] | None = Field(default=None, exclude=True)
235
- # Transient carrier (excluded): the renderer's multimodal sidecar (image tensors + offsets),
236
- # attributed per node by the turn's `commit`, then dropped — never persisted.
237
- multi_modal_data: MultiModalData | None = Field(default=None, exclude=True)
241
+ # Authoritative effective-prompt ranges returned by vLLM, flattened across modalities.
242
+ mm_placeholders: list[tuple[int, int]] | None = Field(default=None, exclude=True)
238
243
  # Transient carrier (excluded): the renderer's special-token id -> modality marker map,
239
244
  # stamped onto `Trace.mm_token_type_id_map` by the turn's `commit`. None unless the
240
245
  # rendering renderer is multimodal.
@@ -158,7 +158,12 @@ async def restore(runtime: Runtime, collected: dict[str, bytes | None]) -> None:
158
158
  # delete content an earlier one just restored. Clearing also drops any file or
159
159
  # symlink the image left at the target.
160
160
  roots = " ".join(shlex.quote(root) for root in collected)
161
- await _run(runtime, f"rm -rf -- {roots}", "clear artifact roots")
161
+ # Container exec requires its configured cwd to exist, including while a
162
+ # submission replaces the entire working directory (or one of its parents).
163
+ workdir = shlex.quote(getattr(runtime.config, "workdir", None) or "/")
164
+ await _run(
165
+ runtime, f"rm -rf -- {roots} && mkdir -p -- {workdir}", "clear artifact roots"
166
+ )
162
167
  for root, archive in collected.items():
163
168
  if archive is None:
164
169
  continue
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: verifiers
3
- Version: 0.3.2.dev149
3
+ Version: 0.3.2.dev151
4
4
  Summary: Verifiers: Environments for LLM Reinforcement Learning
5
5
  Project-URL: Homepage, https://github.com/primeintellect-ai/verifiers
6
6
  Project-URL: Documentation, https://github.com/primeintellect-ai/verifiers
@@ -4,7 +4,7 @@ verifiers/v1/agent.py,sha256=yBCCBe5tVTeQ8OZDlVSZLw_IDvCmEQSnF3wQ1BZcT14,32342
4
4
  verifiers/v1/env.py,sha256=GqrVxxrvv-9WitXhD8Rre0tKQgnQWeD8mTM2gDHMVBs,18040
5
5
  verifiers/v1/episode.py,sha256=UPtA02aGWA-fY8cJKEuYCMtQVnPN1C0BSN13LggNT8U,5986
6
6
  verifiers/v1/errors.py,sha256=QFbKFrma1STrN7abYcU2LBQOQBpj43Rm-nOwL9EFeUM,5743
7
- verifiers/v1/graph.py,sha256=3HtCYL36MDAoiBk7u8yi6FxUBkIqQ7VYn9_IL0NEOog,41559
7
+ verifiers/v1/graph.py,sha256=4EA1TzUBLf9HycfSifKI5J3XkeJzYjbsszZRVwXSGP0,41420
8
8
  verifiers/v1/harness.py,sha256=incFERkwxSPQvBCWc0BVQPg9IAEMTeelBq4USKWLxWo,15669
9
9
  verifiers/v1/judge.py,sha256=uqn-djkK2_7Q3kOjkDo6ipOvV_HOvuLa9y6AoDjRux8,9618
10
10
  verifiers/v1/rollout.py,sha256=WiMCJyAypsR1OaXvzJQx7G-T67v__AnnDNoXotuI_VE,25444
@@ -13,8 +13,8 @@ verifiers/v1/session.py,sha256=qVMgKmqkItemR3uYUDAFuDz7huP593kHQbge8FlRba8,28420
13
13
  verifiers/v1/state.py,sha256=EckF2bWp-vV4b1jYJ9sLI5xrfGuI5spIgYYwW926toI,595
14
14
  verifiers/v1/task.py,sha256=g2MTr1w2TEZXo8g5Z1_SlqqO_mQ9r1VLXrLj4o0Z_VY,12596
15
15
  verifiers/v1/taskset.py,sha256=fp2E0IEhL_Ybj9cegZwljfTmW27p_30zTWHFKw_kXE4,4374
16
- verifiers/v1/trace.py,sha256=8WKgoM0fe9Zdi0m8ffck7PNyi-QF8EzSTV0LM1DqqYc,33029
17
- verifiers/v1/types.py,sha256=xPetvJxksKPYKic389RFnVvy2vN-gD2bPZFZoEyzmp4,10191
16
+ verifiers/v1/trace.py,sha256=tgzkft4nKF4QGeGXwoqkfXjaUbG65Nn1oBFzMUFp-u0,32234
17
+ verifiers/v1/types.py,sha256=-In5plgB3mQLqtRKHC3QDCPDvuruAqVd8WIjKrsGB_I,10578
18
18
  verifiers/v1/acp/__init__.py,sha256=9RySmxFeEMT2XSJy5wYevqEhQdul2jHl9f5XAribG3A,13631
19
19
  verifiers/v1/acp/runner.py,sha256=zPo-2ZXmFMmQchhD7nNzlZUrabG09qiBvpB67KVGm3M,14095
20
20
  verifiers/v1/cli/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
@@ -40,7 +40,7 @@ verifiers/v1/clients/__init__.py,sha256=Ysig0tE_0E4Jsfgfes1XHN-fK1s_RCXdqZD6E7lK
40
40
  verifiers/v1/clients/base.py,sha256=PoDw4GMqrfuPTFVK6n2jDmlD0_lJ5zTrYcNio5yq6K8,1691
41
41
  verifiers/v1/clients/client.py,sha256=zqC_AkiD9pl0kxxIbupNiehbS-YnVdp3EMZ0LOHTd8U,3125
42
42
  verifiers/v1/clients/eval.py,sha256=cJucyUFsbF1JlG8nF0Tlw2EBMSoAF0YMklgYiOSoNpg,8322
43
- verifiers/v1/clients/train.py,sha256=jm_z_Gkwqo4IKEVN8WrDnEyJykXxmswFHzlsnFop7_I,18945
43
+ verifiers/v1/clients/train.py,sha256=lSJTwHjwOET3D_CI_eYZSoXSQcv2H0K5sYJ7kD_kdw4,19489
44
44
  verifiers/v1/configs/__init__.py,sha256=X7u6X7B3ieD1HZl79Tv1kg-yPHbT0QLvAebaGqxVtvA,317
45
45
  verifiers/v1/configs/agent.py,sha256=Vut73R2EW8QuL5c0iTnNFC_g-Obvl8YUKsjlvkABuL0,5311
46
46
  verifiers/v1/configs/client.py,sha256=fPWtBJv5lPkqttwk0AF2ubfLdzxX2XCO0Zr-A9V4mtk,4440
@@ -173,7 +173,7 @@ verifiers/v1/tasksets/textarena/__init__.py,sha256=Os2OlBY_pSH0B1DP32RitumgSLpq2
173
173
  verifiers/v1/tasksets/textarena/taskset.py,sha256=9fa_unlFiFZyuN4I02rW3x4qyRnFWXqwTUZTloQefu0,4433
174
174
  verifiers/v1/utils/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
175
175
  verifiers/v1/utils/aio.py,sha256=0yNFOqB2oO6Qktc6NI_ZSAGuo1aHsIK_wuT1BtmoZ3I,1532
176
- verifiers/v1/utils/artifacts.py,sha256=6XEF6EVuES_zYB_pA2aIWHOvm38XyWXTsLovieA6WYY,9367
176
+ verifiers/v1/utils/artifacts.py,sha256=i9DRhxViMAf-aBPjs3_9Q8RWVRC5pD-_v97VbNPPNcY,9638
177
177
  verifiers/v1/utils/compile.py,sha256=5pxy31wx2Q-3pQjOO5B9SH6OT-wwb8r4F1LAl6HMeQA,5835
178
178
  verifiers/v1/utils/decorators.py,sha256=vhoN6YoKb15-Fkqmc9RaNlUIbt9WRvNK4zAn2bAxfS4,5367
179
179
  verifiers/v1/utils/format.py,sha256=yX160sGewMMZn9iJqTHRbYYt4Z9pHasjMlFT1dXrKD8,2446
@@ -192,8 +192,8 @@ verifiers/v1/utils/scope.py,sha256=bzUEiWDOVdbj7dXOwGOlzsQ_-Xu0j9NR5TalQB-8TS8,8
192
192
  verifiers/v1/utils/score.py,sha256=493yJVMw8teCu9JxapxMFPFzI0hNUqdo0Y2nGW4kckk,6200
193
193
  verifiers/v1/utils/trace_store.py,sha256=RsCDonGR-bs-tvEstcppiJm3jFmAdIiG83fA5KtCbAw,2801
194
194
  verifiers/v1/utils/version.py,sha256=-obEo_-l9-D8FLef4hYxncOe-uJpxrM1g2Hig_37Sgs,1607
195
- verifiers-0.3.2.dev149.dist-info/METADATA,sha256=Pkec0unD3c-vy5nESbIyVvndtWR2NyXBST7fVC700AI,4161
196
- verifiers-0.3.2.dev149.dist-info/WHEEL,sha256=W3fkpkm7-wf9vBI5Z-7s0eWkeM-spu78I8Neb98DeEg,87
197
- verifiers-0.3.2.dev149.dist-info/entry_points.txt,sha256=iugElcdWPKbQM7uFF0lZ8iUpHsNr17-BwEAAjJWxV3U,259
198
- verifiers-0.3.2.dev149.dist-info/licenses/LICENSE,sha256=v0RrUsdV3IDoZhrRce297IXS3xMHNJ-_LdLpFAUWb9k,1072
199
- verifiers-0.3.2.dev149.dist-info/RECORD,,
195
+ verifiers-0.3.2.dev151.dist-info/METADATA,sha256=1VtFlk1XKVSzFTnSb4uqkTEl8Je4hFcEhWUo6h5r3RE,4161
196
+ verifiers-0.3.2.dev151.dist-info/WHEEL,sha256=W3fkpkm7-wf9vBI5Z-7s0eWkeM-spu78I8Neb98DeEg,87
197
+ verifiers-0.3.2.dev151.dist-info/entry_points.txt,sha256=iugElcdWPKbQM7uFF0lZ8iUpHsNr17-BwEAAjJWxV3U,259
198
+ verifiers-0.3.2.dev151.dist-info/licenses/LICENSE,sha256=v0RrUsdV3IDoZhrRce297IXS3xMHNJ-_LdLpFAUWb9k,1072
199
+ verifiers-0.3.2.dev151.dist-info/RECORD,,