From 30b68bd4f249d3801b1d13977cac1906c377e829 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 06:58:21 +0000 Subject: [PATCH 01/10] Refine incompatible native chat histories into proved multi-history streams --- docs/features/additional-histories.mdx | 16 +- src/art/trajectories/_tokenize.py | 330 +++++++++++++++-- .../unit/trajectories/test_native_streams.py | 344 ++++++++++++++++++ .../test_native_streams_training.py | 117 ++++++ 4 files changed, 779 insertions(+), 28 deletions(-) create mode 100644 tests/unit/trajectories/test_native_streams.py create mode 100644 tests/unit/trajectories/test_native_streams_training.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 154c9fd90..9f66d5068 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -159,7 +159,21 @@ for unchanged, complete exchange histories. Recorded logprobs belong to those exact conditioned tokens. Chat, Responses, Messages, and Completions keep their existing protocol-specific projection rules; no separate tokenization API or native-representation option is needed. `multi_history=True` preserves the -histories selected by the trajectory, including their order and model selection. +selected source order and model selection. If a complete unchanged Chat history +cannot use the native fast path and its recorded requests are not one continuous +token prefix, ART can refine it into adjacent native streams. Each stream retains +the final request in that run, all sampled responses' original conditioning and +logprobs, and the original trajectory's reward. The existing first-occurrence +mask still owns each sampled prefix edge once across streams. Earlier responses +outside a run are that request's context; their assistant and STOP roles require +the normal rendering proof and do not become new sampled output. This may load a +tokenizer to prove request roles. Unproved request roles still refuse. + +The internal exchange-training path uses the same refinement and one token weight +per original trajectory. `History.tokenize()` and `multi_history=False` keep +their single-linear representation. Explicit overrides, incomplete native +records, and edited contexts do not authorize splitting. Extra streams may +repeat context; there is no implied memory, speed, or numerical-training claim. Complete unchanged Chat histories end at the final recorded response token, including tool responses and responses stopped by a length limit. ART does not diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 61626c6fc..fe8e9991f 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -36,6 +36,7 @@ AnthropicMessageSource, ChatCompletionsExchange, ChatCompletionsHistory, + ChatCompletionsMessageSource, CompletionsExchange, CompletionsSource, CompletionsStringHistory, @@ -7363,6 +7364,7 @@ def _tokenize_history( _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), _projection_validated: bool = False, _copied_context: bool = False, + _allow_native_streams: bool = False, ) -> TokenizedHistory: if isinstance(history, LegacyHistory): if model is None: @@ -7464,6 +7466,8 @@ def _tokenize_history( ) ): return exact + if _allow_native_streams and not override_requires_render: + _request_native_streams(history) return _tokenize_chat_view( history, base_model=base_model, @@ -7560,7 +7564,28 @@ def tokenize_history( _prior: Sequence[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = (), _projection_validated: bool = False, _context_sources: Sequence[object] | None = None, + _native_stream: bool = False, + _allow_native_streams: bool = False, ) -> TokenizedHistory: + if _native_stream: + assert isinstance(history, ChatCompletionsHistory) + _validate_history_sources(history) + builder = _trace or _TraceBuilder() + # The scope owns the final request, including any historical assistant + # turns. Prove their generic role flags; native IDs alone cannot do so. + tokenized = _tokenize_chat_view( + history, + base_model=base_model, + tokenizer=tokenizer, + chat_template=None, + chat_template_kwargs=None, + _projection_matches=True, + _recorded_boundaries=True, + _trace=builder, + _prior=_prior, + ) + _require_native_stream(history, tokenized, builder) + return tokenized copied = ( list(_context_sources) if _context_sources is not None @@ -7598,6 +7623,8 @@ def tokenize_history( source, prompt, output, logprobs, _prior ) ): + if _allow_native_streams and not override: + _request_native_streams(history) raise ValueError( "A copied response suffix requires its complete original sampled occurrence in the selected trajectory" ) @@ -7613,6 +7640,7 @@ def tokenize_history( _prior=_prior, _projection_validated=_projection_validated, _copied_context=bool(copied), + _allow_native_streams=_allow_native_streams, ) if copied: if trace_builder is None or trace_builder.trace is None: @@ -7644,6 +7672,188 @@ def tokenize_history( return tokenized +class _NativeHistoryStreams(Exception): + def __init__(self, histories: Sequence[History | LegacyHistory]) -> None: + self.histories = list(histories) + self.keys = { + id(history): { + _sampled_source_key(source) + for source in cast(ChatCompletionsHistory, history).message_sources + if source is not None and _source_is_sampled(source) + } + for history in histories + } + + +def _request_native_streams(history: History | LegacyHistory) -> None: + streams = _native_history_streams(history) + if len(streams) > 1: + raise _NativeHistoryStreams(streams) + + +def _native_history_streams( + history: History | LegacyHistory, +) -> list[History | LegacyHistory]: + """Split only complete unchanged Chat sources at incompatible native edges. + + A stream uses its last source's actual request. Sources outside that run are + request context, not newly sampled output. Single-history APIs never call + this helper, and unproved scopes retain the original rendering/refusal path. + """ + from ._history import normalize_chat_message + + if not isinstance(history, ChatCompletionsHistory): + return [history] + _validate_history_sources(history) + state = _history_render_state(history) + if state.context_changed or state.projection_matches is not True: + return [history] + rows = [] + seen = set() + for index, (message, source) in enumerate( + zip(history.messages, history.message_sources, strict=True) + ): + if source is None or not _source_is_sampled(source): + continue + if ( + not isinstance(source, ChatCompletionsMessageSource) + or not isinstance(source.exchange, ChatCompletionsExchange) + or not _source_covers_complete_sampled_message(message, source) + ): + return [history] + key = _sampled_source_key(source) + prompt, output, logprobs = _chat_source_record(source) + choice = _chat_choice(source) + recorded_logprobs = ( + choice_completion_logprobs(choice) + if COMPLETION_LOGPROBS_KEY in (choice.model_extra or {}) + else _logprob_values(_chat_logprob_entries(choice)) + ) + if ( + key in seen + or not prompt + or not output + or len(output) != len(logprobs) + or recorded_logprobs is None + or len(recorded_logprobs) != len(output) + or _source_stop_evidence(source, key)[0] not in {"stop", "length"} + ): + return [history] + seen.add(key) + rows.append((index, source, prompt, output)) + runs: list[ + list[tuple[int, ChatCompletionsMessageSource, list[int], list[int]]] + ] = [] + for row in rows: + if runs: + _, _, prompt, output = runs[-1][-1] + body = [*prompt, *output] + if row[2][: len(body)] == body: + runs[-1].append(row) + continue + runs.append([row]) + if len(runs) < 2: + return [history] + streams: list[History | LegacyHistory] = [] + for run in runs: + end, last, _, _ = run[-1] + exchange = last.exchange + assert isinstance(exchange, ChatCompletionsExchange) + request = exchange.request + messages = [ + *( + normalize_chat_message(message) + for message in request.get("messages", []) + ), + normalize_chat_message( + _chat_choice(last).message.model_dump(mode="python", exclude_none=True) + ), + ] + if ( + messages != history.messages[: end + 1] + or history.tools != request.get("tools") + or history.chat_template != request.get("chat_template") + or history.chat_template_kwargs != request.get("chat_template_kwargs") + ): + return [history] + sources = [ + ChatCompletionsMessageSource(exchange=exchange, request_index=index) + for index in range(end) + ] + [last] + for index, source, _, _ in run: + sources[index] = source + streams.append( + history.model_copy( + update={ + "messages": deepcopy(messages), + "message_sources": sources, + "tools": deepcopy(history.tools), + "chat_template_kwargs": deepcopy(history.chat_template_kwargs), + } + ) + ) + return streams + + +def _require_native_stream( + history: History | LegacyHistory, + value: TokenizedHistory, + builder: _TraceBuilder, + expected_keys: set[_SampledSourceKey] | None = None, +) -> None: + """Certify all scoped sources again after any renderer/tokenizer callbacks.""" + assert isinstance(history, ChatCompletionsHistory) + _validate_history_sources(history) + state = _history_render_state(history) + if state.context_changed or state.projection_matches is not True: + raise ValueError("Native stream request context changed") + trace = builder.trace + if trace is None: + raise ValueError("Native stream requires complete sampled source provenance") + trace.validate(value) + sources = [ + source + for source in history.message_sources + if source is not None and _source_is_sampled(source) + ] + keys = {_sampled_source_key(source) for source in sources} + if ( + set(trace.sources) != keys + or expected_keys is not None + and keys != expected_keys + ): + raise ValueError("Native stream sampled source inventory changed") + for source in sources: + prompt, output, logprobs = _chat_source_record(source) + if ( + prompt is None + or output is None + or not _complete_source_is_represented( + source, prompt, output, logprobs, [(value, trace)] + ) + ): + raise ValueError( + "Native stream does not preserve complete original conditioning" + ) + key = _sampled_source_key(source) + suffix = _sampled_stop_suffix( + output, source=source, source_key=key, tokenizer=builder.tokenizer + ) + observed = [ + index + for index in range(len(output)) + if value.flags[len(prompt) + index] & TokenFlag.STOP + ] + if observed != list(range(len(output) - suffix, len(output))): + raise ValueError("Native stream sampled STOP differs from source authority") + prompt, output, _ = _chat_source_record(sources[-1]) + if prompt is None or output is None or value.tokens != [*prompt, *output]: + raise ValueError("Native stream must end at its final recorded output") + # A string stop reason may invoke an encoder while checking its suffix. + if keys != {_sampled_source_key(source) for source in sources}: + raise ValueError("Native stream source changed while proving STOP") + + def _materialize_trajectory( tokenized: TokenizedHistory, trajectory: Trajectory ) -> TokenizedTrajectory: @@ -7716,32 +7926,69 @@ def tokenize_trajectory( raise ValueError( f"Trajectory tokenization requires exactly one history; found {len(histories)}" ) + allow_streams = ( + multi_history + and not reconcile_text_equivalent_tokenizations + and chat_template is None + and chat_template_kwargs is None + ) context_sources = [_partial_native_context(history) for history in histories] track_context = len(histories) > 1 and any(context_sources) prior: list[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = [] tokenized = [] stop_builders: list[_TraceBuilder | None] = [] collect_stops = tokenizer is None and len(histories) > 1 - for history, copied in zip(histories, context_sources, strict=True): - trace = _TraceBuilder() if track_context or collect_stops else None - result = tokenize_history( - history, - model=model if isinstance(history, LegacyHistory) else history.model, - base_model=base_model, - tokenizer=tokenizer, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - _projection_validated=not isinstance(history, LegacyHistory), - _trace=trace, - _prior=prior, - _context_sources=copied, + scopes: list[tuple[History | LegacyHistory, bool]] = [ + (history, False) for history in histories + ] + scoped_keys: dict[int, set[_SampledSourceKey]] = {} + index = 0 + while index < len(scopes): + history, scoped = scopes[index] + trace = ( + _TraceBuilder() + if scoped or allow_streams or track_context or collect_stops + else None ) + try: + result = tokenize_history( + history, + model=model if isinstance(history, LegacyHistory) else history.model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _projection_validated=not isinstance(history, LegacyHistory), + _trace=trace, + _prior=prior, + _native_stream=scoped, + _allow_native_streams=allow_streams and not scoped, + ) + except _NativeHistoryStreams as planned: + leaf = planned.__traceback__ + while leaf is not None and leaf.tb_next is not None: + leaf = leaf.tb_next + if ( + leaf is None + or leaf.tb_frame.f_code is not _request_native_streams.__code__ + ): + raise + scoped_keys.update(planned.keys) + scopes[index : index + 1] = [(stream, True) for stream in planned.histories] + continue tokenized.append(result) stop_builders.append(trace) - if track_context and trace is not None and trace.trace is not None: + if trace is not None and trace.trace is not None: prior.append((result, trace.trace)) - if collect_stops: + index += 1 + if tokenizer is None and len(tokenized) > 1: _complete_resolved_sampled_stops(tokenized, stop_builders) + for (history, scoped), value, builder in zip( + scopes, tokenized, stop_builders, strict=True + ): + if scoped: + assert builder is not None + _require_native_stream(history, value, builder, scoped_keys[id(history)]) if not multi_history: return _materialize_trajectory(tokenized[0], trajectory) return TokenizedMultiHistoryTrajectory( @@ -7765,33 +8012,62 @@ def _tokenize_trajectory_with_trace( if not trajectory.exchanges: raise ValueError("Private exchange tokenization trace requires exchanges") histories = trajectory.histories(model=model) + scopes: list[tuple[History | LegacyHistory, bool]] = [ + (history, False) for history in histories + ] + scoped_keys: dict[int, set[_SampledSourceKey]] = {} tokenized_histories: list[TokenizedHistory] = [] traces: list[_HistoryTokenizationTrace] = [] builders: list[_TraceBuilder] = [] - for history in histories: + index = 0 + while index < len(scopes): + history, scoped = scopes[index] if isinstance(history, LegacyHistory): raise AssertionError( "Exchange trajectories cannot produce legacy histories" ) trace_builder = _TraceBuilder() - tokenized = tokenize_history( - history, - model=history.model, - base_model=base_model, - tokenizer=tokenizer, - chat_template=chat_template, - chat_template_kwargs=chat_template_kwargs, - _trace=trace_builder, - _projection_validated=True, - _prior=list(zip(tokenized_histories, traces, strict=True)), - ) + try: + tokenized = tokenize_history( + history, + model=history.model, + base_model=base_model, + tokenizer=tokenizer, + chat_template=chat_template, + chat_template_kwargs=chat_template_kwargs, + _trace=trace_builder, + _projection_validated=True, + _prior=list(zip(tokenized_histories, traces, strict=True)), + _native_stream=scoped, + _allow_native_streams=not scoped + and chat_template is None + and chat_template_kwargs is None, + ) + except _NativeHistoryStreams as planned: + leaf = planned.__traceback__ + while leaf is not None and leaf.tb_next is not None: + leaf = leaf.tb_next + if ( + leaf is None + or leaf.tb_frame.f_code is not _request_native_streams.__code__ + ): + raise + scoped_keys.update(planned.keys) + scopes[index : index + 1] = [(stream, True) for stream in planned.histories] + continue if trace_builder.trace is None: raise AssertionError("Exchange tokenization did not produce a source trace") tokenized_histories.append(tokenized) traces.append(trace_builder.trace) builders.append(trace_builder) + index += 1 if tokenizer is None: _complete_resolved_sampled_stops(tokenized_histories, builders) + for (history, scoped), value, builder in zip( + scopes, tokenized_histories, builders, strict=True + ): + if scoped: + _require_native_stream(history, value, builder, scoped_keys[id(history)]) return ( TokenizedMultiHistoryTrajectory( trajectory=trajectory, diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py new file mode 100644 index 000000000..c21227393 --- /dev/null +++ b/tests/unit/trajectories/test_native_streams.py @@ -0,0 +1,344 @@ +from __future__ import annotations + +from copy import deepcopy +import math +import pickle +import struct +from typing import Any, cast + +from openai.types.chat import ChatCompletion +import pytest +from test_tokenize import _CharacterTemplateTokenizer, _chat_exchange + +import art.trajectories as tr +from art.trajectories import _tokenize as module + + +def record(exchange: tr.ChatCompletionsExchange) -> Any: + return cast(Any, exchange.response.choices[0]) + + +def example() -> tuple[tr.Trajectory, Any]: + tokenizer = _CharacterTemplateTokenizer() + first = _chat_exchange(tokenizer._encode("turn 0"), tokenizer._encode("rawanswer§")) + record(first).finish_reason = "length" + second = _chat_exchange( + tokenizer._encode("turn 0answer§turn 1"), tokenizer._encode("answer§"), offset=1 + ) + third = _chat_exchange( + tokenizer._encode("turn 0answer§turn 1answer§turn 2"), + tokenizer._encode("raw terminal tool output"), + offset=2, + ) + payload = third.response.model_dump(mode="python") + payload["choices"][0]["finish_reason"] = "length" + payload["choices"][0]["message"] = { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "public", + "type": "function", + "function": {"name": "lookup", "arguments": "{}"}, + } + ], + } + third.response = ChatCompletion.model_validate(payload) + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second, third]), + reward=3.0, + ) + return trajectory, tokenizer + + +def test_nonnested_streams_cover_native_sources_and_roles() -> None: + trajectory, tokenizer = example() + before = trajectory.model_dump_json() + histories = trajectory.histories() + assert len(histories) == 2 + result, traces = module._tokenize_trajectory_with_trace( + trajectory, tokenizer=tokenizer + ) + assert result.trajectory is trajectory + assert result.reward == trajectory.reward + assert len(result.histories) == 3 + sources = trajectory.exchanges.chat_completions + expected = [ + record(sources[0]), + record(sources[0]), + record(sources[-1]), + ] + for history, choice in zip(result.histories, expected, strict=True): + assert history.tokens == [*choice.prompt_token_ids, *choice.token_ids] + final = result.histories[-1] + assert final.flags[len("turn 0")] & tr.TokenFlag.ASSISTANT + assert not final.flags[len("turn 0")] & (tr.TokenFlag.SAMPLED | tr.TokenFlag.OUTPUT) + assert final.flags[len("turn 0answer")] & tr.TokenFlag.STOP + for exchange in sources: + choice = record(exchange) + matches = [ + (value, trace) + for value, trace in zip(result.histories, traces, strict=True) + if any( + getattr(source, "exchange", None) is exchange + for source in trace.sources.values() + ) + ] + assert matches + for value, trace in matches: + source = next( + source + for source in trace.sources.values() + if getattr(source, "exchange", None) is exchange + ) + assert module._complete_source_is_represented( + source, + choice.prompt_token_ids, + choice.token_ids, + [entry.logprob for entry in choice.logprobs.content], + [(value, trace)], + ) + assert trajectory.model_dump_json() == before + + +def test_old_single_history_refusal_is_retained() -> None: + trajectory, tokenizer = example() + history = trajectory.histories()[-1] + assert isinstance(history, tr.ChatCompletionsHistory) + with pytest.raises(ValueError): + history.tokenize(tokenizer=tokenizer) + + +def test_explicit_rendering_does_not_split(monkeypatch: pytest.MonkeyPatch) -> None: + trajectory, tokenizer = example() + + def forbidden(*args: Any, **kwargs: Any) -> Any: + pytest.fail("explicit rendering cannot split native streams") + + monkeypatch.setattr(module, "_native_history_streams", forbidden) + with pytest.raises(ValueError): + trajectory.tokenize( + tokenizer=tokenizer, multi_history=True, chat_template="override" + ) + + +def test_public_regression_requires_new_streams( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory, tokenizer = example() + monkeypatch.setattr(module, "_native_history_streams", lambda history: [history]) + with pytest.raises(ValueError, match="sampled content boundary"): + trajectory.tokenize(tokenizer=tokenizer, multi_history=True) + + +def finite32(value: float) -> bool: + try: + return math.isfinite(struct.unpack("!f", struct.pack("!f", value))[0]) + except OverflowError: + return False + + +def native_terms(trajectory: tr.Trajectory) -> list[tuple[Any, float]]: + """Independent prefix-edge oracle: claim BEFORE finite filtering.""" + seen = set() + terms = [] + for history in trajectory.histories(): + assert isinstance(history, tr.ChatCompletionsHistory) + for source in history.message_sources: + if source is None or source.choice_index is None: + continue + assert isinstance(source.exchange, tr.ChatCompletionsExchange) + choice: Any = next( + c + for c in source.exchange.response.choices + if c.index == source.choice_index + ) + for index, entry in enumerate(choice.logprobs.content): + key = ( + history.model, + tuple([*choice.prompt_token_ids, *choice.token_ids[: index + 1]]), + ) + if key not in seen: + seen.add(key) + if finite32(entry.logprob): + terms.append((key, entry.logprob)) + return terms + + +def result_terms(result: Any) -> list[tuple[Any, float]]: + terms = [] + for history, mask in zip( + result.histories, + tr.first_occurrence_masks(result.histories, where=tr.TokenFlag.SAMPLED), + strict=True, + ): + for index, (selected, logprob) in enumerate( + zip(mask, history.logprobs, strict=True) + ): + if selected and finite32(logprob): + terms.append( + ((history.model, tuple(history.tokens[: index + 1])), logprob) + ) + return terms + + +@pytest.mark.parametrize("first_logprob", [-0.4, math.nan, 1e100]) +def test_ordered_native_objective_and_compact_roundtrip(first_logprob: float) -> None: + trajectory, tokenizer = example() + first = trajectory.exchanges.chat_completions[0] + record(first).logprobs.content[0].logprob = first_logprob + later = _chat_exchange( + list(record(first).prompt_token_ids), + list(record(first).token_ids), + offset=4, + ) + later.request["messages"] = [{"role": "user", "content": "other branch"}] + record(later).finish_reason = "length" + record(later).logprobs.content[0].logprob = -0.1 + other_model = _chat_exchange([1], [2], model="other/model", offset=5) + other_model.request["messages"] = [{"role": "user", "content": "other model"}] + other_model.response.choices[0].finish_reason = "length" + trajectory.exchanges.chat_completions.extend([later, other_model]) + before = trajectory.model_dump_json() + result = trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + assert result_terms(result) == native_terms(trajectory) + for restored in [ + pickle.loads(pickle.dumps(result)), + tr.compact_validate( + tr.compact_dump(result), type=tr.TokenizedMultiHistoryTrajectory + ), + ]: + assert result_terms(restored) == native_terms(trajectory) + assert restored.model_dump_json() == result.model_dump_json() + assert trajectory.model_dump_json() == before + + +@pytest.mark.parametrize( + "change", + [ + "missing_ids", + "missing_lp", + "edited", + "context", + "unsupported_finish", + "reordered", + ], +) +def test_incomplete_or_edited_history_does_not_authorize_splitting(change: str) -> None: + trajectory, _ = example() + history = trajectory.histories()[-1] + assert isinstance(history, tr.ChatCompletionsHistory) + if change == "missing_ids": + record(trajectory.exchanges.chat_completions[1]).model_extra.pop( + "prompt_token_ids" + ) + elif change == "missing_lp": + trajectory.exchanges.chat_completions[1].response.choices[0].logprobs = None + elif change == "edited": + history.messages[-1]["content"] = "edited response" + history.message_sources[-1] = None + elif change == "context": + history.chat_template_kwargs = {"enable_thinking": False} + elif change == "unsupported_finish": + trajectory.exchanges.chat_completions[1].response.choices[ + 0 + ].finish_reason = "content_filter" + else: + history.messages[1], history.messages[3] = ( + history.messages[3], + history.messages[1], + ) + history.message_sources[1], history.message_sources[3] = ( + history.message_sources[3], + history.message_sources[1], + ) + assert module._native_history_streams(history) == [history] + + +@pytest.mark.parametrize( + "change", ["missing_source", "prompt", "lp", "stop", "extra_stop", "flags"] +) +def test_scoped_final_guard_rejects_corrupt_results(change: str) -> None: + trajectory, tokenizer = example() + result, traces = module._tokenize_trajectory_with_trace( + trajectory, tokenizer=tokenizer + ) + value, trace = result.histories[-1], traces[-1] + builder = module._TraceBuilder(trace=trace, tokenizer=tokenizer) + sampled = [i for i, flag in enumerate(value.flags) if flag & tr.TokenFlag.SAMPLED] + if change == "missing_source": + trace.sources.pop(next(iter(trace.sources))) + elif change == "prompt": + value.tokens[0] += 1 + elif change == "lp": + value.logprobs[sampled[0]] += 1 + elif change == "stop": + stop = next(i for i in sampled if value.flags[i] & tr.TokenFlag.STOP) + value.flags[stop] &= ~tr.TokenFlag.STOP + elif change == "extra_stop": + value.flags[sampled[-1]] |= tr.TokenFlag.STOP + else: + value.flags[sampled[0]] &= ~tr.TokenFlag.SAMPLED + with pytest.raises((ValueError, AssertionError)): + module._require_native_stream(value.history, value, builder) + + +def test_foreign_planning_exception_is_not_swallowed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory, tokenizer = example() + failure = module._NativeHistoryStreams(trajectory.histories()) + + def fail(*args: Any, **kwargs: Any) -> Any: + raise failure + + monkeypatch.setattr(tokenizer, "apply_chat_template", fail) + with pytest.raises(module._NativeHistoryStreams) as caught: + trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + assert caught.value is failure + + +def test_callback_mutation_of_earlier_source_invalidates_scoped_output( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory, tokenizer = example() + original = tokenizer.apply_chat_template + count = 0 + + def mutate(messages: Any, **kwargs: Any) -> Any: + nonlocal count + count += 1 + if len(messages) > 3: + record(trajectory.exchanges.chat_completions[0]).logprobs.content[ + 0 + ].logprob = -99.0 + return original(messages, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", mutate) + with pytest.raises(ValueError, match="inventory changed"): + trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + assert count + + +def test_stop_encoder_cannot_change_already_checked_source( + monkeypatch: pytest.MonkeyPatch, +) -> None: + trajectory, tokenizer = example() + second = trajectory.exchanges.chat_completions[1] + record(second).model_extra["stop_reason"] = "§" + result, traces = module._tokenize_trajectory_with_trace( + trajectory, tokenizer=tokenizer + ) + value, trace = result.histories[-1], traces[-1] + original = tokenizer.__class__.__call__ + + def mutate(self: Any, text: str, **kwargs: Any) -> Any: + if text == "§": + record(second).logprobs.content[0].logprob = -123.0 + return original(self, text, **kwargs) + + monkeypatch.setattr(tokenizer.__class__, "__call__", mutate) + with pytest.raises(ValueError, match="changed while proving STOP"): + module._require_native_stream( + value.history, value, module._TraceBuilder(trace=trace, tokenizer=tokenizer) + ) diff --git a/tests/unit/trajectories/test_native_streams_training.py b/tests/unit/trajectories/test_native_streams_training.py new file mode 100644 index 000000000..ed6ad9f88 --- /dev/null +++ b/tests/unit/trajectories/test_native_streams_training.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +from collections import Counter +from typing import Any + +import pytest +from test_native_streams import example, native_terms +from test_tokenize import _chat_exchange + +from art.preprocessing.pack import packed_tensors_from_tokenized_results +from art.preprocessing.tokenize import tokenize_trajectory_groups +import art.trajectories as tr + + +def test_training_weights_and_packing_keep_each_native_causal_path() -> None: + trajectory, tokenizer = example() + tokenizer.name_or_path = "public/base" + control = tr.Trajectory( + reward=1.0, + exchanges=tr.TrajectoryExchanges( + chat_completions=[_chat_exchange([1, 2], [3])] + ), + ) + group = tr.TrajectoryGroup(trajectories=[trajectory, control]) + before = group.model_dump_json() + results = list( + tokenize_trajectory_groups( + tokenizer, + [group], + allow_training_without_logprobs=False, + scale_rewards=False, + shuffle_group_trajectories=False, + ) + ) + assert len(results) == 3 + terms: list[tuple[tuple[int, ...], float, float, float]] = [] + for original, advantage in [(trajectory, 1.0), (control, -1.0)]: + selected = [r for r in results if r.trajectory is original] + assert selected + expected = native_terms(original) + actual = [] + for result in selected: + assert result.advantage == advantage + assert result.weight == pytest.approx(1 / (len(expected) + 1e-6)) + for index, sampled in enumerate(result.assistant_mask): + if sampled: + prefix = tuple(result.token_ids[: index + 1]) + actual.append((("test/model", prefix), result.logprobs[index])) + terms.append( + (prefix, result.logprobs[index], advantage, result.weight) + ) + assert actual == expected + packed = packed_tensors_from_tokenized_results( + results, + seq_len=256, + truncate_long_results=False, + verbosity=0, + min_prefix_tree_shared_segment_length=1, + ) + # Recover causal ancestors from the actual packed group/parent structure; + # a sibling's tokens must never appear in another source's conditioning. + observed = [] + for row in range(packed["tokens"].shape[0]): + tokens = packed["tokens"][row].tolist() + groups = packed["group_ids"][row].tolist() + parents = packed["parent_ids"][row].tolist() + positions = packed["input_pos"][row].tolist() + for index, sampled in enumerate(packed["assistant_mask"][row].tolist()): + if not sampled: + continue + ancestry = {groups[index]} + current = groups[index] + while True: + parent = parents[groups.index(current)] + if parent in ancestry: + break + ancestry.add(parent) + current = parent + indices = [j for j in range(index + 1) if groups[j] in ancestry] + assert [positions[j] for j in indices] == list(range(positions[index] + 1)) + observed.append( + ( + tuple(tokens[j] for j in indices), + packed["logprobs"][row, index].item(), + packed["advantages"][row, index].item(), + packed["weights"][row, index].item(), + ) + ) + mean_weight = sum(t[3] for t in terms) / len(terms) + advantage_scale = sum(abs(t[2]) * t[3] / mean_weight for t in terms) / len(terms) + + def rounded(values: Any) -> Counter[Any]: + return Counter( + (p, round(lp, 5), round(a, 5), round(w, 5)) for p, lp, a, w in values + ) + + assert rounded(observed) == rounded( + (p, lp, a / advantage_scale, w / mean_weight) for p, lp, a, w in terms + ) + assert group.model_dump_json() == before + + +def test_public_async_dispatch_and_private_trace_agree() -> None: + import asyncio + + from art.trajectories import _tokenize as module + + trajectory, tokenizer = example() + + async def run() -> Any: + return ( + await tr.tokenize([trajectory], multi_history=True, tokenizer=tokenizer) + )[0] + + public = asyncio.run(run()) + private, _ = module._tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) + assert public.model_dump_json() == private.model_dump_json() From fb0b7f450815572f4a5dcf2d40373c28a691ca2c Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 07:01:18 +0000 Subject: [PATCH 02/10] Bind native stream request context across STOP callbacks --- src/art/trajectories/_tokenize.py | 76 ++++++++++++++++--- .../unit/trajectories/test_native_streams.py | 42 ++++++++++ 2 files changed, 107 insertions(+), 11 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index fe8e9991f..8d2084ec7 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -7683,6 +7683,35 @@ def __init__(self, histories: Sequence[History | LegacyHistory]) -> None: } for history in histories } + self.contexts = { + id(history): _native_stream_context(cast(ChatCompletionsHistory, history)) + for history in histories + } + + +def _native_stream_context(history: ChatCompletionsHistory) -> object: + """Keep request order, original source identities and scoped role inputs.""" + return _render_context_key( + [ + history.model, + history.messages, + history.tools, + history.chat_template, + history.chat_template_kwargs, + [ + None + if source is None + else [ + id(type(source)), + id(source.exchange), + source.request_index, + source.choice_index, + ] + for source in history.message_sources + ], + [dict(exchange.request) for exchange in _unique_exchanges(history)], + ] + ) def _request_native_streams(history: History | LegacyHistory) -> None: @@ -7782,16 +7811,19 @@ def _native_history_streams( ] + [last] for index, source, _, _ in run: sources[index] = source - streams.append( - history.model_copy( - update={ - "messages": deepcopy(messages), - "message_sources": sources, - "tools": deepcopy(history.tools), - "chat_template_kwargs": deepcopy(history.chat_template_kwargs), - } - ) + scoped = history.model_copy( + update={ + "messages": deepcopy(messages), + "message_sources": sources, + "tools": deepcopy(history.tools), + "chat_template_kwargs": deepcopy(history.chat_template_kwargs), + } ) + try: + _native_stream_context(scoped) + except (TypeError, RecursionError): + return [history] + streams.append(scoped) return streams @@ -7800,6 +7832,7 @@ def _require_native_stream( value: TokenizedHistory, builder: _TraceBuilder, expected_keys: set[_SampledSourceKey] | None = None, + expected_context: object = None, ) -> None: """Certify all scoped sources again after any renderer/tokenizer callbacks.""" assert isinstance(history, ChatCompletionsHistory) @@ -7807,6 +7840,9 @@ def _require_native_stream( state = _history_render_state(history) if state.context_changed or state.projection_matches is not True: raise ValueError("Native stream request context changed") + context = _native_stream_context(history) + if expected_context is not None and context != expected_context: + raise ValueError("Native stream request context changed") trace = builder.trace if trace is None: raise ValueError("Native stream requires complete sampled source provenance") @@ -7852,6 +7888,8 @@ def _require_native_stream( # A string stop reason may invoke an encoder while checking its suffix. if keys != {_sampled_source_key(source) for source in sources}: raise ValueError("Native stream source changed while proving STOP") + if context != _native_stream_context(history): + raise ValueError("Native stream context changed while proving STOP") def _materialize_trajectory( @@ -7942,6 +7980,7 @@ def tokenize_trajectory( (history, False) for history in histories ] scoped_keys: dict[int, set[_SampledSourceKey]] = {} + scoped_contexts: dict[int, object] = {} index = 0 while index < len(scopes): history, scoped = scopes[index] @@ -7974,6 +8013,7 @@ def tokenize_trajectory( ): raise scoped_keys.update(planned.keys) + scoped_contexts.update(planned.contexts) scopes[index : index + 1] = [(stream, True) for stream in planned.histories] continue tokenized.append(result) @@ -7988,7 +8028,13 @@ def tokenize_trajectory( ): if scoped: assert builder is not None - _require_native_stream(history, value, builder, scoped_keys[id(history)]) + _require_native_stream( + history, + value, + builder, + scoped_keys[id(history)], + scoped_contexts[id(history)], + ) if not multi_history: return _materialize_trajectory(tokenized[0], trajectory) return TokenizedMultiHistoryTrajectory( @@ -8016,6 +8062,7 @@ def _tokenize_trajectory_with_trace( (history, False) for history in histories ] scoped_keys: dict[int, set[_SampledSourceKey]] = {} + scoped_contexts: dict[int, object] = {} tokenized_histories: list[TokenizedHistory] = [] traces: list[_HistoryTokenizationTrace] = [] builders: list[_TraceBuilder] = [] @@ -8053,6 +8100,7 @@ def _tokenize_trajectory_with_trace( ): raise scoped_keys.update(planned.keys) + scoped_contexts.update(planned.contexts) scopes[index : index + 1] = [(stream, True) for stream in planned.histories] continue if trace_builder.trace is None: @@ -8067,7 +8115,13 @@ def _tokenize_trajectory_with_trace( scopes, tokenized_histories, builders, strict=True ): if scoped: - _require_native_stream(history, value, builder, scoped_keys[id(history)]) + _require_native_stream( + history, + value, + builder, + scoped_keys[id(history)], + scoped_contexts[id(history)], + ) return ( TokenizedMultiHistoryTrajectory( trajectory=trajectory, diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index c21227393..6aeed7620 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -342,3 +342,45 @@ def mutate(self: Any, text: str, **kwargs: Any) -> Any: module._require_native_stream( value.history, value, module._TraceBuilder(trace=trace, tokenizer=tokenizer) ) + + +@pytest.mark.parametrize( + "change", ["request", "request_order", "scoped_context", "source_identity"] +) +def test_stop_encoder_cannot_change_role_proof_context( + monkeypatch: pytest.MonkeyPatch, change: str +) -> None: + trajectory, tokenizer = example() + second = trajectory.exchanges.chat_completions[1] + record(second).model_extra["stop_reason"] = "§" + result, traces = module._tokenize_trajectory_with_trace( + trajectory, tokenizer=tokenizer + ) + value, trace = result.histories[-1], traces[-1] + history = value.history + assert isinstance(history, tr.ChatCompletionsHistory) + original = tokenizer.__class__.__call__ + + def mutate(self: Any, text: str, **kwargs: Any) -> Any: + if text == "§": + if change == "request": + second.request["chat_template_kwargs"] = {"changed": True} + elif change == "request_order": + cast(Any, second.request["messages"])[0] = dict( + reversed(list(second.request["messages"][0].items())) + ) + elif change == "scoped_context": + history.chat_template_kwargs = {"changed": True} + else: + source = history.message_sources[3] + assert source is not None + history.message_sources[3] = source.model_copy( + update={"exchange": second.model_copy(deep=True)} + ) + return original(self, text, **kwargs) + + monkeypatch.setattr(tokenizer.__class__, "__call__", mutate) + with pytest.raises(ValueError, match="context changed while proving STOP"): + module._require_native_stream( + history, value, module._TraceBuilder(trace=trace, tokenizer=tokenizer) + ) From ef3742a86b92c54a5df14b9fa928ef611424aceb Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 07:05:56 +0000 Subject: [PATCH 03/10] Recheck all native stream sources after STOP callbacks --- src/art/trajectories/_tokenize.py | 24 +++++++++ .../unit/trajectories/test_native_streams.py | 50 +++++++++++++++++++ 2 files changed, 74 insertions(+) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 8d2084ec7..48fce7c63 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -7706,6 +7706,9 @@ def _native_stream_context(history: ChatCompletionsHistory) -> object: id(source.exchange), source.request_index, source.choice_index, + list(_source_stop_evidence(source, _sampled_source_key(source))) + if _source_is_sampled(source) + else None, ] for source in history.message_sources ], @@ -7892,6 +7895,25 @@ def _require_native_stream( raise ValueError("Native stream context changed while proving STOP") +def _require_native_streams_unchanged( + scopes: Sequence[tuple[History | LegacyHistory, bool]], + keys: Mapping[int, set[_SampledSourceKey]], + contexts: Mapping[int, object], +) -> None: + # A later scope's STOP encoder can change an earlier, already checked source. + # This final barrier invokes no tokenizer/renderer callbacks. + for history, scoped in scopes: + if not scoped: + continue + assert isinstance(history, ChatCompletionsHistory) + if keys[id(history)] != { + _sampled_source_key(source) + for source in history.message_sources + if source is not None and _source_is_sampled(source) + } or contexts[id(history)] != _native_stream_context(history): + raise ValueError("Native stream changed during final STOP validation") + + def _materialize_trajectory( tokenized: TokenizedHistory, trajectory: Trajectory ) -> TokenizedTrajectory: @@ -8035,6 +8057,7 @@ def tokenize_trajectory( scoped_keys[id(history)], scoped_contexts[id(history)], ) + _require_native_streams_unchanged(scopes, scoped_keys, scoped_contexts) if not multi_history: return _materialize_trajectory(tokenized[0], trajectory) return TokenizedMultiHistoryTrajectory( @@ -8122,6 +8145,7 @@ def _tokenize_trajectory_with_trace( scoped_keys[id(history)], scoped_contexts[id(history)], ) + _require_native_streams_unchanged(scopes, scoped_keys, scoped_contexts) return ( TokenizedMultiHistoryTrajectory( trajectory=trajectory, diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index 6aeed7620..1753f9960 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -384,3 +384,53 @@ def mutate(self: Any, text: str, **kwargs: Any) -> Any: module._require_native_stream( history, value, module._TraceBuilder(trace=trace, tokenizer=tokenizer) ) + + +@pytest.mark.parametrize("private", [False, True]) +@pytest.mark.parametrize("change", ["request", "logprob", "stop_reason"]) +def test_later_stop_encoder_cannot_change_an_already_certified_stream( + monkeypatch: pytest.MonkeyPatch, private: bool, change: str +) -> None: + trajectory, tokenizer = example() + first, second, _ = trajectory.exchanges.chat_completions + record(first).finish_reason = "stop" + record(first).model_extra["stop_reason"] = "§" + record(second).model_extra["stop_reason"] = "§" + original_guard = module._require_native_stream + original_encode = tokenizer.__class__.__call__ + armed = False + changed = False + + def guard(history: Any, *args: Any, **kwargs: Any) -> None: + nonlocal armed + # Final scope validation supplies planning snapshots; earlier assembly + # checks do not. Arm only the later stream's actual STOP encoder. + armed = len(args) >= 4 and any( + source is not None and source.exchange is second + for source in history.message_sources + ) + try: + original_guard(history, *args, **kwargs) + finally: + armed = False + + def encode(self: Any, text: str, **kwargs: Any) -> Any: + nonlocal changed + if armed and text == "§": + changed = True + if change == "request": + first.request["chat_template_kwargs"] = {"changed": True} + elif change == "logprob": + record(first).logprobs.content[0].logprob = -123.0 + else: + record(first).model_extra["stop_reason"] = "!" + return original_encode(self, text, **kwargs) + + monkeypatch.setattr(module, "_require_native_stream", guard) + monkeypatch.setattr(tokenizer.__class__, "__call__", encode) + with pytest.raises(ValueError, match="during final STOP validation"): + if private: + module._tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) + else: + trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + assert changed From 13db8a151881bef44ddd97abdfd9eacd675aaee6 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 07:08:27 +0000 Subject: [PATCH 04/10] Certify exact sampled occupancy across native streams --- src/art/trajectories/_tokenize.py | 7 +++++++ tests/unit/trajectories/test_native_streams.py | 8 +++++++- 2 files changed, 14 insertions(+), 1 deletion(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 48fce7c63..ebf81b88a 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -7862,6 +7862,7 @@ def _require_native_stream( and keys != expected_keys ): raise ValueError("Native stream sampled source inventory changed") + expected: list[_SampledSourceKey | None] = [None] * len(value.tokens) for source in sources: prompt, output, logprobs = _chat_source_record(source) if ( @@ -7875,6 +7876,10 @@ def _require_native_stream( "Native stream does not preserve complete original conditioning" ) key = _sampled_source_key(source) + for index in range(len(prompt), len(prompt) + len(output)): + if expected[index] is not None: + raise ValueError("Native stream sampled source spans overlap") + expected[index] = key suffix = _sampled_stop_suffix( output, source=source, source_key=key, tokenizer=builder.tokenizer ) @@ -7885,6 +7890,8 @@ def _require_native_stream( ] if observed != list(range(len(output) - suffix, len(output))): raise ValueError("Native stream sampled STOP differs from source authority") + if trace.source_keys != expected: + raise ValueError("Native stream has sampled tokens outside original sources") prompt, output, _ = _chat_source_record(sources[-1]) if prompt is None or output is None or value.tokens != [*prompt, *output]: raise ValueError("Native stream must end at its final recorded output") diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index 1753f9960..af9c8a72d 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -256,7 +256,8 @@ def test_incomplete_or_edited_history_does_not_authorize_splitting(change: str) @pytest.mark.parametrize( - "change", ["missing_source", "prompt", "lp", "stop", "extra_stop", "flags"] + "change", + ["missing_source", "prompt", "lp", "stop", "extra_stop", "flags", "extra_sample"], ) def test_scoped_final_guard_rejects_corrupt_results(change: str) -> None: trajectory, tokenizer = example() @@ -277,6 +278,11 @@ def test_scoped_final_guard_rejects_corrupt_results(change: str) -> None: value.flags[stop] &= ~tr.TokenFlag.STOP elif change == "extra_stop": value.flags[sampled[-1]] |= tr.TokenFlag.STOP + elif change == "extra_sample": + value.flags[0] |= tr.TokenFlag.SAMPLED + value.logprobs[0] = -0.5 + trace.source_keys[0] = trace.source_keys[sampled[0]] + trace.validate(value) # Coherent trace membership is not complete coverage. else: value.flags[sampled[0]] &= ~tr.TokenFlag.SAMPLED with pytest.raises((ValueError, AssertionError)): From 9273bf5ec21c315c6e43b2f408e3a542e47726d0 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 07:36:22 +0000 Subject: [PATCH 05/10] Isolate prefix warning state in native stream tests --- tests/unit/trajectories/test_native_streams.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index af9c8a72d..3eee639bb 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -14,6 +14,11 @@ from art.trajectories import _tokenize as module +@pytest.fixture(autouse=True) +def isolate_prefix_warning(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(module, "_WARNED_PREFIX_RETOKENIZATION", False) + + def record(exchange: tr.ChatCompletionsExchange) -> Any: return cast(Any, exchange.response.choices[0]) From d1bb8e90e394e56f266f856a3715ca34fa7bddc3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 10:35:43 +0000 Subject: [PATCH 06/10] Verify native-stream losses and gradients against recorded completions --- .../test_native_streams_training.py | 115 ++++++++++++++++-- 1 file changed, 107 insertions(+), 8 deletions(-) diff --git a/tests/unit/trajectories/test_native_streams_training.py b/tests/unit/trajectories/test_native_streams_training.py index ed6ad9f88..9bd589ad2 100644 --- a/tests/unit/trajectories/test_native_streams_training.py +++ b/tests/unit/trajectories/test_native_streams_training.py @@ -1,27 +1,47 @@ from __future__ import annotations from collections import Counter +import math from typing import Any import pytest -from test_native_streams import example, native_terms +from test_native_streams import example, record from test_tokenize import _chat_exchange +import torch +from art.loss import LossInputs, loss_fn from art.preprocessing.pack import packed_tensors_from_tokenized_results from art.preprocessing.tokenize import tokenize_trajectory_groups import art.trajectories as tr -def test_training_weights_and_packing_keep_each_native_causal_path() -> None: +@pytest.mark.parametrize("ppo", [False, True]) +def test_training_loss_gradients_and_packing_keep_each_native_causal_path( + ppo: bool, +) -> None: trajectory, tokenizer = example() tokenizer.name_or_path = "public/base" control = tr.Trajectory( reward=1.0, exchanges=tr.TrajectoryExchanges( - chat_completions=[_chat_exchange([1, 2], [3])] + chat_completions=[ + _chat_exchange( + list( + record( + trajectory.exchanges.chat_completions[0] + ).prompt_token_ids + ), + [3, 4, 5], + ) + ] ), ) group = tr.TrajectoryGroup(trajectories=[trajectory, control]) + # Include clipped and unclipped ratios for both signs of advantage. + for original in group.trajectories: + for exchange in original.exchanges.chat_completions: + for index, entry in enumerate(record(exchange).logprobs.content): + entry.logprob = (-4.8, -5.5, -7.4)[index % 3] before = group.model_dump_json() results = list( tokenize_trajectory_groups( @@ -37,7 +57,21 @@ def test_training_weights_and_packing_keep_each_native_causal_path() -> None: for original, advantage in [(trajectory, 1.0), (control, -1.0)]: selected = [r for r in results if r.trajectory is original] assert selected - expected = native_terms(original) + # Enumerate original completions, independent of histories, masks or trie. + expected = [] + claimed = set() + for exchange in original.exchanges.chat_completions: + choice = record(exchange) + for index, entry in enumerate(choice.logprobs.content): + prefix = tuple(choice.prompt_token_ids + choice.token_ids[: index + 1]) + key = (exchange.request["model"], prefix) + if key not in claimed: + claimed.add(key) + expected.append((key, entry.logprob)) + terms.extend( + (prefix, lp, advantage, 1 / (len(expected) + 1e-6)) + for (_, prefix), lp in expected + ) actual = [] for result in selected: assert result.advantage == advantage @@ -46,9 +80,6 @@ def test_training_weights_and_packing_keep_each_native_causal_path() -> None: if sampled: prefix = tuple(result.token_ids[: index + 1]) actual.append((("test/model", prefix), result.logprobs[index])) - terms.append( - (prefix, result.logprobs[index], advantage, result.weight) - ) assert actual == expected packed = packed_tensors_from_tokenized_results( results, @@ -57,9 +88,30 @@ def test_training_weights_and_packing_keep_each_native_causal_path() -> None: verbosity=0, min_prefix_tree_shared_segment_length=1, ) + stats = packed["prefix_tree_packing_stats"] + assert stats["physical_tokens"] < stats["logical_tokens"] # Recover causal ancestors from the actual packed group/parent structure; # a sibling's tokens must never appear in another source's conditioning. observed = [] + # A small causal softmax predictor: shared parameters make wrong context or + # repeated ownership observable in gradients, without a transformer/backend. + theta = torch.linspace(-0.2, 0.2, 4 * 256, dtype=torch.float64).reshape(4, 256) + theta.requires_grad_() + + def features(prefix: tuple[int, ...]) -> torch.Tensor: + context = prefix[:-1] + return theta.new_tensor( + [ + 1, + len(context) / 50, + sum(context) / 10_000, + sum((i + 1) * token for i, token in enumerate(context)) / 100_000, + ] + ) + + # Ignored positions depend on parameters too: the real loss mask must remove + # their gradients. Selected target j is predicted at packed position j - 1. + predictions = theta.sum().expand(packed["tokens"].shape).clone() for row in range(packed["tokens"].shape[0]): tokens = packed["tokens"][row].tolist() groups = packed["group_ids"][row].tolist() @@ -78,9 +130,14 @@ def test_training_weights_and_packing_keep_each_native_causal_path() -> None: current = parent indices = [j for j in range(index + 1) if groups[j] in ancestry] assert [positions[j] for j in indices] == list(range(positions[index] + 1)) + prefix = tuple(tokens[j] for j in indices) + assert index > 0 + predictions[row, index - 1] = (features(prefix) @ theta).log_softmax(0)[ + prefix[-1] + ] observed.append( ( - tuple(tokens[j] for j in indices), + prefix, packed["logprobs"][row, index].item(), packed["advantages"][row, index].item(), packed["weights"][row, index].item(), @@ -97,6 +154,48 @@ def rounded(values: Any) -> Counter[Any]: assert rounded(observed) == rounded( (p, lp, a / advantage_scale, w / mean_weight) for p, lp, a, w in terms ) + actual_loss = loss_fn( + LossInputs(inputs=packed), predictions, None, None, {"ppo": ppo} + ) + actual_loss.policy_loss.backward() + assert theta.grad is not None and theta.grad.norm() > 0 + + # Closed-form per-completion loss/Jacobian: no packed fields, ART loss helper, + # masks or autograd are used by this oracle. Only finite public LPs are used. + expected_loss = 0.0 + expected_gradient = torch.zeros_like(theta) + active_signs = set() + clipped_signs = set() + for prefix, old_lp, advantage, weight in terms: + feature = features(prefix) + probabilities = (feature @ theta.detach()).softmax(0) + new_lp = probabilities[prefix[-1]].log().item() + # The recorded LP crosses the normal float32 packing boundary. + old_lp = torch.tensor(old_lp, dtype=torch.float32).item() + ratio = math.exp(new_lp - old_lp) + scale = advantage / advantage_scale * weight / mean_weight / len(terms) + if ppo: + clipped = (advantage > 0 and ratio > 1.2) or (advantage < 0 and ratio < 0.8) + expected_loss -= scale * (min(1.2, max(0.8, ratio)) if clipped else ratio) + coefficient = 0.0 if clipped else scale * ratio + else: + clipped = ratio > 5.0 + coefficient = scale * min(5.0, ratio) + expected_loss -= coefficient * new_lp + if coefficient: + active_signs.add(advantage) + if clipped: + clipped_signs.add(advantage) + jacobian = -probabilities + jacobian[prefix[-1]] += 1 + expected_gradient -= coefficient * feature[:, None] * jacobian[None, :] + assert active_signs == clipped_signs == {-1.0, 1.0} + # Packing normalizes coefficients in float32; the independent predictor and + # analytic sum use float64, including cancellation between opposite signs. + assert actual_loss.policy_loss.item() == pytest.approx( + expected_loss, rel=2e-6, abs=1e-6 + ) + torch.testing.assert_close(theta.grad, expected_gradient, rtol=2e-6, atol=2e-8) assert group.model_dump_json() == before From c1b63b6add0fb6ffd607042fd34ba12e0d030880 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 19:27:10 +0000 Subject: [PATCH 07/10] Cover cross-history request context mutation --- .../unit/trajectories/test_native_streams.py | 49 +++++++++++++++++-- 1 file changed, 44 insertions(+), 5 deletions(-) diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index c70fdee9c..1165774f4 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -495,7 +495,16 @@ def mutate(self: Any, text: str, **kwargs: Any) -> Any: @pytest.mark.parametrize("private", [False, True]) @pytest.mark.parametrize( - "change", ["request", "logprob", "stop_reason", "unscoped_logprob"] + "change", + [ + "request", + "logprob", + "stop_reason", + "unscoped_logprob", + "unscoped_request_role", + "unscoped_request_tools", + "unscoped_history_role", + ], ) def test_later_stop_encoder_cannot_change_an_already_certified_stream( monkeypatch: pytest.MonkeyPatch, private: bool, change: str @@ -506,11 +515,17 @@ def test_later_stop_encoder_cannot_change_an_already_certified_stream( record(first).model_extra["stop_reason"] = "§" record(second).model_extra["stop_reason"] = "§" earlier = first - if change == "unscoped_logprob": + if change.startswith("unscoped_"): earlier = _chat_exchange( - tokenizer._encode("separate"), tokenizer._encode("answer§"), offset=-1 + tokenizer._encode("separatehistorical§again"), + tokenizer._encode("answer§"), + offset=-1, ) - earlier.request["messages"] = [{"role": "user", "content": "separate"}] + earlier.request["messages"] = [ + {"role": "user", "content": "separate"}, + {"role": "assistant", "content": "historical"}, + {"role": "user", "content": "again"}, + ] trajectory.exchanges.chat_completions.insert(0, earlier) # Request-owned roles make the old sampled-only boundary shortcut decline, # so the later stream's validation callback is actually reached. @@ -524,9 +539,22 @@ def test_later_stop_encoder_cannot_change_an_already_certified_stream( ] = tokenizer._encode("bridgerequest-only§") original_guard = module._require_native_stream original_encode = tokenizer.__class__.__call__ + original_history = module.tokenize_history + completed = [] armed = False changed = False + def history_call(history: Any, *args: Any, **kwargs: Any) -> Any: + value = original_history(history, *args, **kwargs) + if change.startswith("unscoped_") and any( + source is not None and source.exchange is earlier + for source in history.message_sources + ): + assert value.flags[len("separate")] & tr.TokenFlag.ASSISTANT + assert not value.flags[len("separate")] & tr.TokenFlag.SAMPLED + completed.append(value) + return value + def guard(history: Any, *args: Any, **kwargs: Any) -> None: nonlocal armed # Final scope validation supplies planning snapshots; earlier assembly @@ -548,15 +576,25 @@ def encode(self: Any, text: str, **kwargs: Any) -> Any: first.request["chat_template_kwargs"] = {"changed": True} elif change in {"logprob", "unscoped_logprob"}: record(earlier).logprobs.content[0].logprob = -123.0 + elif change == "unscoped_request_role": + earlier.request["messages"][1]["role"] = "user" + elif change == "unscoped_request_tools": + earlier.request["tools"] = [ + {"type": "function", "function": {"name": "changed"}} + ] + elif change == "unscoped_history_role": + assert len(completed) == 1 + completed[0].history.messages[1]["role"] = "user" else: record(first).model_extra["stop_reason"] = "!" return original_encode(self, text, **kwargs) monkeypatch.setattr(module, "_require_native_stream", guard) + monkeypatch.setattr(module, "tokenize_history", history_call) monkeypatch.setattr(tokenizer.__class__, "__call__", encode) message = ( "Sampled source changed during tokenization callback" - if change == "unscoped_logprob" + if change.startswith("unscoped_") else "during final STOP validation" ) with pytest.raises(ValueError, match=message): @@ -564,4 +602,5 @@ def encode(self: Any, text: str, **kwargs: Any) -> Any: module._tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) else: trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + assert changed assert changed From cbec1abf39d3e01384f913efbf3d09cf43342499 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 19:45:41 +0000 Subject: [PATCH 08/10] Retain original history context through native stream refinement --- src/art/trajectories/_tokenize.py | 330 +++++++++++++++--- .../unit/trajectories/test_native_streams.py | 26 +- 2 files changed, 302 insertions(+), 54 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 8e9523677..9b3a099e4 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -945,6 +945,8 @@ class _TraceBuilder: tokenizer: Tokenizer | None = None rendered_outputs: tuple[tuple[int, int, object], ...] = () validate_sources: Callable[[_SampledSourceKey | None], None] | None = None + validate_context: Callable[[bool], None] | None = None + track_sources: bool = True def set( self, @@ -960,7 +962,8 @@ def set( trace.validate(tokenized) self.trace = trace self.rendered_outputs = rendered_outputs - self.validate_sources = _sampled_source_validator(sources) + if self.track_sources: + self.validate_sources = _sampled_source_validator(sources) def _fingerprint(value: object) -> str: @@ -4295,6 +4298,69 @@ def _source_has_no_materialized_output( return False +def _tokenization_context(value: object) -> object: + """Snapshot ordered semantic inputs without serializing sampled responses. + + History and source fields remain typed, including protocol selectors and + request-owned context. Sampled response evidence has its own source-key + validator; retaining entire response objects here would duplicate it. + """ + exchanges: dict[int, object] = {} + + def snapshot(item: object) -> object: + kind = type(item) + if kind in (str, int, bool, bytes, type(None), datetime): + return kind, item + if kind is float: + return kind, repr(item) + if kind in (list, tuple): + return kind, tuple(snapshot(child) for child in cast(Sequence, item)) + if isinstance(item, Mapping): + return kind, tuple( + (snapshot(key), snapshot(child)) for key, child in item.items() + ) + if isinstance(item, Exchange): + if id(item) not in exchanges: + exchanges[id(item)] = ( + kind, + id(item), + item.model, + snapshot(item.request), + ) + return exchanges[id(item)] + if isinstance(item, BaseModel): + return ( + kind, + tuple( + (name, snapshot(getattr(item, name))) + for name in type(item).model_fields + ), + snapshot(item.model_extra), + ) + raise TypeError("Unsupported mutable tokenization context") + + return snapshot(value) + + +def _tokenization_context_validator(value: object) -> Callable[[bool], None]: + try: + expected = _tokenization_context(value) + except TypeError: + # An opaque, complete native history still works without a tokenizer. + # A callback-bearing path cannot claim an uncheckable context proof. + expected = None + + def validate(require_supported: bool) -> None: + if expected is None and not require_supported: + return + if expected is None or _tokenization_context(value) != expected: + raise ValueError( + "Tokenization context changed during tokenization callback" + ) + + return validate + + def _sampled_source_validator( sources: Mapping[_SampledSourceKey, object], ) -> Callable[[_SampledSourceKey | None], None]: @@ -4308,11 +4374,12 @@ def _sampled_source_validator( exchange, exchange.model, _source_stop_evidence(source, key), + _tokenization_context_validator(exchange.request), ) def validate(selected: _SampledSourceKey | None) -> None: for key in expected if selected is None else (selected,): - source, exchange, model, stop = expected[key] + source, exchange, model, stop, validate_request = expected[key] current = ( _exchange_sampled_source_key(source) if isinstance(source, Exchange) @@ -4325,6 +4392,7 @@ def validate(selected: _SampledSourceKey | None) -> None: or _source_stop_evidence(source, key) != stop ): raise ValueError("Sampled source changed during tokenization callback") + validate_request(True) return validate @@ -4927,7 +4995,7 @@ def _tokenize_recorded_chat_boundaries( if not callable(decode) or not messages or messages[-1].get("role") != "assistant": return None entries: list[tuple[int, object, list[int], list[int], list[float]]] = [] - seen: set[_SampledSourceKey] = set() + sources: dict[_SampledSourceKey, object] = {} for index, (message, source) in enumerate( zip(messages, history.message_sources, strict=True) ): @@ -4937,9 +5005,9 @@ def _tokenize_recorded_chat_boundaries( return None key = _sampled_source_key(source) prompt, output, logprobs = _chat_source_record(source) - if key in seen or prompt is None or output is None: + if key in sources or prompt is None or output is None: return None - seen.add(key) + sources[key] = source entries.append((index, source, prompt, output, logprobs)) if not entries: return None @@ -4967,73 +5035,120 @@ def _tokenize_recorded_chat_boundaries( ): return None boundaries: dict[_SampledSourceKey, _RenderedLengthStopBoundary] = {} - terminators = _terminator_ids(tokenizer) + # Entries already consumed every native record. No callback may replace + # that evidence before a later source or the final assembler reads it again. + validate_sources = _sampled_source_validator(sources) + validate_context = _tokenization_context_validator(history) + keys = tuple(sources) + selected_key: _SampledSourceKey | None = None + + def checked(call: Callable[[], Any], *, optional_decode: bool = False) -> Any: + try: + value = call() + except (TypeError, KeyError, NotImplementedError): + # These capability failures are caught by the caller and may + # resume tokenization. A changed source must not reach fallback. + validate_sources(selected_key) + validate_context(True) + raise + except ValueError: + if not optional_decode: + raise + value = None + # Other callback exceptions propagate unchanged; no result or fallback + # can consume their potentially changed inputs. + validate_sources(selected_key) + validate_context(True) + return value + + def decline() -> None: + validate_sources(None) + validate_context(True) + + terminators = checked(lambda: _terminator_ids(tokenizer)) if not terminators: - return None + return decline() # The last native output ends the recorded history. No later prompt proves # an additional footer, so do not reconstruct one from a lossy projection. for ordinal, (index, source, prompt, output, _) in enumerate(entries[:-1]): - key = _sampled_source_key(source) + selected_key = key = keys[ordinal] + validate_sources(key) stop, _ = _source_stop_evidence(source, key) if stop not in {"stop", "length"}: - return None - if stop == "stop" and _sampled_stop_suffix( - output, source=source, source_key=key, tokenizer=tokenizer + return decline() + if stop == "stop" and checked( + lambda: _sampled_stop_suffix( + output, source=source, source_key=key, tokenizer=tokenizer + ) ): continue try: - try: - body = decode( + body = checked( + lambda: decode( output, skip_special_tokens=False, clean_up_tokenization_spaces=False, - ) - except ValueError: - return None - generation = render(messages[:index], add_generation_prompt=True) - completed = render(messages[: index + 1], add_generation_prompt=False) + ), + optional_decode=True, + ) + if body is None: + return decline() + generation = checked( + lambda: render(messages[:index], add_generation_prompt=True) + ) + completed = checked( + lambda: render(messages[: index + 1], add_generation_prompt=False) + ) # Only the actually sampled body anchors the tail. Literal content, # tool JSON and reasoning are never searched for or re-tokenized. if not isinstance(body, str) or not completed.startswith(generation + body): - return None + return decline() suffix = completed[len(generation) + len(body) :] - tail = _ids(tokenizer(suffix, add_special_tokens=False)) if suffix else [] + tail = ( + _ids(checked(lambda: tokenizer(suffix, add_special_tokens=False))) + if suffix + else [] + ) stops = [i for i, token in enumerate(tail) if token in terminators] if len(stops) != 1: - return None + return decline() terminator = stops[0] - try: - trailing = decode( + trailing = checked( + lambda: decode( tail[terminator + 1 :], skip_special_tokens=False, clean_up_tokenization_spaces=False, - ) - except ValueError: - return None + ), + optional_decode=True, + ) if not isinstance(trailing, str) or trailing and not trailing.isspace(): - return None + return decline() following: list[int] = [] if ordinal + 1 < len(entries): next_index, _, next_prompt, _, _ = entries[ordinal + 1] - next_generation = render( - messages[:next_index], add_generation_prompt=True + next_generation = checked( + lambda: render(messages[:next_index], add_generation_prompt=True) ) if not next_generation.startswith(completed): - return None + return decline() gap = suffix + next_generation[len(completed) :] - gap_ids = _ids(tokenizer(gap, add_special_tokens=False)) + gap_ids = _ids( + checked(lambda: tokenizer(gap, add_special_tokens=False)) + ) if ( gap_ids[: len(tail)] != tail or next_prompt[len(prompt) + len(output) :] != gap_ids ): - return None + return decline() following = gap_ids[len(tail) :] boundaries[key] = _RenderedLengthStopBoundary( tail=tuple(tail[: terminator + 1]), following=tuple([*tail[terminator + 1 :], *following]), ) except (TypeError, KeyError, NotImplementedError): - return None + return decline() + validate_sources(None) + validate_context(True) return _tokenize_exact_projected_chat_history( history, tokenizer=tokenizer, @@ -6371,16 +6486,80 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: _trace=_trace, ) ) - and _merge_recorded_request_roles( + ): + if _merge_recorded_request_roles( exact, rendered, assistant_mask, output_mask, stop_mask, length_stop_mask, - ) - ): - return exact + ): + return exact + if ( + _recorded_boundaries + and chat_template is None + and chat_template_kwargs is None + ): + # Later request-only assistants can have historical rendering + # different from today's normalized template too. Prove the + # entire final recorded request, not only the first prompt. + source = history.message_sources[sampled_message_indices[-1]] + assert source is not None + exchange = _source_exchange(source) + assert isinstance( + exchange, + (ChatCompletionsExchange, MessagesExchange, ResponsesExchange), + ) + prompt = source_prompt_tokens(source) + signature = _source_signature(source) + validate_context = _tokenization_context_validator(history) + masks = None + if prompt and exact.tokens[: len(prompt)] == prompt: + request_messages, request_tools = _request_messages(exchange) + try: + masks = _recorded_prompt_role_masks( + request_messages, + [None] * len(request_messages), + prompt, + tokenizer=resolved_tokenizer, + template=original_template, + tools=request_tools, + kwargs=kwargs, + ) + except (TypeError, KeyError, NotImplementedError): + pass + validate_context(True) + prompt_cache.clear() + output_cache.clear() + if _source_signature(source) != signature: + raise ValueError( + "Sampled source changed while proving recorded request roles" + ) + if ( + masks is not None + and all( + flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) + or not (flag & TokenFlag.ASSISTANT) + or assistant + for flag, assistant in zip(exact.flags, masks[0]) + ) + and all( + flag & (TokenFlag.SAMPLED | TokenFlag.OUTPUT) + or not (flag & TokenFlag.STOP) + or stop + for flag, stop in zip(exact.flags, masks[1]) + ) + ): + for index, (assistant, stop) in enumerate(zip(*masks, strict=True)): + if not exact.flags[index] & ( + TokenFlag.SAMPLED | TokenFlag.OUTPUT + ): + exact.flags[index] |= _rendered_flag(assistant, False, stop) + return exact + raise ValueError( + "Cannot preserve request roles across exact native prompt replacement" + ) sampled_message_count = sum( message.get("role") == "assistant" @@ -7678,10 +7857,18 @@ def tokenize_history( _native_stream: bool = False, _allow_native_streams: bool = False, ) -> TokenizedHistory: + trace_builder = _trace or _TraceBuilder(track_sources=False) + if trace_builder.validate_context is None: + trace_builder.validate_context = _tokenization_context_validator( + [history, chat_template, chat_template_kwargs] + ) + else: + trace_builder.validate_context(False) if _native_stream: assert isinstance(history, ChatCompletionsHistory) _validate_history_sources(history) - builder = _trace or _TraceBuilder() + builder = trace_builder + builder.track_sources = True # The scope owns the final request, including any historical assistant # turns. Prove their generic role flags; native IDs alone cannot do so. tokenized = _tokenize_chat_view( @@ -7696,6 +7883,8 @@ def tokenize_history( _prior=_prior, ) _require_native_stream(history, tokenized, builder) + if builder.tokenizer is not None: + trace_builder.validate_context(True) return tokenized copied = ( list(_context_sources) @@ -7703,6 +7892,7 @@ def tokenize_history( else _partial_native_context(history) ) if copied: + trace_builder.track_sources = True history = cast(History, history) _validate_history_sources(history) state = None if _projection_validated else _history_render_state(history) @@ -7739,7 +7929,6 @@ def tokenize_history( raise ValueError( "A copied response suffix requires its complete original sampled occurrence in the selected trajectory" ) - trace_builder = _trace or (_TraceBuilder() if copied else None) tokenized = _tokenize_history( history, model=model, @@ -7753,6 +7942,8 @@ def tokenize_history( _copied_context=bool(copied), _allow_native_streams=_allow_native_streams, ) + if trace_builder.tokenizer is not None: + trace_builder.validate_context(True) if copied: if trace_builder is None or trace_builder.trace is None: raise ValueError( @@ -8052,8 +8243,11 @@ def _validate_completed_sources(builders: Sequence[_TraceBuilder | None]) -> Non # Later callbacks may edit an earlier completed history. Check its # original source keys and stop evidence without calling a tokenizer. for builder in builders: - if builder is not None and builder.validate_sources is not None: - builder.validate_sources(None) + if builder is not None: + if builder.validate_context is not None: + builder.validate_context(True) + if builder.validate_sources is not None: + builder.validate_sources(None) def _complete_resolved_sampled_stops( @@ -8078,6 +8272,8 @@ def _complete_resolved_sampled_stops( and (tokenizer := resolved.get(value.model)) is not None ): assert builder.validate_sources is not None + if builder.validate_context is not None: + builder.validate_context(True) builder.validate_sources(None) _mark_sampled_stops( value.tokens, @@ -8129,6 +8325,17 @@ def tokenize_trajectory( prior: list[tuple[TokenizedHistory, _HistoryTokenizationTrace]] = [] tokenized = [] stop_builders: list[_TraceBuilder | None] = [] + scope_builders = [ + _TraceBuilder( + track_sources=allow_streams or track_context or len(histories) > 1, + validate_context=_tokenization_context_validator( + [history, chat_template, chat_template_kwargs] + ), + ) + for history in histories + ] + # Refinement must not discard the original pre-callback context proof. + all_builders = list(scope_builders) scopes: list[tuple[History | LegacyHistory, bool]] = [ (history, False) for history in histories ] @@ -8137,11 +8344,7 @@ def tokenize_trajectory( index = 0 while index < len(scopes): history, scoped = scopes[index] - trace = ( - _TraceBuilder() - if scoped or allow_streams or track_context or len(histories) > 1 - else None - ) + trace = scope_builders[index] try: result = tokenize_history( history, @@ -8168,6 +8371,16 @@ def tokenize_trajectory( scoped_keys.update(planned.keys) scoped_contexts.update(planned.contexts) scopes[index : index + 1] = [(stream, True) for stream in planned.histories] + planned_builders = [ + _TraceBuilder( + validate_context=_tokenization_context_validator( + [stream, chat_template, chat_template_kwargs] + ) + ) + for stream in planned.histories + ] + scope_builders[index : index + 1] = planned_builders + all_builders.extend(planned_builders) continue tokenized.append(result) stop_builders.append(trace) @@ -8188,7 +8401,7 @@ def tokenize_trajectory( scoped_contexts[id(history)], ) _require_native_streams_unchanged(scopes, scoped_keys, scoped_contexts) - _validate_completed_sources(stop_builders) + _validate_completed_sources(all_builders) if not multi_history: return _materialize_trajectory(tokenized[0], trajectory) return TokenizedMultiHistoryTrajectory( @@ -8220,6 +8433,15 @@ def _tokenize_trajectory_with_trace( tokenized_histories: list[TokenizedHistory] = [] traces: list[_HistoryTokenizationTrace] = [] builders: list[_TraceBuilder] = [] + scope_builders = [ + _TraceBuilder( + validate_context=_tokenization_context_validator( + [history, chat_template, chat_template_kwargs] + ) + ) + for history in histories + ] + all_builders = list(scope_builders) index = 0 while index < len(scopes): history, scoped = scopes[index] @@ -8227,7 +8449,7 @@ def _tokenize_trajectory_with_trace( raise AssertionError( "Exchange trajectories cannot produce legacy histories" ) - trace_builder = _TraceBuilder() + trace_builder = scope_builders[index] try: tokenized = tokenize_history( history, @@ -8256,6 +8478,16 @@ def _tokenize_trajectory_with_trace( scoped_keys.update(planned.keys) scoped_contexts.update(planned.contexts) scopes[index : index + 1] = [(stream, True) for stream in planned.histories] + planned_builders = [ + _TraceBuilder( + validate_context=_tokenization_context_validator( + [stream, chat_template, chat_template_kwargs] + ) + ) + for stream in planned.histories + ] + scope_builders[index : index + 1] = planned_builders + all_builders.extend(planned_builders) continue if trace_builder.trace is None: raise AssertionError("Exchange tokenization did not produce a source trace") @@ -8276,7 +8508,7 @@ def _tokenize_trajectory_with_trace( scoped_contexts[id(history)], ) _require_native_streams_unchanged(scopes, scoped_keys, scoped_contexts) - _validate_completed_sources(builders) + _validate_completed_sources(all_builders) return ( TokenizedMultiHistoryTrajectory( trajectory=trajectory, diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index 1165774f4..b49bbeaac 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -504,6 +504,7 @@ def mutate(self: Any, text: str, **kwargs: Any) -> Any: "unscoped_request_role", "unscoped_request_tools", "unscoped_history_role", + "expanded_original_role", ], ) def test_later_stop_encoder_cannot_change_an_already_certified_stream( @@ -540,10 +541,20 @@ def test_later_stop_encoder_cannot_change_an_already_certified_stream( original_guard = module._require_native_stream original_encode = tokenizer.__class__.__call__ original_history = module.tokenize_history + original_planner = module._native_history_streams completed = [] + expanded = [] armed = False changed = False + def plan_streams(history: Any) -> Any: + if any( + source is not None and source.exchange is second + for source in history.message_sources + ): + expanded.append(history) + return original_planner(history) + def history_call(history: Any, *args: Any, **kwargs: Any) -> Any: value = original_history(history, *args, **kwargs) if change.startswith("unscoped_") and any( @@ -585,18 +596,23 @@ def encode(self: Any, text: str, **kwargs: Any) -> Any: elif change == "unscoped_history_role": assert len(completed) == 1 completed[0].history.messages[1]["role"] = "user" + elif change == "expanded_original_role": + assert len(expanded) == 1 + expanded[0].messages[0]["role"] = "assistant" else: record(first).model_extra["stop_reason"] = "!" return original_encode(self, text, **kwargs) monkeypatch.setattr(module, "_require_native_stream", guard) monkeypatch.setattr(module, "tokenize_history", history_call) + monkeypatch.setattr(module, "_native_history_streams", plan_streams) monkeypatch.setattr(tokenizer.__class__, "__call__", encode) - message = ( - "Sampled source changed during tokenization callback" - if change.startswith("unscoped_") - else "during final STOP validation" - ) + if change == "unscoped_logprob": + message = "Sampled source changed during tokenization callback" + elif change.startswith("unscoped_") or change == "expanded_original_role": + message = "Tokenization context changed during tokenization callback" + else: + message = "during final STOP validation" with pytest.raises(ValueError, match=message): if private: module._tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) From 6b5115dd6761b40f2484ec972158df704c678d87 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 20:01:59 +0000 Subject: [PATCH 09/10] Preserve the late native role proof before stream refinement --- src/art/trajectories/_tokenize.py | 6 +-- .../unit/trajectories/test_native_streams.py | 53 +++++++++++++++++++ 2 files changed, 56 insertions(+), 3 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 4ffa4169b..5598b5541 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -5476,9 +5476,6 @@ def render_text( ): return recorded - if _allow_native_streams: - _request_native_streams(history) - prefix_render_cache = _PrefixChatRenderCache(render_normalized_text) def segmented_render( @@ -6563,6 +6560,9 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: "Cannot preserve request roles across exact native prompt replacement" ) + if _allow_native_streams: + _request_native_streams(history) + sampled_message_count = sum( message.get("role") == "assistant" and source is not None diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index b49bbeaac..9044b6aaa 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -164,6 +164,59 @@ def changed_renderer(messages: Any, **kwargs: Any) -> Any: assert trajectory.model_dump_json() == before +def test_existing_late_native_role_proof_precedes_stream_refinement( + monkeypatch: pytest.MonkeyPatch, +) -> None: + tokenizer = _CharacterTemplateTokenizer() + messages = [ + {"role": "user", "content": "intro"}, + {"role": "assistant", "content": "history"}, + {"role": "user", "content": "turn0"}, + ] + prompt = "introhistory§turn0" + first = _chat_exchange(tokenizer._encode(prompt), tokenizer._encode("rawanswer§")) + first.request["messages"] = deepcopy(messages) + record(first).message.content = "answer" + second = _chat_exchange( + tokenizer._encode(prompt + "answer§turn1"), + tokenizer._encode("terminal§"), + offset=1, + ) + second.request["messages"] = [ + *deepcopy(messages), + {"role": "assistant", "content": "answer"}, + {"role": "user", "content": "turn1"}, + ] + record(second).message.content = "terminal" + for exchange in (first, second): + record(exchange).finish_reason = "stop" + record(exchange).model_extra["stop_reason"] = tokenizer.eos_token_id + trajectory = tr.Trajectory( + exchanges=tr.TrajectoryExchanges(chat_completions=[first, second]), reward=2.0 + ) + render = tokenizer.apply_chat_template + + def changed_renderer(selected: Any, **kwargs: Any) -> Any: + copied = deepcopy(selected) + for message in copied: + if ( + message.get("role") == "assistant" + and message.get("content") == "answer" + ): + message["content"] = "!answer" + return render(copied, **kwargs) + + monkeypatch.setattr(tokenizer, "apply_chat_template", changed_renderer) + before = trajectory.model_dump_json() + with monkeypatch.context() as old_route: + old_route.setattr(module, "_native_history_streams", lambda history: [history]) + existing = trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + result = trajectory.tokenize(multi_history=True, tokenizer=tokenizer) + assert result.model_dump_json() == existing.model_dump_json() + assert result_terms(result) == native_terms(trajectory) + assert trajectory.model_dump_json() == before + + @pytest.mark.parametrize("middle_kind", ["tool", "reasoning"]) def test_native_stream_keeps_complete_nonterminal_structured_output( middle_kind: str, From df2b5167dcd3e72619a0555369b7082303b96fc3 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Sat, 26 Sep 2026 23:37:13 +0000 Subject: [PATCH 10/10] Type native-stream role mutation fixtures explicitly --- tests/unit/trajectories/test_native_streams.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/tests/unit/trajectories/test_native_streams.py b/tests/unit/trajectories/test_native_streams.py index a3c9bf505..a2942b811 100644 --- a/tests/unit/trajectories/test_native_streams.py +++ b/tests/unit/trajectories/test_native_streams.py @@ -6,7 +6,7 @@ import struct from typing import Any, cast -from openai.types.chat import ChatCompletion +from openai.types.chat import ChatCompletion, ChatCompletionMessageParam import pytest from test_tokenize import _CharacterTemplateTokenizer, _chat_exchange @@ -168,7 +168,7 @@ def test_existing_late_native_role_proof_precedes_stream_refinement( monkeypatch: pytest.MonkeyPatch, ) -> None: tokenizer = _CharacterTemplateTokenizer() - messages = [ + messages: list[ChatCompletionMessageParam] = [ {"role": "user", "content": "intro"}, {"role": "assistant", "content": "history"}, {"role": "user", "content": "turn0"}, @@ -645,7 +645,7 @@ def encode(self: Any, text: str, **kwargs: Any) -> Any: elif change in {"logprob", "unscoped_logprob"}: record(earlier).logprobs.content[0].logprob = -123.0 elif change == "unscoped_request_role": - earlier.request["messages"][1]["role"] = "user" + cast(dict[str, Any], earlier.request["messages"][1])["role"] = "user" elif change == "unscoped_request_tools": earlier.request["tools"] = [ {"type": "function", "function": {"name": "changed"}}