Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 14 additions & 0 deletions src/art/trajectories/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
) -> (
Expand Down
43 changes: 34 additions & 9 deletions src/art/trajectories/_tokenize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand All @@ -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"
Expand Down Expand Up @@ -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(
Expand Down
16 changes: 16 additions & 0 deletions tests/unit/trajectories/test_tensorized_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
import pickle
import subprocess
import sys
from typing import assert_type

from openai.types.chat import ChatCompletion
import pytest
Expand Down Expand Up @@ -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}
Expand All @@ -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()
Expand Down
67 changes: 67 additions & 0 deletions tests/unit/trajectories/test_tokenize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
Loading