From 8279c2a4c5f7517c4b5620d65152d97bf0ce80d4 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 00:12:18 +0000 Subject: [PATCH 1/2] Preserve prompt masks across merged whitespace tokens --- src/art/trajectories/_tokenize.py | 43 +++++++++++---- tests/unit/trajectories/test_tokenize.py | 67 ++++++++++++++++++++++++ 2 files changed, 101 insertions(+), 9 deletions(-) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 58fe012df..007380922 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -480,7 +480,11 @@ def _assistant_stop_masks( def _translate_token_mask( - source: Sequence[int], target: Sequence[int], mask: Sequence[bool] + source: Sequence[int], + target: Sequence[int], + mask: Sequence[bool], + *, + tokenizer: Tokenizer | None = None, ) -> list[bool]: """Translate a token mask across a prefix replacement without guessing.""" @@ -490,13 +494,28 @@ def _translate_token_mask( translated = [False] * len(target) mapped = [False] * len(source) - for source_start, target_start, length in SequenceMatcher( + decode = getattr(tokenizer, "decode", None) + for tag, start, end, target_start, target_end in SequenceMatcher( None, source, target, autojunk=False - ).get_matching_blocks(): - for offset in range(length): - if mask[source_start + offset]: - translated[target_start + offset] = True - mapped[source_start + offset] = True + ).get_opcodes(): + if tag == "equal": + translated[target_start:target_end] = mask[start:end] + mapped[start:end] = [True] * (end - start) + elif tag == "replace" and callable(decode) and any(mask[start:end]): + # A merged whitespace token inherits the mask when any of its + # characters participate. Multiple target tokens require a uniform + # source mask. This translates flags without changing served IDs. + if target_end - target_start != 1 and not all(mask[start:end]): + continue + kwargs = dict(skip_special_tokens=False, clean_up_tokenization_spaces=False) + text = decode(list(source[start:end]), **kwargs) + if text.isspace() and text == decode( + list(target[target_start:target_end]), **kwargs + ): + translated[target_start:target_end] = [True] * ( + target_end - target_start + ) + mapped[start:end] = [True] * (end - start) if any(selected and not retained for selected, retained in zip(mask, mapped)): raise ValueError( "Cannot preserve assistant boundaries across exact prompt token replacement" @@ -4785,10 +4804,16 @@ def source_matches_context(source: object) -> bool: direct_bounds or None, ) assistant_mask = _translate_token_mask( - canonical_rendered, rendered, canonical_assistant_mask + canonical_rendered, + rendered, + canonical_assistant_mask, + tokenizer=resolved_tokenizer, ) output_mask = _translate_token_mask( - canonical_rendered, rendered, canonical_output_mask + canonical_rendered, + rendered, + canonical_output_mask, + tokenizer=resolved_tokenizer, ) stop_mask = _translate_token_mask(canonical_rendered, rendered, canonical_stop_mask) length_stop_mask = _translate_token_mask( diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 5858aa2bb..69f727247 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -646,6 +646,73 @@ def test_exact_output_boundaries_survive_prefix_order_drift_and_length_stop() -> history.tokenize(tokenizer=tokenizer, chat_template="explicit override") +def test_length_stop_preserves_served_prompt_with_retokenized_demonstration() -> None: + class Tokenizer(_BoundaryTokenizer): + def decode(self, token_ids: list[int], **kwargs: object) -> str: + return "".join({2: "\n", 30: "\n\n"}.get(t, str(t)) for t in token_ids) + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + add_generation_prompt: bool, + **kwargs: object, + ) -> list[int]: + tokens = super().apply_chat_template( + messages, add_generation_prompt=add_generation_prompt, **kwargs + ) + # The first newline belongs to the generation prefix; the next + # belongs to the assistant. The served prompt merges them. + return [*tokens, 2] if add_generation_prompt else tokens + + exchange = _chat_exchange([1, 30, 3, 2], [4], offset=1) + exchange.response.choices[0].finish_reason = "length" + trajectory = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange]) + ) + tokenizer = Tokenizer( + ("user", [1]), ("assistant", [2, 2]), ("user", [3]), ("assistant", [2, 4, 9]) + ) + tokenized = trajectory.tokenize(tokenizer=tokenizer) + assert tokenized.tokens == [1, 30, 3, 2, 4] + assert not tokenized.flags[1] & (tr.TokenFlag.SAMPLED | tr.TokenFlag.OUTPUT) + assert tokenized.flags[4] == _SAMPLED_ASSISTANT_OUTPUT + assert tokenized.logprobs[4] == -0.4 + + +@pytest.mark.parametrize( + "text,mask", + [("different", [True, True]), ("\n\n", [True, False]), ("\ufffd", [True, True])], +) +def test_mask_translation_rejects_changed_text_or_ambiguous_boundaries( + text: str, mask: list[bool] +) -> None: + from art.trajectories._tokenize import _translate_token_mask + + tokenizer = cast( + tr.Tokenizer, + SimpleNamespace( + decode=lambda tokens, **kwargs: "\n\n" if tokens == [1, 2] else text + ), + ) + with pytest.raises(ValueError, match="Cannot preserve assistant boundaries"): + _translate_token_mask([1, 2], [3, 4], mask, tokenizer=tokenizer) + + +@pytest.mark.parametrize( + "mask", [[False, True], [True, False], [True, True], [False, False]] +) +def test_merged_whitespace_token_inherits_mask_of_its_characters( + mask: list[bool], +) -> None: + from art.trajectories._tokenize import _translate_token_mask + + tokenizer = cast( + tr.Tokenizer, SimpleNamespace(decode=lambda tokens, **kwargs: "\n\n") + ) + assert _translate_token_mask([1, 2], [3], mask, tokenizer=tokenizer) == [any(mask)] + + def test_exact_length_boundary_with_multiple_parts_and_prefix_drift() -> None: first = _chat_exchange([1], [2, 9]) second = _chat_exchange([1, 2, 9, 3], [4, 5], offset=1) From 2283a7cbfed95b58c929f057789a544624d32d15 Mon Sep 17 00:00:00 2001 From: Brad Hilton Date: Wed, 16 Sep 2026 00:42:29 +0000 Subject: [PATCH 2/2] Preserve trajectory group types through tensorization --- src/art/trajectories/__init__.py | 14 ++++++++++++++ .../unit/trajectories/test_tensorized_models.py | 16 ++++++++++++++++ 2 files changed, 30 insertions(+) diff --git a/src/art/trajectories/__init__.py b/src/art/trajectories/__init__.py index 730c3d096..43626c952 100644 --- a/src/art/trajectories/__init__.py +++ b/src/art/trajectories/__init__.py @@ -1350,6 +1350,20 @@ def compact_dump(self) -> CompactTrajectoryPayload: return dump_tokenized_trajectory_group(self) + @overload + def tensorize( + self: TokenizedTrajectoryGroup[TokenizedTrajectory], + *, + device: torch.device | str | None = None, + ) -> TensorizedTrajectoryGroup[TensorizedTrajectory]: ... + + @overload + def tensorize( + self: TokenizedTrajectoryGroup[TokenizedMultiHistoryTrajectory], + *, + device: torch.device | str | None = None, + ) -> TensorizedTrajectoryGroup[TensorizedMultiHistoryTrajectory]: ... + def tensorize( self, *, device: torch.device | str | None = None ) -> ( diff --git a/tests/unit/trajectories/test_tensorized_models.py b/tests/unit/trajectories/test_tensorized_models.py index b8336dacc..26cbcc85b 100644 --- a/tests/unit/trajectories/test_tensorized_models.py +++ b/tests/unit/trajectories/test_tensorized_models.py @@ -6,6 +6,7 @@ import pickle import subprocess import sys +from typing import assert_type from openai.types.chat import ChatCompletion import pytest @@ -253,6 +254,7 @@ def test_trajectory_and_group_tensorize_retain_mutable_sources() -> None: trajectories=[tokenized], ) tensorized_group = tokenized_group.tensorize() + assert_type(tensorized_group, tr.TensorizedTrajectoryGroup[tr.TensorizedTrajectory]) assert tensorized_group.trajectory_group is source_group assert tensorized_group.trajectories[0].trajectory is trajectory group_metrics: dict[str, float | int | bool] = {"batch": 2} @@ -261,6 +263,20 @@ def test_trajectory_and_group_tensorize_retain_mutable_sources() -> None: assert source_group.metrics == {"batch": 2} assert source_group.metadata == {"name": "updated"} + multi_group = tr.TokenizedTrajectoryGroup[tr.TokenizedMultiHistoryTrajectory]( + trajectory_group=source_group, + trajectories=[ + tr.TokenizedMultiHistoryTrajectory( + trajectory=trajectory, histories=[tokenized] + ) + ], + ).tensorize() + assert_type( + multi_group, tr.TensorizedTrajectoryGroup[tr.TensorizedMultiHistoryTrajectory] + ) + assert multi_group.trajectory_group is source_group + assert len(multi_group.trajectories[0].histories) == 1 + def test_tensorized_compact_round_trips_inferred_and_typed() -> None: value = _tokenized_history().tensorize()