diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 61626c6fc..a51e2b436 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -113,7 +113,12 @@ def __init__(self, render: _ChatRender) -> None: self.bytes = 0 def for_messages( - self, messages: list[dict[str, Any]], text: str, *, settings: object = None + self, + messages: list[dict[str, Any]], + text: str, + *, + settings: object = None, + add_generation_prompt: bool | None = None, ) -> _ChatRender: try: context = tuple(_render_context_key(message) for message in messages) @@ -130,10 +135,23 @@ def for_messages( break common += 1 + rendered_generation_prompt = add_generation_prompt + def render( selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool ) -> str: count = len(selected_messages) + # This exact full-context render was just evaluated by the caller. + # A changed probe still proves every other generation/completion. + if ( + count == len(messages) + and add_generation_prompt == rendered_generation_prompt + and all( + original is current + for original, current in zip(messages, selected_messages) + ) + ): + return text if count > common or any( original is not current for original, current in zip(messages, selected_messages) @@ -5277,6 +5295,7 @@ def segmented_render( span_render = prefix_render_cache.for_messages( selected_messages, text, + add_generation_prompt=add_generation_prompt, settings=_render_context_key( [ history.tools, diff --git a/tests/unit/trajectories/test_prefix_render_cache.py b/tests/unit/trajectories/test_prefix_render_cache.py index 05760040a..a285edb96 100644 --- a/tests/unit/trajectories/test_prefix_render_cache.py +++ b/tests/unit/trajectories/test_prefix_render_cache.py @@ -129,7 +129,10 @@ def render(selected_messages, *, add_generation_prompt): assert cache.for_messages([{"content": value}], "unchanged") is render -def test_later_generation_split_is_recomputed_when_completed_suffix_is_equal(): +@pytest.mark.parametrize("seed_completed", [False, True]) +def test_later_generation_split_is_recomputed_when_completed_suffix_is_equal( + seed_completed, +): messages = [ {"role": "user", "content": "x"}, {"role": "assistant", "content": "A"}, @@ -150,7 +153,11 @@ def render(selected_messages, *, add_generation_prompt): tokenization._assistant_char_spans( messages, original, - cache.for_messages(messages, original), + cache.for_messages( + messages, + original, + add_generation_prompt=False if seed_completed else None, + ), add_generation_prompt=False, ) probe = deepcopy(messages) @@ -158,7 +165,12 @@ def render(selected_messages, *, add_generation_prompt): text = render(probe, add_generation_prompt=False) assert text == original.replace(">A!", ">B!") actual = tokenization._assistant_char_spans( - probe, text, cache.for_messages(probe, text), add_generation_prompt=False + probe, + text, + cache.for_messages( + probe, text, add_generation_prompt=False if seed_completed else None + ), + add_generation_prompt=False, ) expected = tokenization._assistant_char_spans( probe, text, render, add_generation_prompt=False @@ -168,6 +180,29 @@ def render(selected_messages, *, add_generation_prompt): assert text[start:end] == "ELLO!" +def test_completed_render_reuse_requires_same_messages_and_generation(): + calls = [] + + def render(selected_messages, *, add_generation_prompt): + calls.append((deepcopy(selected_messages), add_generation_prompt)) + return json.dumps(selected_messages) + str(add_generation_prompt) + + messages = [{"role": "assistant", "content": "original"}] + cache = tokenization._PrefixChatRenderCache(render) + cache.for_messages(messages, render(messages, add_generation_prompt=False)) + probe = [{"role": "assistant", "content": "changed"}] + text = render(probe, add_generation_prompt=False) + cached = cache.for_messages(probe, text, add_generation_prompt=False) + before = len(calls) + assert cached(probe[:], add_generation_prompt=False) == text + assert len(calls) == before + assert cached(probe, add_generation_prompt=True) != text + assert len(calls) == before + 1 + assert cached(deepcopy(probe), add_generation_prompt=False) == text + assert len(calls) == before + 2 + assert cached([], add_generation_prompt=False) == "[]False" + + _TEMPLATE = """{% for message in messages %}<{{ message.role }}>{{ message.content or '' }} {% if message.reasoning_content and (enable_thinking is not defined or enable_thinking) %}{{ message.reasoning_content }}{% endif %} {% for call in message.tool_calls or [] %}{% set tool_call = call.function %}{{ tool_call.name }}({% for k, v in tool_call.arguments.items() %}{{ k }}={{ v|tojson }};{% endfor %}){% endfor %} @@ -395,4 +430,4 @@ def counted(*args, **kwargs): if tools: assert cached_calls <= calls - turns * (turns - 1) else: - assert cached_calls == calls # Plain text already avoids tool probes. + assert cached_calls == calls - 1 # The completed full render is reused.