diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 92824eec0..c65bd7b1b 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -4182,6 +4182,23 @@ def _tokenize_exact_projected_chat_history( if _source_stop_evidence(source, source_key)[0] == "length": boundary = (length_stop_boundaries or {}).get(source_key) next_prompt = _chat_source_prompt_tokens(sampled_sources[index + 1]) + if boundary is not None and next_prompt is not None: + rendered_boundary = [*boundary.tail, *boundary.following] + native_boundary = next_prompt[end:] + extra = len(native_boundary) - len(rendered_boundary) + decode = getattr(tokenizer, "decode", None) + if ( + extra > 0 + and native_boundary[extra:] == rendered_boundary + and callable(decode) + and decode(native_boundary[:extra]).isspace() + ): + # Services may insert whitespace before a truncated turn's + # proven stop tail. Keep those served, nonsampled tokens. + boundary = _RenderedLengthStopBoundary( + tail=(*native_boundary[:extra], *boundary.tail), + following=boundary.following, + ) boundary_end = ( end + len(boundary.tail) + len(boundary.following) if boundary is not None @@ -5706,6 +5723,36 @@ def differing_span(probe: Sequence[int]) -> tuple[int, int] | None: multi_generation_response or len(parts) != 1 or parts[0][0] != "content" ): start = generation_start + if _sampled_stop_suffix( + full_exact, + source=source, + source_key=_sampled_source_key(source), + tokenizer=resolved_tokenizer, + ): + # Adjacent assistants can share a role mask. Prove this message's + # end before replacing its rendered closing markup and stop. + completed = probe_render( + messages[: message_index + 1], add_generation_prompt=False + ) + rendered_completed = ( + canonical_render_to_rendered(completed) + if completed is not None + else None + ) + if ( + rendered_completed is not None + and rendered[: len(rendered_completed)] == rendered_completed + ): + tail_mask, tail_stops = _assistant_stop_masks( + rendered_completed, + assistant_mask[: len(rendered_completed)], + resolved_tokenizer, + ) + tail_end = end + while tail_end < len(tail_mask) and tail_mask[tail_end]: + tail_end += 1 + if tail_end > end and tail_stops[tail_end - 1]: + end = tail_end replacements.append( ( start, diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index f1ff287e3..ac5d8b14b 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -5831,6 +5831,155 @@ def apply_chat_template( assert tokenized.flags[1] == (_SAMPLED_ASSISTANT_OUTPUT) +@pytest.mark.parametrize("reasoning", ["a", "a§"]) +def test_sampled_tail_does_not_consume_adjacent_assistant( + reasoning: str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr( + "art.trajectories._tokenize._WARNED_PREFIX_RETOKENIZATION", False + ) + + class Tokenizer(_CharacterTemplateTokenizer): + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tokenize: bool = True, + add_generation_prompt: bool, + **kwargs: object, + ) -> str | list[int]: + text = "".join( + str(message.get("reasoning") or "") + + str(message.get("content") or "") + + ("§" if message["role"] == "assistant" else "") + for message in messages + ) + return self._encode(text) if tokenize else text + + tokenizer = Tokenizer() + prompt = tokenizer._encode("question") + output = [7001, *tokenizer._encode(reasoning[1:] + "b§")] + first = _chat_exchange(prompt, output) + first.request["messages"] = [{"role": "user", "content": "question"}] + data = first.response.model_dump(mode="python") + data["choices"][0]["message"] = { + "role": "assistant", + "reasoning": reasoning, + "content": "b", + } + first.response = ChatCompletion.model_validate(data) + second_output = tokenizer._encode("cd§") + second = _chat_exchange([*prompt, *output], second_output, offset=1) + second.request["messages"] = [ + *first.request["messages"], + cast(ChatCompletionMessageParam, data["choices"][0]["message"]), + ] + second.response.choices[0].message.content = "cd" + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[first, second]) + ).chat_completions_history() + history.chat_template = "rerender" + + with pytest.warns(UserWarning, match="preserved the original sampled token IDs"): + tokenized = history.tokenize(tokenizer=tokenizer) + + assert tokenized.tokens == [*prompt, *output, *second_output] + assert tokenized.logprobs[len(prompt) :] == [ + -token / 10 for token in [*output, *second_output] + ] + assert tokenized.flags[len(prompt) :] == [ + _SAMPLED_ASSISTANT_OUTPUT + | (tr.TokenFlag.STOP if index == len(tokens) - 1 else 0) + for tokens in (output, second_output) + for index in range(len(tokens)) + ] + + +def test_complete_sampled_tool_call_replaces_rendered_closing_markup( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + "art.trajectories._tokenize._WARNED_PREFIX_RETOKENIZATION", False + ) + + class ToolTokenizer(_CharacterTemplateTokenizer): + def __call__(self, text: str, **kwargs: object) -> dict[str, object]: + return {"input_ids": self._encode(text)} + + def apply_chat_template( + self, + messages: list[dict[str, Any]], + *, + tokenize: bool = True, + add_generation_prompt: bool, + **kwargs: object, + ) -> str | list[int]: + text = "" + for message in messages: + if message["role"] == "user": + text += f"{message['content']}" + else: + function = message["tool_calls"][0]["function"] + text += ( + f"{message.get('reasoning', '')}" + f"" + f"{function['arguments']}§" + ) + if add_generation_prompt: + text += "" + return self._encode(text) if tokenize else text + + tokenizer = ToolTokenizer() + prompt = tokenizer._encode("turn 0") + output = tokenizer._encode('thought\n{"x":1}§') + exchange = _chat_exchange(prompt, output) + data = exchange.response.model_dump(mode="python") + data["choices"][0].update( + finish_reason="tool_calls", + message={ + "role": "assistant", + "content": None, + "reasoning": "thought\n", + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": {"name": "lookup", "arguments": '{"x":1}'}, + } + ], + }, + ) + exchange.response = ChatCompletion.model_validate(data) + next_prompt = [*prompt, *output, *tokenizer._encode("continue")] + following = _chat_exchange(next_prompt, output, offset=1) + following.request["messages"] = [ + {"role": "user", "content": "turn 0"}, + cast(ChatCompletionMessageParam, data["choices"][0]["message"]), + {"role": "user", "content": "continue"}, + ] + following_data = following.response.model_dump(mode="python") + following_data["choices"][0].update( + finish_reason="tool_calls", message=data["choices"][0]["message"] + ) + following.response = ChatCompletion.model_validate(following_data) + history = art.Trajectory( + exchanges=TrajectoryExchanges(chat_completions=[exchange, following]) + ).chat_completions_history() + history.chat_template = "rerender" + + with pytest.warns(UserWarning, match="preserved the original sampled token IDs"): + tokenized = history.tokenize(tokenizer=tokenizer) + + assert tokenized.tokens == [*next_prompt, *output] + assert tokenized.logprobs[len(prompt) : len(prompt) + len(output)] == [ + -token / 10 for token in output + ] + assert tokenized.flags[len(prompt) : len(prompt) + len(output)] == [ + _SAMPLED_ASSISTANT_OUTPUT | (tr.TokenFlag.STOP if token == 9 else 0) + for token in output + ] + + def test_template_change_preserves_complete_exact_sampled_suffix() -> None: exchange = _chat_exchange([1], [2, 3]) history = art.Trajectory( @@ -8312,6 +8461,8 @@ def tokenize(*, allow_missing: bool) -> list[TokenizedResult]: ("7", None), ("0", "missing_stop"), ("0", "wrong_stop"), + ("0", "extra_boundary_newline"), + ("0", "extra_boundary_text"), ("0", "changed_sampled_token"), ("7", "changed_sampled_token"), ("7", "unrendered_sampled_token"), @@ -8389,13 +8540,20 @@ def apply_chat_template( prompt.remove(9) elif corruption == "wrong_stop": prompt[prompt.index(9)] = tokenizer._encode("!")[0] + elif corruption in {"extra_boundary_newline", "extra_boundary_text"}: + boundary_start = len(first_choice.model_extra["prompt_token_ids"]) + len( + first_choice.model_extra["token_ids"] + ) + prompt[boundary_start:boundary_start] = tokenizer._encode( + "\n" if corruption == "extra_boundary_newline" else "!" + ) elif corruption in {"changed_sampled_token", "unrendered_sampled_token"}: prefix = first_choice.model_extra["prompt_token_ids"] prompt[len(prefix)] = tokenizer._encode("!")[0] trajectory = art.Trajectory( exchanges=TrajectoryExchanges(chat_completions=exchanges) ) - if corruption in {"missing_stop", "wrong_stop"}: + if corruption in {"missing_stop", "wrong_stop", "extra_boundary_text"}: with pytest.raises(ValueError, match="Could not uniquely locate"): _tokenize_trajectory_with_trace(trajectory, tokenizer=tokenizer) return