Skip to content
Draft
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
21 changes: 20 additions & 1 deletion src/art/trajectories/_tokenize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
Expand Down Expand Up @@ -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,
Expand Down
43 changes: 39 additions & 4 deletions tests/unit/trajectories/test_prefix_render_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"},
Expand All @@ -150,15 +153,24 @@ 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)
probe[1]["content"] = "B"
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
Expand All @@ -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) %}<think>{{ message.reasoning_content }}</think>{% endif %}
{% for call in message.tool_calls or [] %}{% set tool_call = call.function %}<call>{{ tool_call.name }}({% for k, v in tool_call.arguments.items() %}{{ k }}={{ v|tojson }};{% endfor %})</call>{% endfor %}
Expand Down Expand Up @@ -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.
Loading