From efeb640fcfc04dc55a9fba181faa3ff205f2121d Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 03:19:44 +0000 Subject: [PATCH 1/3] Preserve complete sampled conditioning in exact history recovery --- src/art/trajectories/_tokenize.py | 391 +++++++++++++++++- .../test_exact_source_recovery.py | 278 +++++++++++++ 2 files changed, 650 insertions(+), 19 deletions(-) create mode 100644 tests/unit/trajectories/test_exact_source_recovery.py diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 9b692894c..a8672e59c 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -479,6 +479,55 @@ def _assistant_stop_masks( return trimmed, stop +def _prove_completed_sampled_stop_tail( + rendered: list[int], + completed: list[int] | None, + assistant_mask: list[bool], + bounds: tuple[int, int], + full_exact: list[int], + *, + source: object, + tokenizer: Tokenizer, +) -> int | None: + """Bound a rendered closing tail by this message's completed prefix.""" + if ( + completed is None + or not 0 <= bounds[0] < bounds[1] <= len(completed) <= len(rendered) + or rendered[: len(completed)] != completed + ): + return None + count = _sampled_stop_suffix( + full_exact, + source=source, + source_key=_sampled_source_key(source), + tokenizer=tokenizer, + ) + if ( + not count + or count != _stop_suffix(full_exact, None, tokenizer) + or full_exact[:-count] != rendered[bounds[0] : bounds[1]] + ): + return None + mask, stops = _assistant_stop_masks( + completed, assistant_mask[: len(completed)], tokenizer + ) + if not all(mask[bounds[0] : bounds[1]]): + return None + end = bounds[1] + while end < len(mask) and mask[end]: + end += 1 + tail = completed[bounds[1] : end] + terminators = _terminator_ids(tokenizer) + if ( + end == bounds[1] + or not stops[end - 1] + or tail[-count:] != full_exact[-count:] + or sum(token in terminators for token in tail) != count + ): + return None + return end + + def _translate_token_mask( source: Sequence[int], target: Sequence[int], @@ -4057,6 +4106,7 @@ def _tokenize_exact_projected_chat_history( length_stop_boundaries: Mapping[_SampledSourceKey, _RenderedLengthStopBoundary] | None = None, projection_validated: bool = False, + _native_prompt_context: bool = False, _trace: _TraceBuilder | None = None, ) -> TokenizedHistory | None: if not projection_validated and not _history_matches_projection(history): @@ -4199,6 +4249,24 @@ def _tokenize_exact_projected_chat_history( if boundary is not None and next_prompt is not None: rendered_boundary = [*boundary.tail, *boundary.following] native_boundary = next_prompt[end:] + if ( + _native_prompt_context + and projection_validated is True + and retained_ids == output + and len(output_logprobs) == len(output) + and next_prompt[:end] == [*prompt, *output] + and boundary.tail + and native_boundary[: len(boundary.tail)] == list(boundary.tail) + and native_boundary != rendered_boundary + ): + # The complete request prefix authenticates this non-loss + # context. Keep the entire independently proved closing + # tail at the exact output end; never trim a common prefix. + boundary = _RenderedLengthStopBoundary( + tail=boundary.tail, + following=tuple(native_boundary[len(boundary.tail) :]), + ) + rendered_boundary = [*boundary.tail, *boundary.following] extra = len(native_boundary) - len(rendered_boundary) decode = getattr(tokenizer, "decode", None) if ( @@ -4464,6 +4532,93 @@ def _source_covers_complete_sampled_message( ) == normalize_chat_message(projected[0]) +def _require_exact_chat_source_edges( + history: ChatCompletionsHistory, + tokenized: TokenizedHistory, + trace: _HistoryTokenizationTrace | None, + tokenizer: Tokenizer, +) -> None: + def refuse() -> None: + raise ValueError( + "Exact source boundary retry lacks complete conditioned source proof" + ) + + if ( + trace is None + or tokenized.history is not history + or len(tokenized.tokens) != len(tokenized.logprobs) + or len(tokenized.tokens) != len(tokenized.flags) + ): + refuse() + assert trace is not None + trace.validate(tokenized) + expected = { + _sampled_source_key(source): source + for message, source in zip( + history.messages, history.message_sources, strict=True + ) + if message.get("role") == "assistant" + and source is not None + and _source_is_sampled(source) + } + positions: dict[_SampledSourceKey, list[int]] = {} + for index, key in enumerate(trace.source_keys): + if key is not None: + positions.setdefault(key, []).append(index) + if ( + not expected + or expected.keys() != positions.keys() + or expected.keys() != trace.sources.keys() + ): + refuse() + required = ( + TokenFlag.EXACT | TokenFlag.SAMPLED | TokenFlag.ASSISTANT | TokenFlag.OUTPUT + ) + for key, source in expected.items(): + indices = positions[key] + start, end = indices[0], indices[-1] + 1 + prompt = _chat_source_prompt_tokens(source) + output = _source_output_tokens(source, key) + lp_ids, logprobs = _chat_source_full_tokens(source) + if ( + indices != list(range(start, end)) + or prompt is None + or output is None + or tokenized.tokens[:start] != prompt + or tokenized.tokens[start:end] != output + or lp_ids != output + or len(logprobs) != end - start + or not all( + a == b or (math.isnan(a) and math.isnan(b)) + for a, b in zip(tokenized.logprobs[start:end], logprobs, strict=True) + ) + or any(flag & required != required for flag in tokenized.flags[start:end]) + ): + refuse() + assert output is not None + stop_count = _sampled_stop_suffix( + output, source=source, source_key=key, tokenizer=tokenizer + ) + if ( + any( + bool(tokenized.flags[index] & TokenFlag.STOP) + != (index >= end - stop_count) + for index in indices + ) + or ( + _source_stop_evidence(source, key)[0] == "length" + and any(tokenized.flags[index] & TokenFlag.STOP for index in indices) + ) + or ( + stop_count + and end < len(tokenized.tokens) + and tokenized.flags[end] & TokenFlag.STOP + and not tokenized.flags[end] & TokenFlag.SAMPLED + ) + ): + refuse() + + def _tokenize_chat_view( history: ChatCompletionsHistory, *, @@ -4473,7 +4628,9 @@ def _tokenize_chat_view( chat_template_kwargs: Mapping[str, object] | None, _projection_matches: bool | None = None, _trace: _TraceBuilder | None = None, + _exact_source_boundary_retry: bool = False, ) -> TokenizedHistory: + original_tokenizer = tokenizer _validate_history_sources(history) config = ( _TokenizerConfig(base_model or history.model or "") @@ -4510,6 +4667,12 @@ def _tokenize_chat_view( **default_chat_template_kwargs_for_template(template), **explicit_kwargs, } + if ( + _exact_source_boundary_retry + and "enable_thinking" not in explicit_kwargs + and kwargs.get("enable_thinking") is False + ): + kwargs.pop("enable_thinking") ends_with_assistant = bool(messages) and messages[-1].get("role") == "assistant" segmented = False @@ -4783,7 +4946,8 @@ def source_matches_context(source: object) -> bool: exact_prefix_length = 0 canonical_prefix_length = 0 if ( - chat_template is None + not _exact_source_boundary_retry + and chat_template is None and chat_template_kwargs is None and _projection_matches is True ): @@ -4818,22 +4982,63 @@ def source_matches_context(source: object) -> bool: canonical_assistant_mask, direct_bounds or None, ) - assistant_mask = _translate_token_mask( - canonical_rendered, - rendered, - canonical_assistant_mask, - tokenizer=resolved_tokenizer, - ) - output_mask = _translate_token_mask( - canonical_rendered, - rendered, - canonical_output_mask, - tokenizer=resolved_tokenizer, - ) - stop_mask = _translate_token_mask(canonical_rendered, rendered, canonical_stop_mask) - length_stop_mask = _translate_token_mask( - canonical_rendered, rendered, canonical_length_stop_mask - ) + try: + assistant_mask = _translate_token_mask( + canonical_rendered, + rendered, + canonical_assistant_mask, + tokenizer=resolved_tokenizer, + ) + output_mask = _translate_token_mask( + canonical_rendered, + rendered, + canonical_output_mask, + tokenizer=resolved_tokenizer, + ) + stop_mask = _translate_token_mask( + canonical_rendered, rendered, canonical_stop_mask + ) + length_stop_mask = _translate_token_mask( + canonical_rendered, rendered, canonical_length_stop_mask + ) + except ValueError as error: + origin = error.__traceback__ + while origin is not None and origin.tb_next is not None: + origin = origin.tb_next + if ( + _exact_source_boundary_retry + or original_tokenizer is not None + or _projection_matches is not True + or chat_template is not None + or chat_template_kwargs is not None + or not isinstance(history.model, str) + or not history.model.startswith("wandb-artifact:///") + or not _history_has_length_stop(history) + or type(error) is not ValueError + or origin is None + or origin.tb_frame.f_code is not _translate_token_mask.__code__ + or str(error) + != "Cannot preserve assistant boundaries across exact prompt token replacement" + ): + raise + retry_trace = _TraceBuilder() + exact = _tokenize_chat_view( + history, + base_model=base_model, + tokenizer=original_tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _projection_matches=_projection_matches, + _trace=retry_trace, + _exact_source_boundary_retry=True, + ) + _require_exact_chat_source_edges( + history, exact, retry_trace.trace, resolved_tokenizer + ) + if _trace is not None: + assert retry_trace.trace is not None + _trace.set(exact, retry_trace.trace.source_keys, retry_trace.trace.sources) + return exact positions_by_first_token: dict[int, list[int]] = {} for index, token_id in enumerate(rendered): positions_by_first_token.setdefault(token_id, []).append(index) @@ -5431,12 +5636,30 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: tokenizer=resolved_tokenizer, length_stop_boundaries=length_stop_boundaries, projection_validated=True, + _native_prompt_context=( + _exact_source_boundary_retry + and _projection_matches is True + and chat_template is None + and chat_template_kwargs is None + and all( + source_matches_context(history.message_sources[index]) + and _source_covers_complete_sampled_message( + history.messages[index], history.message_sources[index] + ) + for index in sampled_message_indices + ) + ), _trace=_trace, ) ) ): return exact + if _exact_source_boundary_retry: + raise ValueError( + "Exact source boundary retry lacks a complete renderer boundary proof" + ) + sampled_message_count = sum( message.get("role") == "assistant" and source is not None @@ -5467,7 +5690,10 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: authoritative_prompt = ( source_prompt_tokens(source) if sampled and source is not None else None ) - initial_proven_bounds = marked_bounds.get(message_index) or probed_bounds.get( + initial_marked_bounds = marked_bounds.get(message_index) + completed_message_render: list[int] | None = None + recovered_stop_end: int | None = None + initial_proven_bounds = initial_marked_bounds or probed_bounds.get( message_index ) exact_output_matches: list[tuple[int, int]] | None = None @@ -5519,6 +5745,7 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: if completed is not None else None ) + completed_message_render = rendered_completed corrected_bounds = ( canonical_span_to_rendered(len(prefix), len(completed)) if prefix is not None and completed is not None @@ -5608,6 +5835,38 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: ) generation_start = len(rendered_prompt) sampled_bounds = (generation_start, len(rendered)) + elif ( + initial_marked_bounds is not None + and complete_sampled_message + and full_exact + and len(full_logprobs) == len(full_exact) + and _projection_matches is True + and source_context_matches + and exact_prefix_length == 0 + and encoded_ids == canonical_rendered == rendered + and len(parts) == 1 + and len(full_exact) != len(part_ids(parts[0][1])) + and rendered[initial_marked_bounds[0] : initial_marked_bounds[1]] + == part_ids(parts[0][1]) + and ( + recovered_stop_end := _prove_completed_sampled_stop_tail( + rendered, + completed_message_render, + assistant_mask, + initial_marked_bounds, + full_exact, + source=source, + tokenizer=resolved_tokenizer, + ) + ) + is not None + ): + # Full-context marker offsets still prove this interval when + # a truncated-prefix probe cannot extend its closing markup. + # Downstream replacement uses the complete source's exact + # IDs/logprobs, retaining its content/stop/overlap checks. + sampled_bounds = initial_marked_bounds + content_bounds_proven = True else: raise ValueError( "Could not prove a sampled history message boundary with this " @@ -5745,7 +6004,10 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: multi_generation_response or len(parts) != 1 or parts[0][0] != "content" ): start = generation_start - if _sampled_stop_suffix( + if recovered_stop_end is not None: + # Consume exactly the independently proved current-message tail. + end = recovered_stop_end + elif _sampled_stop_suffix( full_exact, source=source, source_key=_sampled_source_key(source), @@ -5826,6 +6088,45 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: if match[1] <= upper ] if len(bounded_matches) != 1: + if ( + not bounded_matches + and not _exact_source_boundary_retry + and original_tokenizer is None + and _projection_matches is True + and chat_template is None + and chat_template_kwargs is None + and isinstance(history.model, str) + and history.model.startswith("wandb-artifact:///") + and _history_has_length_stop(history) + and complete_sampled_message + and full_exact + and len(full_logprobs) == len(full_exact) + and len(parts) == 1 + and len(local) == len(full_exact) + and upper - lower < len(local) + ): + retry_trace = _TraceBuilder() + exact = _tokenize_chat_view( + history, + base_model=base_model, + tokenizer=original_tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _projection_matches=_projection_matches, + _trace=retry_trace, + _exact_source_boundary_retry=True, + ) + _require_exact_chat_source_edges( + history, exact, retry_trace.trace, resolved_tokenizer + ) + if _trace is not None: + assert retry_trace.trace is not None + _trace.set( + exact, + retry_trace.trace.source_keys, + retry_trace.trace.sources, + ) + return exact raise ValueError( "Could not uniquely locate a sampled history message in " "the rendered history" @@ -6125,6 +6426,58 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: flags[index] |= TokenFlag.EXACT if history.model is None: raise ValueError("History tokenization requires a model") + # The complete source prompt is the conditioning authority. A known + # mismatch must not be silently accepted by ordinary rendered assembly. + checked_prefixes: set[_SampledSourceKey] = set() + for start, source_key in enumerate(source_keys): + if source_key is None or source_key in checked_prefixes: + continue + checked_prefixes.add(source_key) + source = sources[source_key] + source_prompt = source_prompt_tokens(source) + if source_prompt is None or token_ids[:start] == source_prompt: + # Missing evidence is not an alignment proof; preserve the existing + # API path unless an actual full-prompt mismatch is established. + continue + if ( + _exact_source_boundary_retry + or original_tokenizer is not None + or _projection_matches is not True + or chat_template is not None + or chat_template_kwargs is not None + or not isinstance(history.model, str) + or not history.model.startswith("wandb-artifact:///") + or not _history_has_length_stop(history) + or not all( + source_matches_context(history.message_sources[index]) + and _source_covers_complete_sampled_message( + history.messages[index], history.message_sources[index] + ) + for index in sampled_message_indices + ) + ): + raise ValueError( + "Exact source prefix mismatch lacks unchanged source authority" + ) + retry_trace = _TraceBuilder() + exact = _tokenize_chat_view( + history, + base_model=base_model, + tokenizer=original_tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _projection_matches=_projection_matches, + _trace=retry_trace, + _exact_source_boundary_retry=True, + ) + _require_exact_chat_source_edges( + history, exact, retry_trace.trace, resolved_tokenizer + ) + if _trace is not None: + assert retry_trace.trace is not None + _trace.set(exact, retry_trace.trace.source_keys, retry_trace.trace.sources) + return exact + _mark_sampled_stops( token_ids, flags, diff --git a/tests/unit/trajectories/test_exact_source_recovery.py b/tests/unit/trajectories/test_exact_source_recovery.py new file mode 100644 index 000000000..bd776db63 --- /dev/null +++ b/tests/unit/trajectories/test_exact_source_recovery.py @@ -0,0 +1,278 @@ +"""Exercise the recovery proof with public primitive source fixtures. + +Extract the production functions to keep protocol parsing and model imports out +of these focused tests. The existing tokenizer suite covers those integrations. +""" + +from __future__ import annotations + +import ast +from copy import deepcopy +from dataclasses import dataclass +from enum import IntFlag +from functools import lru_cache +import math +from pathlib import Path +from types import SimpleNamespace as NS +from typing import Any + +import pytest + + +@lru_cache +def _source(): + path = Path(__file__).resolve().parents[3] / "src/art/trajectories/_tokenize.py" + return ast.parse(path.read_text()) + + +def _function(name): + return next(node for node in _source().body if getattr(node, "name", None) == name) + + +def _compile(nodes, namespace): + nodes = [ + ast.ImportFrom("__future__", [ast.alias("annotations")], 0), + *deepcopy(nodes), + ] + exec( + compile(ast.fix_missing_locations(ast.Module(nodes, [])), __file__, "exec"), + namespace, + ) + + +def _mask_entry(): + helper = _function("_translate_token_mask") + # Deepest helper code identifies this one refusal, independently of line + # offsets introduced by formatting or other upstream changes. + raises = [n for n in ast.walk(helper) if isinstance(n, ast.Raise)] + assert len(raises) == 1 + assert isinstance(raises[0].exc, ast.Call) + assert isinstance(raises[0].exc.func, ast.Name) + assert raises[0].exc.func.id == "ValueError" + assert ast.literal_eval(raises[0].exc.args[0]) == ( + "Cannot preserve assistant boundaries across exact prompt token replacement" + ) + block = next( + n + for n in _function("_tokenize_chat_view").body + if isinstance(n, ast.Try) + and any( + isinstance(x, ast.Name) and x.id == "_translate_token_mask" + for x in ast.walk(n) + ) + ) + entry = ast.parse("def entry():\n pass").body[0] + assert isinstance(entry, ast.FunctionDef) + entry.body = [block, ast.parse("return assistant_mask").body[0]] + env = {} + _compile([helper, entry], env) + calls = [] + caller = NS(trace="untouched") + + def retry(history, **kwargs): + calls.append(kwargs) + assert kwargs["tokenizer"] is None + assert kwargs["_exact_source_boundary_retry"] is True + kwargs["_trace"].trace = NS(source_keys=[], sources={}) + return "recovered" + + class Trace: + def __init__(self): + self.trace = None + + env.update( + canonical_rendered=[1, 1], + rendered=[2, 1], + canonical_assistant_mask=[False, True], + canonical_output_mask=[False, True], + canonical_stop_mask=[False, False], + canonical_length_stop_mask=[False, False], + _exact_source_boundary_retry=False, + original_tokenizer=None, + _projection_matches=True, + chat_template=None, + chat_template_kwargs=None, + history=NS(model="wandb-artifact:///public/project/model"), + base_model="public", + _history_has_length_stop=lambda _: True, + resolved_tokenizer=None, + _TraceBuilder=Trace, + _tokenize_chat_view=retry, + _require_exact_chat_source_edges=lambda *args: None, + _trace=None, + ) + return env, calls, caller + + +def test_own_mask_refusal_retries_after_source_line_changes(): + env, calls, _ = _mask_entry() + assert env["entry"]() == "recovered" + assert len(calls) == 1 + + +def test_decoder_same_text_error_is_not_a_mask_refusal(): + env, calls, _ = _mask_entry() + error = ValueError( + "Cannot preserve assistant boundaries across exact prompt token replacement" + ) + + class Decoder: + def decode(self, *args, **kwargs): + raise error + + env.update( + canonical_rendered=[1], + rendered=[2], + canonical_assistant_mask=[True], + resolved_tokenizer=Decoder(), + ) + with pytest.raises(ValueError) as caught: + env["entry"]() + assert caught.value is error + assert not calls + + +@pytest.mark.parametrize( + "change", + [ + {"_exact_source_boundary_retry": True}, + {"original_tokenizer": object()}, + {"_projection_matches": False}, + {"chat_template_kwargs": {}}, + ], +) +def test_ineligible_mask_refusal_propagates(change): + env, calls, _ = _mask_entry() + env.update(change) + with pytest.raises(ValueError, match="Cannot preserve assistant"): + env["entry"]() + assert not calls + + +def test_successful_translation_does_not_retry(): + env, calls, _ = _mask_entry() + env["rendered"] = [1, 1] + assert env["entry"]() == [False, True] + assert not calls + + +class Flag(IntFlag): + EXACT = 1 + SAMPLED = 2 + ASSISTANT = 4 + STOP = 8 + OUTPUT = 16 + + +def _builder_fixture(): + env: dict[str, Any] = dict( + __name__=__name__, + dataclass=dataclass, + math=math, + TokenFlag=Flag, + TokenizedHistory=NS, + _history_matches_projection=lambda h: h.projection, + _source_signature=lambda s: s.key if s else None, + _source_is_sampled=lambda s: s.sampled, + _sampled_source_key=lambda s: s.key, + _chat_source_prompt_tokens=lambda s: s.prompt, + _chat_source_full_tokens=lambda s: (s.output, s.lp), + _source_output_tokens=lambda s, k: s.output, + _source_stop_evidence=lambda s, k: (s.kind,), + _sampled_stop_suffix=lambda ids, **kw: int(kw["source"].kind == "stop"), + ) + names = ( + "_RenderedLengthStopBoundary", + "_HistoryTokenizationTrace", + "_TraceBuilder", + "_retained_output_suffix", + "_mark_sampled_stops", + "_require_exact_chat_source_edges", + "_tokenize_exact_projected_chat_history", + ) + _compile([_function(n) for n in names], env) + a = NS( + key="A", + prompt=[10], + output=[20] * 4096, + lp=[-0.1] * 4096, + kind="length", + sampled=True, + ) + b = NS( + key="B", + prompt=[10, *a.output, 90, 91, 70, 71], + output=[30], + lp=[-0.2], + kind="stop", + sampled=True, + ) + history = NS( + model="public", + projection=True, + messages=[{"role": "assistant"}] * 2, + message_sources=[a, b], + ) + boundaries = {"A": env["_RenderedLengthStopBoundary"]((90, 91), (70,))} + return env, history, boundaries + + +def _build(env, history, boundaries, native=True): + trace = env["_TraceBuilder"]() + value = env["_tokenize_exact_projected_chat_history"]( + history, + tokenizer=None, + projection_validated=True, + length_stop_boundaries=boundaries, + _native_prompt_context=native, + _trace=trace, + ) + return value, trace.trace + + +def test_native_context_requires_complete_tail_and_conditioning(): + env, history, boundaries = _builder_fixture() + assert _build(env, history, boundaries, native=False)[0] is None + value, trace = _build(env, history, boundaries) + env["_require_exact_chat_source_edges"](history, value, trace, None) + assert value.tokens == [*history.message_sources[1].prompt, 30] + assert sum(bool(f & Flag.SAMPLED) for f in value.flags) == 4097 + assert not any(f & Flag.STOP for f in value.flags[1:4097]) + assert value.flags[4098] == Flag.EXACT | Flag.STOP + assert all(math.isnan(value.logprobs[i]) for i in range(4097, 4101)) + assert boundaries["A"].following == (70,) + + +@pytest.mark.parametrize("case", ["tail", "prompt", "logprobs", "suffix_only"]) +def test_native_context_refuses_incomplete_authority(case): + env, history, boundaries = _builder_fixture() + a, b = history.message_sources + if case == "tail": + b.prompt[4098] = 92 + elif case == "prompt": + b.prompt[0] = 11 + elif case == "logprobs": + a.lp.pop() + else: + b.prompt.pop(1) + assert _build(env, history, boundaries)[0] is None + + +@pytest.mark.parametrize( + "case", ["prefix", "ids", "logprobs", "stop", "missing_output"] +) +def test_full_source_guard_rejects_changed_training_edges(case): + env, history, boundaries = _builder_fixture() + value, trace = _build(env, history, boundaries) + if case == "prefix": + history.message_sources[0].prompt = [11] + elif case == "ids": + value.tokens[1] = 22 + elif case == "logprobs": + value.logprobs[1] = -0.3 + elif case == "missing_output": + history.message_sources[0].output = None + else: + value.flags[1] |= Flag.STOP + with pytest.raises(ValueError, match="conditioned source proof"): + env["_require_exact_chat_source_edges"](history, value, trace, None) From 8f84a7e7376f4c4c3e531c6fce4f1c97ec591ab2 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 03:38:39 +0000 Subject: [PATCH 2/3] Document exact sampled conditioning and update compatibility tests --- src/art/trajectories/__init__.py | 14 ++++++ tests/unit/trajectories/test_tokenize.py | 64 ++++++++++++++---------- 2 files changed, 51 insertions(+), 27 deletions(-) diff --git a/src/art/trajectories/__init__.py b/src/art/trajectories/__init__.py index 04a9ab825..788eb63fe 100644 --- a/src/art/trajectories/__init__.py +++ b/src/art/trajectories/__init__.py @@ -876,6 +876,20 @@ def tokenize( chat_template: str | None = None, chat_template_kwargs: Mapping[str, object] | None = None, ) -> TokenizedTrajectory | TokenizedMultiHistoryTrajectory: + """Tokenize histories while retaining their sampled-source evidence. + + A recorded sampled logprob belongs to its complete original token + prefix. A known different prefix raises ValueError unless exact-source + reconstruction can preserve every sampled source's conditioning. + Explicit template/tokenizer overrides and text-equivalent reconciliation + do not permit carrying sampled logprobs into different conditioning. + + For divergent captured prompts, use ``multi_history=True`` with the + default ``reconcile_text_equivalent_tokenizations=False`` to retain + separate authoritative histories. This does not recondition samples or + discard their evidence. Missing prompt metadata follows the existing + fallback and is not certified aligned by this known-mismatch check. + """ from ._tokenize import tokenize_trajectory return tokenize_trajectory( diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 0e04a5268..c6f856841 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -5573,7 +5573,7 @@ def apply_chat_template( ).tokens == [10, 20, 11, 30] -def test_chat_prefix_retokenization_splits_unless_reconciled( +def test_chat_prefix_retokenization_splits_and_refuses_sampled_reconciliation( monkeypatch: pytest.MonkeyPatch, ) -> None: first = _chat_exchange([1], [101, 102]) @@ -5620,23 +5620,27 @@ def apply_chat_template( reconcile_text_equivalent_tokenizations=True ) - with pytest.warns(UserWarning, match="preserved the original sampled token IDs"): - tokenized = history.tokenize(base_model="base/model") - - assert tokenized.tokens == [1, 101, 102, 3, 4] - assert tokenized.logprobs[1:3] == [-10.1, -10.2] - assert all(tokenized.flags[index] & tr.TokenFlag.EXACT for index in (1, 2, 4)) + # Equal decoded text does not make [101, 102] the conditioning [500] + # recorded for the second response's sampled logprob. + with ( + pytest.warns(UserWarning, match="preserved the original sampled token IDs"), + pytest.raises(ValueError, match="Exact source prefix mismatch"), + ): + history.tokenize(base_model="base/model") + for value in (trajectory, art.TrajectoryGroup([trajectory])): + with pytest.raises(ValueError, match="Exact source prefix mismatch"): + value.tokenize( + reconcile_text_equivalent_tokenizations=True, + base_model="base/model", + ) - direct = trajectory.tokenize( - reconcile_text_equivalent_tokenizations=True, - base_model="base/model", - ) - grouped = art.TrajectoryGroup([trajectory]).tokenize( - reconcile_text_equivalent_tokenizations=True, - base_model="base/model", - ) - assert direct.tokens == tokenized.tokens - assert grouped.trajectories[0].tokens == tokenized.tokens + separate = trajectory.tokenize(multi_history=True, base_model="base/model") + assert [history.tokens for history in separate.histories] == [ + [1, 101, 102], + [1, 500, 3, 4], + ] + assert separate.histories[0].logprobs[1:] == [-10.1, -10.2] + assert separate.histories[1].logprobs[-1] == -0.4 @pytest.mark.parametrize("length_changing_prompt", (False, True)) @@ -7139,7 +7143,7 @@ def apply_chat_template( assert tokenized.logprobs[-2:] == [-0.7, -0.2] -def test_chat_view_preserves_initial_prompt_and_ignores_later_disagreement() -> None: +def test_chat_view_refuses_later_sampled_prompt_disagreement() -> None: first = _chat_exchange([1], [2]) second = _chat_exchange([9, 8, 7], [3], offset=1) trajectory = art.Trajectory( @@ -7162,9 +7166,15 @@ def apply_chat_template( history = trajectory.chat_completions_history( reconcile_text_equivalent_tokenizations=True ) - tokenized = history.tokenize(tokenizer=Tokenizer()) + with pytest.raises(ValueError, match="Exact source prefix mismatch"): + history.tokenize(tokenizer=Tokenizer()) - assert tokenized.tokens == [1, 2, 7, 3] + separate = trajectory.tokenize(multi_history=True, tokenizer=Tokenizer()) + assert [history.tokens for history in separate.histories] == [ + [1, 2], + [9, 8, 7, 3], + ] + assert [history.logprobs[-1] for history in separate.histories] == [-0.2, -0.3] def test_reasoning_stripped_chat_histories_tokenize_authoritative_views() -> None: @@ -8402,7 +8412,7 @@ def counted(history: tr.History) -> bool: assert calls == 0 -def test_explicit_template_override_rerenders_exact_exchange_scaffold() -> None: +def test_explicit_template_override_refuses_changed_sampled_conditioning() -> None: trajectory = art.Trajectory( exchanges=TrajectoryExchanges(chat_completions=[_chat_exchange([1], [2])]) ) @@ -8429,13 +8439,13 @@ def apply_chat_template( assert add_generation_prompt return [10] - tokenized = trajectory.tokenize( - tokenizer=Tokenizer(), - chat_template="custom", - ) + # The override would attach the logprob sampled after [1] to prefix [10]. + with pytest.raises(ValueError, match="Exact source prefix mismatch"): + trajectory.tokenize(tokenizer=Tokenizer(), chat_template="custom") - assert tokenized.tokens == [10, 2, 30] - assert tokenized.logprobs[1] == -0.2 + original = trajectory.tokenize(tokenizer=Tokenizer()) + assert original.tokens == [1, 2] + assert original.logprobs[1] == -0.2 def test_responses_external_context_requires_or_uses_exact_prompt_tokens() -> None: From 775b2ea04bbe6e66ff37068e2ed5876f2fee3472 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 04:00:55 +0000 Subject: [PATCH 3/3] Add sampled-conditioning migration guidance and safe mismatch coordinates --- docs/features/additional-histories.mdx | 39 ++++++++++++++++++++++++ src/art/trajectories/__init__.py | 6 ++-- src/art/trajectories/_tokenize.py | 18 ++++++++++- tests/unit/trajectories/test_tokenize.py | 7 ++++- 4 files changed, 66 insertions(+), 4 deletions(-) diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 24802ace7..97f905e22 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -134,6 +134,45 @@ trajectory = Trajectory( ) ``` +## Migrating captured histories with different token prefixes + +Tokenization now rejects a known change to the complete token prefix of a +captured sampled response unless exact-source reconstruction preserves its +conditioning. Earlier versions could carry the response's sampled logprobs into +a different rendered prefix. Matching decoded text is not enough: the token +sequence determines the conditioning. Explicit tokenizer or template overrides +and `reconcile_text_equivalent_tokenizations=True` do not waive this check. + +When captured turns have different authoritative prefixes, keep separate +histories where the protocol supports them: + +```python +tokenized = trajectory.tokenize( + multi_history=True, + reconcile_text_equivalent_tokenizations=False, +) +``` + +This preserves separate histories that the protocol already produces, with each +history's original prompt, sampled outputs and logprobs. It does not automatically +split every incompatible captured history. Every resulting history must still +satisfy the conditioning checks; incomplete or inconsistent source evidence is +not made valid. Remove an override that changes sampled conditioning, or supply +matching source evidence. +If you deliberately want to train on edited or approximate text, construct +source-less/manual histories or use the [SFT message format](/fundamentals/sft-training#data-format) +as a separate supervised dataset. Do not carry the old sampled logprobs or source +references into that approximation. This is not an opt-in to recondition +captured RL samples or a new mode that substitutes NaN logprobs. + +The `Exact source prefix mismatch` error reports only structural coordinates: +the source's zero-based message index (or `None` if unavailable), assembled and +recorded prefix lengths, and the first differing token offset. If the prefixes +differ only in length, the offset is the shorter length. It contains no token IDs, message +text, request IDs or prefix hashes. Use those coordinates to locate the source +and rendering override in your own data. Missing prompt metadata still follows +the existing fallback; absence of this error is not proof of alignment. + ## How It Works ### Tokenization Process diff --git a/src/art/trajectories/__init__.py b/src/art/trajectories/__init__.py index 788eb63fe..2c5b9585c 100644 --- a/src/art/trajectories/__init__.py +++ b/src/art/trajectories/__init__.py @@ -886,8 +886,10 @@ def tokenize( For divergent captured prompts, use ``multi_history=True`` with the default ``reconcile_text_equivalent_tokenizations=False`` to retain - separate authoritative histories. This does not recondition samples or - discard their evidence. Missing prompt metadata follows the existing + separate authoritative histories. This preserves existing separate + histories; it does not force a split, and each history must still pass + the conditioning checks. This does not recondition samples or discard + their evidence. Missing prompt metadata follows the existing fallback and is not certified aligned by this known-mismatch check. """ from ._tokenize import tokenize_trajectory diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index a8672e59c..b205fa414 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -6456,8 +6456,24 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: for index in sampled_message_indices ) ): + message_index = next( + (i for i, item in enumerate(history.message_sources) if item is source), + None, + ) + mismatch = next( + ( + i + for i, (a, b) in enumerate(zip(token_ids[:start], source_prompt)) + if a != b + ), + min(start, len(source_prompt)), + ) raise ValueError( - "Exact source prefix mismatch lacks unchanged source authority" + "Exact source prefix mismatch lacks unchanged source authority " + f"(source_message_index={message_index}, " + f"assembled_prefix_tokens={start}, " + f"recorded_prefix_tokens={len(source_prompt)}, " + f"first_mismatch_offset={mismatch})" ) retry_trace = _TraceBuilder() exact = _tokenize_chat_view( diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index c6f856841..e74dc4480 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -8440,8 +8440,13 @@ def apply_chat_template( return [10] # The override would attach the logprob sampled after [1] to prefix [10]. - with pytest.raises(ValueError, match="Exact source prefix mismatch"): + with pytest.raises(ValueError, match="Exact source prefix mismatch") as caught: trajectory.tokenize(tokenizer=Tokenizer(), chat_template="custom") + assert str(caught.value) == ( + "Exact source prefix mismatch lacks unchanged source authority " + "(source_message_index=1, assembled_prefix_tokens=1, " + "recorded_prefix_tokens=1, first_mismatch_offset=0)" + ) original = trajectory.tokenize(tokenizer=Tokenizer()) assert original.tokens == [1, 2]