From 5e36a7f0a928eefe6334b0dda3266951003528bf Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 16:25:02 +0000 Subject: [PATCH 1/2] Add an explicit native representation for sampled tokenization --- docs/features/additional-histories.mdx | 42 +- src/art/trajectories/__init__.py | 24 +- src/art/trajectories/_parallel.py | 58 +- src/art/trajectories/_sampled_native.py | 269 ++++++++ .../unit/trajectories/test_sampled_native.py | 581 ++++++++++++++++++ 5 files changed, 952 insertions(+), 22 deletions(-) create mode 100644 src/art/trajectories/_sampled_native.py create mode 100644 tests/unit/trajectories/test_sampled_native.py diff --git a/docs/features/additional-histories.mdx b/docs/features/additional-histories.mdx index 756eea7d7..405042d04 100644 --- a/docs/features/additional-histories.mdx +++ b/docs/features/additional-histories.mdx @@ -160,7 +160,7 @@ for history in histories: Use `await art.tokenize_sampled(trajectories)` (or trajectory groups) when you need complete native Chat Completions output with model-bound STOP flags. This -opt-in API first performs ordinary `multi_history=True` tokenization without +default `representation="rendered"` mode first performs ordinary `multi_history=True` tokenization without renderer overrides or text reconciliation. After that succeeds, it resolves the exact history model's tokenizer configuration, including its revision, to certify each selected sampled source's nonempty original conditioning, output IDs and @@ -187,15 +187,49 @@ a different `base_model` does not supply authority. Incomplete native evidence, changed conditioning, unsupported sampled protocols, and extra or incorrect STOP flags raise an error, even if ordinary tokenization -succeeded. Missing or null logprob carriers are not recorded NaNs; explicitly -recorded raw NaNs remain valid evidence. Unsupported finish reasons such as -`content_filter` are not certified as an absence of STOP. Nonsampled histories are retained. This API does not recover failed +succeeded. Nonsampled histories are retained. This default mode does not recover failed rendering, split or join histories, or establish SFT equivalence. Generic `art.tokenize` remains unchanged: a native-only path can avoid loading a tokenizer, so missing STOP flags there do not prove that a terminating suffix is absent. Use the explicit API when that distinction matters. +### Using recorded native conditioning instead of rendering + +For a loss defined on **SAMPLED first occurrences, selected before filtering +nonfinite float32 logprobs**, you can explicitly select a native representation: + +```python +tokenized = await art.tokenize_sampled( + trajectories, model="my-policy", representation="native" +) +``` + +This mode never renders messages or catches an ordinary tokenizer failure. It +requires complete original Chat Completions prompt/output tokens and logprobs, +validates the entire selected source inventory, and uses each source model's +resolved STOP authority. Content, literal reasoning delimiters, structured +reasoning and tool calls remain in the original captured objects. Missing +evidence is an error; a different renderer is not substituted. + +Consecutive generations share a history only when each complete earlier native +prompt and output exactly prefixes the later request and their captured message +views agree. Other generations retain separate histories. Source encounters +remain in their original canonical order, including repeats across histories; +there is no global source deduplication. Native request gaps are exact, +**nonsampled context**, without guessed assistant/OUTPUT roles or synthetic +renderer STOP tails. IDs, conditioning, logprobs, sampled flags and STOP ownership +are checked again on constructed output. + +This preserves the ordered sampled objective, not generic OUTPUT/SFT loss, +arbitrary per-history weighting, packing, or floating-point execution order. +Mixed protocols, legacy/additional histories, incomplete source messages and +edited projections are unsupported and raise errors. This includes assistant +turns echoed without their original structured reasoning and response IDs reused +within one model. Original trajectory/group +metadata and source objects are retained. The same current-model-configuration +limitation on historical STOP authority applies to both representations. + ### Data Structure The legacy `LegacyHistory` payload structure: diff --git a/src/art/trajectories/__init__.py b/src/art/trajectories/__init__.py index f25c0f454..44733ea67 100644 --- a/src/art/trajectories/__init__.py +++ b/src/art/trajectories/__init__.py @@ -1613,6 +1613,7 @@ async def tokenize_sampled( *, model: str | None = None, base_model: str | None = None, + representation: Literal["rendered", "native"] = "rendered", ) -> list[TokenizedMultiHistoryTrajectory]: ... @@ -1622,6 +1623,7 @@ async def tokenize_sampled( *, model: str | None = None, base_model: str | None = None, + representation: Literal["rendered", "native"] = "rendered", ) -> list[TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory]]: ... @@ -1630,14 +1632,15 @@ async def tokenize_sampled( *, model: str | None = None, base_model: str | None = None, + representation: Literal["rendered", "native"] = "rendered", ) -> ( list[TokenizedMultiHistoryTrajectory] | list[TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory]] ): - """Tokenize ordinary histories, then certify native sampled output and STOP. + """Certify sampled output using ordinary rendering or explicit native sources. - This opt-in API uses ``multi_history=True`` without renderer overrides or - text reconciliation. It requires complete Chat Completions source messages, + The default ``representation="rendered"`` uses ``multi_history=True`` without + renderer overrides or text reconciliation. It requires complete Chat Completions source messages, nonempty original conditioning, output IDs and logprobs for every sampled span. Unsupported or incomplete sampled histories refuse; nonsampled histories are retained. Ordinary tokenization failures propagate without recovery. @@ -1651,9 +1654,23 @@ async def tokenize_sampled( rendering, provide SFT equivalence or repartition histories. Generic :func:`tokenize` retains its native-only, no-load behavior; absent STOP flags there do not imply that a terminating suffix is known to be absent. + + ``representation="native"`` instead constructs complete original Chat + Completions sources without rendering. Consecutive sources join only when + every earlier prompt/output is an exact prefix of the later native request; + otherwise they remain separate histories. Repeated encounters across + canonical histories remain in order. This mode preserves the objective only + for SAMPLED first-occurrence ownership chosen BEFORE float32 finite-logprob + filtering. It does not preserve OUTPUT/SFT loss, per-history weighting, + layout, packing or floating-point execution order. Native request gaps are + exact nonsampled context; synthetic renderer STOP tails are not invented. + Mixed protocols, additional/legacy histories, incomplete or edited sources + refuse. STOP authority still follows each resolved source-model configuration. """ from ._parallel import transform + if representation not in {"rendered", "native"}: + raise ValueError("Unknown sampled representation") return cast( Any, await transform( @@ -1667,6 +1684,7 @@ async def tokenize_sampled( chat_template=None, chat_template_kwargs=None, _sampled=True, + _native_sampled=representation == "native", ), ) diff --git a/src/art/trajectories/_parallel.py b/src/art/trajectories/_parallel.py index 670fc2b27..34516442a 100644 --- a/src/art/trajectories/_parallel.py +++ b/src/art/trajectories/_parallel.py @@ -536,6 +536,7 @@ class _ProcessOptions: chat_template: str | None chat_template_kwargs: Mapping[str, object] | None sampled: bool = False + native_sampled: bool = False class _ProcessTransferError(RuntimeError): @@ -555,21 +556,40 @@ def _tokenize_process_payload(payload: bytes) -> bytes: raise _ProcessTransferError( f"could not deserialize process input: {type(error).__name__}: {error}" ) from None - tokenized = trajectory.tokenize( - multi_history=options.multi_history, - reconcile_text_equivalent_tokenizations=( - options.reconcile_text_equivalent_tokenizations - ), - model=options.model, - base_model=options.base_model, - tokenizer=None, - chat_template=options.chat_template, - chat_template_kwargs=options.chat_template_kwargs, - ) - if options.sampled: - from ._sampled import reconcile_sampled_stops + if options.native_sampled: + if ( + not options.sampled + or not options.multi_history + or options.reconcile_text_equivalent_tokenizations + or options.chat_template is not None + or options.chat_template_kwargs is not None + ): + raise ValueError( + "Native representation requires unmodified sampled options" + ) + from ._sampled_native import tokenize_native + + tokenized = tokenize_native( + trajectory, model=options.model, base_model=options.base_model + ) + else: + tokenized = trajectory.tokenize( + multi_history=options.multi_history, + reconcile_text_equivalent_tokenizations=( + options.reconcile_text_equivalent_tokenizations + ), + model=options.model, + base_model=options.base_model, + tokenizer=None, + chat_template=options.chat_template, + chat_template_kwargs=options.chat_template_kwargs, + ) + if options.sampled: + from ._sampled import reconcile_sampled_stops - tokenized = reconcile_sampled_stops(tokenized, base_model=options.base_model) + tokenized = reconcile_sampled_stops( + tokenized, base_model=options.base_model + ) try: return pickle.dumps(tokenized, protocol=pickle.HIGHEST_PROTOCOL) except Exception as error: @@ -724,7 +744,10 @@ async def transform( chat_template_kwargs: Mapping[str, object] | None, device: Any = None, _sampled: bool = False, + _native_sampled: bool = False, ) -> list[object]: + if _native_sampled and not _sampled: + raise ValueError("Native representation requires sampled tokenization") if _sampled and ( operation != "tokenize" or not multi_history @@ -745,6 +768,10 @@ async def transform( ) def convert(trajectory: Trajectory) -> object: + if _native_sampled: + from ._sampled_native import tokenize_native + + return tokenize_native(trajectory, model=model, base_model=base_model) tokenized = trajectory.tokenize( multi_history=multi_history, reconcile_text_equivalent_tokenizations=reconcile_text_equivalent_tokenizations, @@ -775,7 +802,7 @@ def convert(trajectory: Trajectory) -> object: capacity=capacity, ) if _sampled: - key = (*key, "sampled_stops") + key = (*key, "sampled_native" if _native_sampled else "sampled_stops") use_processes = _supports_processes( capacity=capacity, size=len(leaves), tokenizer=tokenizer ) and _processes_enabled(key) @@ -790,6 +817,7 @@ def convert(trajectory: Trajectory) -> object: chat_template=chat_template, chat_template_kwargs=chat_template_kwargs, sampled=_sampled, + native_sampled=_native_sampled, ) try: workers = _process_workers(key, capacity=capacity, size=len(leaves)) diff --git a/src/art/trajectories/_sampled_native.py b/src/art/trajectories/_sampled_native.py new file mode 100644 index 000000000..e8dd07073 --- /dev/null +++ b/src/art/trajectories/_sampled_native.py @@ -0,0 +1,269 @@ +"""Explicit native sampled representation; never a generic rendering fallback.""" + +from __future__ import annotations + +from copy import deepcopy +from dataclasses import dataclass +import math + +from ..preprocessing.dynamo_tokens import ( + COMPLETION_LOGPROBS_KEY, + choice_completion_logprobs, +) +from . import ( + ChatCompletionsExchange, + ChatCompletionsHistory, + ChatCompletionsMessageSource, + TokenFlag, + TokenizedHistory, + TokenizedMultiHistoryTrajectory, + Tokenizer, + Trajectory, +) +from . import _tokenize as original +from ._history import normalize_chat_message +from ._sampled import _require_exact_chat_source_edges +from ._serialization import _equal_with_nan + + +@dataclass +class _Span: + source: ChatCompletionsMessageSource + key: original._SampledSourceKey + value: TokenizedHistory + start: int + + +def _singleton(source: ChatCompletionsMessageSource, bound: Tokenizer) -> _Span: + exchange = source.exchange + assert isinstance(exchange, ChatCompletionsExchange) + prompt = original._chat_source_prompt_tokens(source) + output, logprobs = original._chat_source_full_tokens(source) + choice = original._chat_choice(source) + recorded_logprobs = ( + choice_completion_logprobs(choice) + if COMPLETION_LOGPROBS_KEY in (choice.model_extra or {}) + else original._logprob_values(original._chat_logprob_entries(choice)) + ) + if ( + not prompt + or not output + or len(output) != len(logprobs) + or recorded_logprobs is None + or len(recorded_logprobs) != len(output) + ): + raise ValueError( + "Native sampled sources require complete prompt, output and logprobs" + ) + key = original._sampled_source_key(source) + if original._source_stop_evidence(source, key)[0] not in {"stop", "length"}: + raise ValueError("Native sampled sources require supported STOP evidence") + request = exchange.request.get("messages") + message = original._chat_choice_message(source) + if not isinstance(request, list) or message is None: + raise ValueError("Native sampled source message is unavailable") + history = ChatCompletionsHistory( + model=exchange.model, + messages=[ + *(normalize_chat_message(m) for m in request), + normalize_chat_message(message), + ], + message_sources=[ + *( + ChatCompletionsMessageSource(exchange=exchange, request_index=i) + for i in range(len(request)) + ), + source, + ], + tools=deepcopy(exchange.request.get("tools")), + chat_template=exchange.request.get("chat_template"), + chat_template_kwargs=deepcopy(exchange.request.get("chat_template_kwargs")), + ) + original._validate_history_sources(history) + trace = original._TraceBuilder() + value = original._tokenize_exact_projected_chat_history( + history, tokenizer=bound, _trace=trace + ) + if value is None or value.tokens != [*prompt, *output]: + raise ValueError("Native sampled source cannot be represented exactly") + _require_exact_chat_source_edges(history, value, trace.trace, bound) + return _Span(source, key, value, len(prompt)) + + +def _join( + history: ChatCompletionsHistory, run: list[_Span], bound: Tokenizer +) -> TokenizedHistory | None: + if len(run) == 1: + return run[0].value + last = run[-1].value.history + assert isinstance(last, ChatCompletionsHistory) + # Token nesting does not authorize reassigning a captured message's source. + # Use the final request's real view only when it also matches the canonical + # message prefix in which these source encounters occurred. + if not _equal_with_nan(last.messages, history.messages[: len(last.messages)]): + return None + selected = {span.key for span in run} + message_sources = list(last.message_sources) + for i, source in enumerate(history.message_sources[: len(last.messages)]): + if ( + source is not None + and original._source_is_sampled(source) + and original._sampled_source_key(source) in selected + ): + message_sources[i] = source + scoped = last.model_copy(update={"message_sources": message_sources}) + original._validate_history_sources(scoped) + tokens = list(run[-1].value.tokens) + flags = [TokenFlag.EXACT] * len(tokens) + logprobs = [math.nan] * len(tokens) + keys: list[original._SampledSourceKey | None] = [None] * len(tokens) + sources: dict[original._SampledSourceKey, object] = {} + previous_end = 0 + for span in run: + start, end = span.start, len(span.value.tokens) + if start < previous_end or tokens[:end] != span.value.tokens: + raise ValueError("Native sampled chain lost complete conditioning") + flags[start:end] = span.value.flags[start:end] + logprobs[start:end] = span.value.logprobs[start:end] + keys[start:end] = [span.key] * (end - start) + sources[span.key] = span.source + previous_end = end + value = TokenizedHistory( + history=scoped, + model=run[-1].value.model, + tokens=tokens, + logprobs=logprobs, + flags=flags, + ) + trace = original._HistoryTokenizationTrace(keys, sources) + _require_exact_chat_source_edges(scoped, value, trace, bound) + return value + + +def tokenize_native( + trajectory: Trajectory, *, model: str | None, base_model: str | None +) -> TokenizedMultiHistoryTrajectory: + """Construct complete native Chat sources in canonical encounter order. + + Only SAMPLED first-occurrence-before-finite-filter loss is represented. + Nonsampled request gaps are EXACT context, without invented assistant roles, + OUTPUT flags or rendered STOP tails. No renderer is invoked and no ordinary + exception is caught. Complete source validation precedes returning any value. + """ + exchanges = trajectory.exchanges + if ( + not exchanges.chat_completions + or exchanges.completions + or exchanges.responses + or exchanges.messages + or trajectory.messages_and_choices + or trajectory.additional_histories + or trajectory.tools is not None + ): + raise ValueError( + "Native sampled representation requires unmixed Chat Completions sources" + ) + histories: list[ChatCompletionsHistory] = [] + for history in trajectory.histories(model=model): + if not isinstance(history, ChatCompletionsHistory) or not history.model: + raise ValueError( + "Native sampled representation requires selected Chat histories" + ) + histories.append(history) + if not histories: + raise ValueError( + "Native sampled representation requires selected Chat histories" + ) + selected_models = {h.model for h in histories} + expected: dict[ + tuple[str, original._SampledSourceKey], ChatCompletionsMessageSource + ] = {} + identities: set[tuple[str, str, int]] = set() + for exchange in exchanges.chat_completions: + if exchange.model not in selected_models: + continue + if not exchange.model or not exchange.response.id: + raise ValueError("Native sampled source identity is unavailable") + for choice in exchange.response.choices: + identity = (exchange.model, exchange.response.id, choice.index) + if identity in identities: + raise ValueError("Native sampled source identity is duplicated") + identities.add(identity) + source = ChatCompletionsMessageSource( + exchange=exchange, choice_index=choice.index + ) + expected[(exchange.model, original._sampled_source_key(source))] = source + encountered: list[ + tuple[ChatCompletionsHistory, list[ChatCompletionsMessageSource]] + ] = [] + seen: set[tuple[str, original._SampledSourceKey]] = set() + for history in histories: + assert history.model is not None + original._validate_history_sources(history) + sources: dict[original._SampledSourceKey, ChatCompletionsMessageSource] = {} + for message, source in zip( + history.messages, history.message_sources, strict=True + ): + if ( + message.get("role") != "assistant" + or source is None + or not original._source_is_sampled(source) + ): + continue + key = original._sampled_source_key(source) + authority = expected.get((history.model, key)) + if ( + not isinstance(source.exchange, ChatCompletionsExchange) + or authority is None + or source.exchange is not authority.exchange + or history.model != source.exchange.model + or not original._source_covers_complete_sampled_message(message, source) + ): + raise ValueError( + "Native sampled source projection is not complete and unchanged" + ) + sources.setdefault(key, source) + seen.add((history.model, key)) + if not sources: + raise ValueError( + "Native sampled history contains no supported sampled sources" + ) + encountered.append((history, list(sources.values()))) + if not expected or seen != expected.keys(): + raise ValueError("Native sampled source inventory is incomplete") + + resolved: dict[str, Tokenizer] = {} + assembled: list[TokenizedHistory] = [] + for history, source_list in encountered: + selected_model = history.model + assert selected_model is not None + if selected_model not in resolved: + config = original._tokenizer_config(selected_model, None) + if base_model is not None and config.base_model != base_model: + raise ValueError( + "Native STOP authority differs from the requested base model" + ) + bound = original._load_tokenizer(config) + if bound is None: + raise ValueError("Native sampled STOP authority is unavailable") + resolved[selected_model] = bound + bound = resolved[selected_model] + runs: list[list[_Span]] = [] + for source in source_list: + span = _singleton(source, bound) + previous = runs[-1][-1] if runs else None + if ( + previous is not None + and len(previous.value.tokens) <= span.start + and span.value.tokens[: len(previous.value.tokens)] + == previous.value.tokens + ): + runs[-1].append(span) + else: + runs.append([span]) + for run in runs: + joined = _join(history, run, bound) + assembled.extend( + [joined] if joined is not None else [span.value for span in run] + ) + return TokenizedMultiHistoryTrajectory(trajectory=trajectory, histories=assembled) diff --git a/tests/unit/trajectories/test_sampled_native.py b/tests/unit/trajectories/test_sampled_native.py new file mode 100644 index 000000000..ea2f619dd --- /dev/null +++ b/tests/unit/trajectories/test_sampled_native.py @@ -0,0 +1,581 @@ +from __future__ import annotations + +from copy import deepcopy +from dataclasses import replace +from datetime import datetime, timedelta +import math +import pickle +import struct +from typing import Any, cast + +from openai.types.chat import ChatCompletion +from openai.types.responses import Response +import pytest + +import art +from art import trajectories as tr +from art.preprocessing.dynamo_tokens import COMPLETION_LOGPROBS_KEY +from art.trajectories import _parallel, _sampled_native, _tokenize +from art.trajectories._serialization import _equal_with_nan + + +class StopOnly: + eos_token_id = 9 + all_special_tokens: list[str] = [] + special_tokens_map: dict[str, str] = {} + + def apply_chat_template(self, *args: Any, **kwargs: Any) -> Any: + raise AssertionError("Native sampled construction must not render") + + def __call__(self, text: str, **kwargs: Any) -> list[int]: + raise AssertionError("This numeric STOP authority must not encode content") + + +def recorded( + *, + terminal_tool: bool = False, + nested: bool = True, + repeated: bool = False, + model: str = "policy", + lp: float = -0.2, +) -> tr.Trajectory: + messages: list[dict[str, Any]] = [{"role": "user", "content": "public question"}] + prompts = [[1], [1, 20, 21, 99, 2], [1, 20, 21, 99, 2, 30, 31, 9, 3]] + if not nested: + prompts[-1][6] = 32 + outputs = [[20, 21], [30, 31], [40, 9]] + replies = [ + {"role": "assistant", "content": "public truncated answer"}, + { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call-public", + "type": "function", + "function": {"name": "lookup", "arguments": '{"key":"public"}'}, + } + ], + }, + {"role": "assistant", "content": "public final answer"}, + ] + exchanges = [] + for i in range(2 if terminal_tool else 3): + response = ChatCompletion.model_validate( + { + "id": f"public-{i}", + "object": "chat.completion", + "created": i, + "model": model, + "choices": [ + { + "index": 0, + "finish_reason": "length" if i == 0 else "stop", + "message": replies[i], + "prompt_token_ids": prompts[i], + "token_ids": outputs[i], + "logprobs": { + "content": [ + { + "token": f"token_id:{token}", + "logprob": lp, + "bytes": [], + "top_logprobs": [], + } + for token in outputs[i] + ] + }, + } + ], + } + ) + now = datetime(2026, 1, 1) + timedelta(seconds=i) + exchanges.append( + tr.ChatCompletionsExchange( + request=tr.ChatCompletionsRequest( + model=model, + messages=cast(Any, deepcopy(messages)), + chat_template_kwargs={ + "enable_thinking": False, + "preserve_thinking": not (repeated and i == 1), + }, + ), + response=response, + start_time=now, + end_time=now, + ) + ) + messages.append(deepcopy(replies[i])) + messages.append({"role": "user", "content": f"public follow-up {i}"}) + return tr.Trajectory( + reward=0.75, + metrics={"retained": 2}, + metadata={"public": True}, + exchanges=tr.TrajectoryExchanges(chat_completions=exchanges), + ) + + +@pytest.fixture +def authority(monkeypatch: pytest.MonkeyPatch) -> list[str]: + calls: list[str] = [] + + def config(model: str, base: str | None) -> _tokenize._TokenizerConfig: + assert base is None + calls.append(model) + return _tokenize._TokenizerConfig(model, "public-revision:" + model) + + monkeypatch.setattr(_tokenize, "_tokenizer_config", config) + monkeypatch.setattr(_tokenize, "_load_tokenizer", lambda config: StopOnly()) + monkeypatch.setattr(_parallel, "_cpu_capacity", lambda: 2) + monkeypatch.setattr(_parallel, "_supports_processes", lambda **_: False) + return calls + + +def source_terms(t: tr.Trajectory, model: str | None = None) -> list[tuple]: + """Independent original native inventory in canonical encounter order.""" + rows = [] + for history in t.histories(model=model): + assert isinstance(history, tr.ChatCompletionsHistory) + seen = set() + for message, source in zip( + history.messages, history.message_sources, strict=True + ): + if ( + message.get("role") != "assistant" + or source is None + or source.choice_index is None + ): + continue + identity = (id(source.exchange), source.choice_index) + if identity in seen: + continue + seen.add(identity) + assert isinstance(source.exchange, tr.ChatCompletionsExchange) + choice = next( + choice + for choice in source.exchange.response.choices + if choice.index == source.choice_index + ) + extra = choice.model_extra or {} + prompt, output = extra["prompt_token_ids"], extra["token_ids"] + lp = extra.get(COMPLETION_LOGPROBS_KEY) + if lp is None: + assert ( + choice.logprobs is not None and choice.logprobs.content is not None + ) + lp = [entry.logprob for entry in choice.logprobs.content] + for i, (token, prob) in enumerate(zip(output, lp, strict=True)): + rows.append( + ( + history.model, + tuple([*prompt, *output[:i]]), + token, + prob, + choice.finish_reason != "length" + and i == len(output) - 1 + and token == 9, + ) + ) + return rows + + +def result_terms(value: tr.TokenizedMultiHistoryTrajectory) -> list[tuple]: + return [ + (h.model, tuple(h.tokens[:i]), token, lp, bool(flag & tr.TokenFlag.STOP)) + for h in value.histories + for i, (token, lp, flag) in enumerate( + zip(h.tokens, h.logprobs, h.flags, strict=True) + ) + if flag & tr.TokenFlag.SAMPLED + ] + + +def claims(rows: list[tuple]) -> list[tuple]: + seen = set() + result = [] + for model, prompt, token, lp, stop in rows: + key = (model, prompt, token) + first = key not in seen + seen.add(key) + try: + finite = math.isfinite(struct.unpack("!f", struct.pack("!f", lp))[0]) + except OverflowError: + finite = False + result.append((key, first, first and finite, stop)) + return result + + +@pytest.mark.parametrize("terminal_tool", [False, True]) +@pytest.mark.parametrize("nested", [False, True]) +async def test_tool_synthetic_stop_uses_native_authority_without_rendering( + authority: list, + terminal_tool: bool, + nested: bool, + monkeypatch: pytest.MonkeyPatch, +) -> None: + t = recorded(terminal_tool=terminal_tool, nested=nested) + before = t.model_dump() + + def fail(*args: Any, **kwargs: Any) -> Any: + raise AssertionError( + "The explicit native route must not call ordinary tokenization" + ) + + monkeypatch.setattr(tr.Trajectory, "tokenize", fail) + value = (await art.tokenize_sampled([t], representation="native"))[0] + assert value.trajectory is t + assert _equal_with_nan(before, t.model_dump()) + assert result_terms(value) == source_terms(t) + assert len(value.histories) == (2 if not terminal_tool and not nested else 1) + assert authority == ["policy"] + # The tool response has no sampled terminator: no synthetic sampled STOP + # or renderer-owned OUTPUT tail may be invented for it. + tool_rows = [row for row in result_terms(value) if row[2] in {30, 31}] + assert len(tool_rows) == 2 and all(not row[-1] for row in tool_rows) + required = ( + tr.TokenFlag.EXACT + | tr.TokenFlag.SAMPLED + | tr.TokenFlag.ASSISTANT + | tr.TokenFlag.OUTPUT + ) + assert all( + flag == tr.TokenFlag.EXACT or flag & required == required + for h in value.histories + for flag in h.flags + ) + + +@pytest.mark.parametrize("lp", [-0.2, math.nan, -math.inf, 1e100]) +async def test_repeated_sources_keep_first_ownership_before_float32_filter( + authority: list, + lp: float, +) -> None: + t = recorded(repeated=True, lp=lp) + original = source_terms(t) + assert len(original) > 6 + value = (await art.tokenize_sampled([t], representation="native"))[0] + assert _equal_with_nan(result_terms(value), original) + assert claims(result_terms(value)) == claims(original) + masks = tr.first_occurrence_masks(value.histories, where=tr.TokenFlag.SAMPLED) + actual = [ + take + for h, mask in zip(value.histories, masks, strict=True) + for flag, take in zip(h.flags, mask, strict=True) + if flag & tr.TokenFlag.SAMPLED + ] + assert actual == [row[1] for row in claims(original)] + + +async def test_two_models_with_same_ids_and_evidence_do_not_collide( + authority: list, +) -> None: + a, b = recorded(model="a"), recorded(model="b") + a.exchanges.chat_completions.extend(b.exchanges.chat_completions) + value = (await art.tokenize_sampled([a], model="*", representation="native"))[0] + assert result_terms(value) == source_terms(a, "*") + assert set(authority) == {"a", "b"} + + +async def test_multiple_nonpositional_choice_indices_keep_complete_inventory( + authority: list, +) -> None: + exchange = recorded().exchanges.chat_completions[0] + data = exchange.response.model_dump() + first = data["choices"][0] + first["index"] = 7 + second = deepcopy(first) + second["index"] = 3 + second["token_ids"] = [22, 23] + for entry, token in zip(second["logprobs"]["content"], [22, 23], strict=True): + entry["token"] = f"token_id:{token}" + data["choices"] = [first, second] + exchange.response = ChatCompletion.model_validate(data) + t = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[exchange])) + value = (await art.tokenize_sampled([t], representation="native"))[0] + assert len(value.histories) == 2 and result_terms(value) == source_terms(t) + indices = set() + for h in value.histories: + assert isinstance(h.history, tr.ChatCompletionsHistory) + indices.update( + source.choice_index + for source in h.history.message_sources + if source is not None and source.choice_index is not None + ) + assert indices == {3, 7} + + +@pytest.mark.parametrize("first_lp", [math.nan, 1e100]) +async def test_nonfinite_first_source_owns_edge_before_later_finite_source( + authority: list, + first_lp: float, +) -> None: + a = recorded(lp=first_lp).exchanges.chat_completions[0] + b = recorded(lp=-0.2).exchanges.chat_completions[0] + b.response.id = "distinct-finite-source" + b.start_time += timedelta(seconds=1) + b.end_time += timedelta(seconds=1) + t = tr.Trajectory(exchanges=tr.TrajectoryExchanges(chat_completions=[a, b])) + value = (await art.tokenize_sampled([t], representation="native"))[0] + assert len(value.histories) == 2 + assert tr.first_occurrence_masks(value.histories, where=tr.TokenFlag.SAMPLED) == [ + [False, True, True], + [False, False, False], + ] + assert value.histories[1].logprobs[1:] == [-0.2, -0.2] + assert not any(row[2] for row in claims(result_terms(value))) + + +async def test_mixed_protocols_refuse_before_loading_stop_authority( + authority: list, +) -> None: + t = recorded() + now = datetime(2026, 1, 1) + response = Response.model_validate( + { + "id": "public-response", + "object": "response", + "created_at": 0, + "model": "policy", + "output": [], + "parallel_tool_calls": True, + "tool_choice": "auto", + "tools": [], + } + ) + t.exchanges.responses.append( + tr.ResponsesExchange( + request={"model": "policy", "input": "public"}, + response=response, + start_time=now, + end_time=now, + ) + ) + with pytest.raises(ValueError, match="unmixed Chat"): + await art.tokenize_sampled([t], representation="native") + assert authority == [] + + +async def test_packed_completion_logprobs_and_compact_roundtrip( + authority: list, +) -> None: + t = recorded() + for exchange in t.exchanges.chat_completions: + choice = exchange.response.choices[0] + assert choice.model_extra is not None + choice.model_extra[COMPLETION_LOGPROBS_KEY] = [-0.2] * len( + choice.model_extra["token_ids"] + ) + choice.logprobs = None + packed = tr.compact_dump(t) + restored = tr.compact_validate(packed, type=tr.Trajectory) + value = (await art.tokenize_sampled([restored], representation="native"))[0] + assert result_terms(value) == source_terms(restored) + roundtrip = tr.compact_validate( + tr.compact_dump(value), type=tr.TokenizedMultiHistoryTrajectory + ) + assert _equal_with_nan(value.model_dump(), roundtrip.model_dump()) + + +@pytest.mark.parametrize("thinking", [False, True, None]) +@pytest.mark.parametrize("reasoning_key", ["reasoning", "reasoning_content"]) +async def test_native_mode_does_not_infer_thinking_or_rewrite_literal_content( + authority: list, + thinking: bool | None, + reasoning_key: str, +) -> None: + t = recorded() + message = { + "role": "assistant", + "content": "public literal nested ", + reasoning_key: "recorded structured reasoning", + } + first = t.exchanges.chat_completions[0] + data = first.response.model_dump() + data["choices"][0]["message"] = message + first.response = ChatCompletion.model_validate(data) + for i, exchange in enumerate(t.exchanges.chat_completions): + kwargs = exchange.request["chat_template_kwargs"] + if thinking is None: + kwargs.pop("enable_thinking") + else: + kwargs["enable_thinking"] = thinking + if i: + exchange.request["messages"][1] = cast(Any, deepcopy(message)) + before = t.model_dump() + value = (await art.tokenize_sampled([t], representation="native"))[0] + assert result_terms(value) == source_terms(t) + assert _equal_with_nan(before, t.model_dump()) + + +async def test_complete_logprob_ids_remain_valid_without_separate_output_ids( + authority: list, +) -> None: + t = recorded() + for exchange in t.exchanges.chat_completions: + extra = exchange.response.choices[0].model_extra + assert extra is not None + extra.pop("token_ids") + value = (await art.tokenize_sampled([t], representation="native"))[0] + assert [row[2] for row in result_terms(value)] == [20, 21, 30, 31, 40, 9] + + +@pytest.mark.parametrize( + "bad", + [ + "prompt", + "output", + "lp", + "lp_length", + "packed_absent", + "packed_length", + "duplicate", + "additional", + "legacy", + "edited", + ], +) +async def test_incomplete_or_edited_authority_refuses( + authority: list, + monkeypatch: pytest.MonkeyPatch, + bad: str, +) -> None: + t = recorded() + choice = t.exchanges.chat_completions[0].response.choices[0] + assert choice.model_extra is not None + if bad == "prompt": + choice.model_extra.pop("prompt_token_ids") + elif bad == "output": + choice.model_extra["token_ids"] = [999] + elif bad == "lp": + choice.logprobs = None + elif bad == "lp_length": + assert choice.logprobs is not None and choice.logprobs.content is not None + choice.logprobs.content.pop() + elif bad == "packed_absent": + choice.model_extra[COMPLETION_LOGPROBS_KEY] = None + elif bad == "packed_length": + choice.model_extra[COMPLETION_LOGPROBS_KEY] = [-0.2] + elif bad == "duplicate": + t.exchanges.chat_completions.append(deepcopy(t.exchanges.chat_completions[0])) + elif bad == "additional": + t = tr.Trajectory( + additional_histories=[ + tr.LegacyHistory( + model="policy", + messages_and_choices=[{"role": "user", "content": "public"}], + ) + ] + ) + elif bad == "legacy": + t = tr.Trajectory(messages_and_choices=[{"role": "user", "content": "public"}]) + elif bad == "edited": + histories = t.histories() + assert isinstance(histories[0], tr.ChatCompletionsHistory) + histories[0].messages[-1]["content"] = "edited public source" + monkeypatch.setattr( + tr.Trajectory, "histories", lambda self, **kwargs: histories + ) + before = t.model_dump() + with pytest.raises((ValueError, TypeError)): + await art.tokenize_sampled([t], representation="native") + assert _equal_with_nan(before, t.model_dump()) + + +@pytest.mark.parametrize("bad", ["tokens", "lp", "stop", "ownership"]) +async def test_constructed_output_is_independently_checked( + authority: list, + monkeypatch: pytest.MonkeyPatch, + bad: str, +) -> None: + original = _tokenize._tokenize_exact_projected_chat_history + + def corrupt(*args: Any, **kwargs: Any) -> Any: + value = original(*args, **kwargs) + assert value is not None + if bad == "tokens": + value.tokens[0] += 1 + elif bad == "lp": + value.logprobs[-1] -= 1 + elif bad == "stop": + value.flags[-1] ^= tr.TokenFlag.STOP + else: + kwargs["_trace"].trace.source_keys[-1] = None + return value + + monkeypatch.setattr(_tokenize, "_tokenize_exact_projected_chat_history", corrupt) + with pytest.raises((ValueError, AssertionError)): + await art.tokenize_sampled([recorded()], representation="native") + + +async def test_group_process_dispatch_rebinds_original_objects( + authority: list, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(_parallel, "_supports_processes", lambda **_: True) + monkeypatch.setattr(_parallel, "_processes_enabled", lambda _: True) + options = [] + + async def process_map(payloads: list[bytes], trajectories: list, **_: Any) -> list: + options.extend(pickle.loads(payload)[1] for payload in payloads) + return [ + _parallel._deserialize_process_result( + _parallel._tokenize_process_payload(payload), source + ) + for payload, source in zip(payloads, trajectories, strict=True) + ] + + monkeypatch.setattr(_parallel, "_ordered_process_map", process_map) + t = recorded(repeated=True) + group = tr.TrajectoryGroup([t], metadata={"order": 7}, metrics={"retained": 2}) + result = (await art.tokenize_sampled([group], representation="native"))[0] + assert result.metadata == group.metadata and result.metrics == group.metrics + value = result.trajectories[0] + assert value.trajectory is t and result_terms(value) == source_terms(t) + assert all(option.sampled and option.native_sampled for option in options) + exchanges = {id(e) for e in t.exchanges.chat_completions} + for h in value.histories: + assert isinstance(h.history, tr.ChatCompletionsHistory) + assert all( + id(s.exchange) in exchanges + for s in h.history.message_sources + if s is not None + ) + + +@pytest.mark.parametrize( + "change", + [ + {"sampled": False}, + {"multi_history": False}, + {"reconcile_text_equivalent_tokenizations": True}, + {"chat_template": "override"}, + {"chat_template_kwargs": {}}, + ], +) +def test_native_process_options_reject_overrides(authority: list, change: dict) -> None: + options = _parallel._ProcessOptions( + True, False, None, None, None, None, sampled=True, native_sampled=True + ) + with pytest.raises(ValueError, match="unmodified sampled options"): + _parallel._tokenize_process_payload( + pickle.dumps((recorded(), replace(options, **change))) + ) + assert authority == [] + + +async def test_native_model_selection_and_empty_containers(authority: list) -> None: + assert await art.tokenize_sampled([], representation="native") == [] + t = recorded(model="a") + t.exchanges.chat_completions.extend(recorded(model="b").exchanges.chat_completions) + value = ( + await art.tokenize_sampled( + [t], model="b", base_model="b", representation="native" + ) + )[0] + assert result_terms(value) == source_terms(t, "b") and authority == ["b"] + with pytest.raises(ValueError, match="requested base"): + await art.tokenize_sampled( + [t], model="b", base_model="a", representation="native" + ) + with pytest.raises(ValueError, match="Unknown sampled representation"): + await art.tokenize_sampled([], representation=cast(Any, "unknown")) From 3c28a3fe913928b7aa09f059437d8b6d6eddb748 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Fri, 25 Sep 2026 18:06:03 +0000 Subject: [PATCH 2/2] Normalize native tool views and clarify sampled authority errors --- src/art/trajectories/_sampled_native.py | 18 +-- .../unit/trajectories/test_sampled_native.py | 125 ++++++++++++++++++ 2 files changed, 131 insertions(+), 12 deletions(-) diff --git a/src/art/trajectories/_sampled_native.py b/src/art/trajectories/_sampled_native.py index e8dd07073..2795b2aaf 100644 --- a/src/art/trajectories/_sampled_native.py +++ b/src/art/trajectories/_sampled_native.py @@ -21,8 +21,8 @@ Trajectory, ) from . import _tokenize as original -from ._history import normalize_chat_message -from ._sampled import _require_exact_chat_source_edges +from ._history import _TOOLS, normalize_chat_message +from ._sampled import _load_sampled_stop_tokenizer, _require_exact_chat_source_edges from ._serialization import _equal_with_nan @@ -75,7 +75,7 @@ def _singleton(source: ChatCompletionsMessageSource, bound: Tokenizer) -> _Span: ), source, ], - tools=deepcopy(exchange.request.get("tools")), + tools=deepcopy(_TOOLS.validate_python(exchange.request.get("tools"))), chat_template=exchange.request.get("chat_template"), chat_template_kwargs=deepcopy(exchange.request.get("chat_template_kwargs")), ) @@ -238,15 +238,9 @@ def tokenize_native( selected_model = history.model assert selected_model is not None if selected_model not in resolved: - config = original._tokenizer_config(selected_model, None) - if base_model is not None and config.base_model != base_model: - raise ValueError( - "Native STOP authority differs from the requested base model" - ) - bound = original._load_tokenizer(config) - if bound is None: - raise ValueError("Native sampled STOP authority is unavailable") - resolved[selected_model] = bound + resolved[selected_model] = _load_sampled_stop_tokenizer( + selected_model, base_model=base_model + ) bound = resolved[selected_model] runs: list[list[_Span]] = [] for source in source_list: diff --git a/tests/unit/trajectories/test_sampled_native.py b/tests/unit/trajectories/test_sampled_native.py index ea2f619dd..c3b921ff9 100644 --- a/tests/unit/trajectories/test_sampled_native.py +++ b/tests/unit/trajectories/test_sampled_native.py @@ -579,3 +579,128 @@ async def test_native_model_selection_and_empty_containers(authority: list) -> N ) with pytest.raises(ValueError, match="Unknown sampled representation"): await art.tokenize_sampled([], representation=cast(Any, "unknown")) + + +@pytest.mark.parametrize("terminal_tool", [False, True]) +@pytest.mark.parametrize("extension", [False, True]) +async def test_native_request_tools_use_canonical_normalization( + authority: list, terminal_tool: bool, extension: bool +) -> None: + source = recorded(terminal_tool=terminal_tool) + tool: dict[str, Any] = { + "type": "function", + "function": { + "name": "lookup", + "description": "Public lookup", + "parameters": {"type": "object", "properties": {"key": {"type": "string"}}}, + }, + } + if extension: + tool["x_vendor"] = {"public": True} + tool["function"]["x_vendor"] = "public" + for exchange in source.exchanges.chat_completions: + exchange.request["tools"] = cast(Any, deepcopy([tool])) + before = source.model_dump() + canonical = source.histories() + value = (await art.tokenize_sampled([source], representation="native"))[0] + assert len(value.histories) == 1 + assert isinstance(value.histories[0].history, tr.ChatCompletionsHistory) + assert isinstance(canonical[-1], tr.ChatCompletionsHistory) + assert value.histories[0].history.tools == canonical[-1].tools + assert result_terms(value) == source_terms(source) + assert claims(result_terms(value)) == claims(source_terms(source)) + assert _equal_with_nan(before, source.model_dump()) + roundtrip = tr.compact_validate( + value.compact_dump(), type=tr.TokenizedMultiHistoryTrajectory + ) + assert _equal_with_nan(roundtrip.model_dump(), value.model_dump()) + + +@pytest.mark.parametrize("terminal_tool", [False, True]) +async def test_native_string_stop_marks_complete_sampled_suffix( + authority: list, monkeypatch: pytest.MonkeyPatch, terminal_tool: bool +) -> None: + class TextStop(StopOnly): + def __call__(self, text: str, **kwargs: Any) -> list[int]: + assert text == "END" + assert kwargs == {"add_special_tokens": False} + return [30, 31] + + monkeypatch.setattr(_tokenize, "_load_tokenizer", lambda config: TextStop()) + source = recorded(terminal_tool=terminal_tool) + choice = source.exchanges.chat_completions[1].response.choices[0] + assert choice.model_extra is not None + choice.model_extra["stop_reason"] = "END" + before = source.model_dump() + value = (await art.tokenize_sampled([source], representation="native"))[0] + assert len(value.histories) == 1 + actual = result_terms(value) + assert [r[:-1] for r in actual] == [r[:-1] for r in source_terms(source)] + assert [r[2] for r in actual if r[-1]] == ( + [30, 31] if terminal_tool else [30, 31, 9] + ) + assert _equal_with_nan(source.model_dump(), before) + + +def test_join_keeps_singletons_when_message_view_differs() -> None: + source = recorded() + history = source.histories()[0] + assert isinstance(history, tr.ChatCompletionsHistory) + sources = [ + s + for s in history.message_sources + if s is not None and s.choice_index is not None + ] + spans = [_sampled_native._singleton(s, StopOnly()) for s in sources] + messages = deepcopy(history.messages) + messages[0]["content"] = "Different public context view" + other_view = history.model_copy(update={"messages": messages}) + before = [span.value.model_dump() for span in spans] + assert _sampled_native._join(other_view, spans, StopOnly()) is None + assert _equal_with_nan([span.value.model_dump() for span in spans], before) + # Direct defensive join control: retaining these complete singletons keeps + # original sampled conditioning and encounter order without editing context. + singles = tr.TokenizedMultiHistoryTrajectory( + trajectory=source, histories=[s.value for s in spans] + ) + assert result_terms(singles) == source_terms(source) + + +@pytest.mark.parametrize("unexpected", [False, True]) +async def test_native_stop_loader_reports_authority_without_base_fallback( + authority: list, monkeypatch: pytest.MonkeyPatch, unexpected: bool +) -> None: + error = ( + RuntimeError("public unrelated failure") + if unexpected + else ValueError("pass base_model explicitly") + ) + + def fail(config: Any) -> Any: + raise error + + monkeypatch.setattr(_tokenize, "_load_tokenizer", fail) + with pytest.raises(type(error)) as caught: + await art.tokenize_sampled([recorded()], representation="native") + if unexpected: + assert caught.value is error + else: + assert caught.value.__cause__ is error + assert "loadable tokenizer model ID" in str(caught.value) + assert "base_model" in str(caught.value) + assert authority == ["policy"] + + +@pytest.mark.parametrize("reasoning_field", ["reasoning", "reasoning_content"]) +async def test_reasoning_stripped_followup_refuses_complete_source_certification( + authority: list, reasoning_field: str +) -> None: + source = recorded() + message = source.exchanges.chat_completions[0].response.choices[0].message + assert message.model_extra is not None + message.model_extra[reasoning_field] = "Public structured reasoning" + before = source.model_dump() + with pytest.raises(ValueError, match="projection is not complete and unchanged"): + await art.tokenize_sampled([source], representation="native") + assert authority == [] + assert _equal_with_nan(before, source.model_dump())