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
219 changes: 219 additions & 0 deletions src/art/trajectories/_render_cache.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,219 @@
"""Conservative eligibility for invocation-local chat-render reuse."""

import inspect
import math
import sys
from types import CodeType
from typing import cast


def _render_context_key(value: object) -> object:
"""Snapshot plain JSON without losing mapping order or scalar types."""
kind = type(value)
if kind in (str, int, bool, type(None)):
return kind, value
if kind is float and math.isfinite(cast(float, value)):
return kind, repr(value)
if kind is list:
return kind, tuple(_render_context_key(item) for item in cast(list, value))
if kind is dict and all(type(key) is str for key in cast(dict, value)):
return kind, tuple(
(key, _render_context_key(item)) for key, item in cast(dict, value).items()
)
raise TypeError("Not a plain JSON rendering context")


def _code_contains(code, candidate, name):
return (code is candidate and code.co_name == name) or any(
_code_contains(child, candidate, name)
for child in code.co_consts
if isinstance(child, CodeType)
)


def cacheable_chat_template(tokenizer, template, tools, kwargs, messages) -> bool:
"""Admit a finite, nonmutating Jinja subset, never arbitrary renderer purity.

Callers must key all effective context and keep this cache local to one
tokenization. Unsupported syntax, context or overrides retain normal rendering.
No Transformers import or global eligibility memo is performed here.
"""
try:
base_module = sys.modules.get("transformers.tokenization_utils_base")
chat = sys.modules.get("transformers.utils.chat_template_utils")
if base_module is None or chat is None or type(template) is not str:
return False
base = base_module.PreTrainedTokenizerBase
cls = type(tokenizer)
module = sys.modules.get(cls.__module__)
if (
not isinstance(tokenizer, base)
or base_module.render_jinja_template is not chat.render_jinja_template
or not cls.__module__.startswith("transformers.")
or getattr(module, cls.__name__, None) is not cls
or type(tokenizer.chat_template) not in (str, type(None))
or inspect.getattr_static(cls, "special_tokens_map")
is not inspect.getattr_static(base, "special_tokens_map")
):
return False
for name in ("apply_chat_template", "get_chat_template"):
method = getattr(tokenizer, name)
if method.__self__ is not tokenizer or method.__func__ is not getattr(
base, name
):
return False
# Check before the stock property calls str(): custom objects can hide
# mutable state even if their resulting special-token strings look plain.
added_token = sys.modules["tokenizers"].AddedToken
special = tokenizer._special_tokens_map
if type(special) is not dict or any(
type(key) is not str or type(value) not in (str, type(None), added_token)
for key, value in special.items()
):
return False
_render_context_key([messages, tools, kwargs, tokenizer.special_tokens_map])
if type(kwargs) is not dict or kwargs.get("continue_final_message"):
return False

from jinja2 import Environment, Undefined, defaults, nodes
from jinja2.runtime import Context, LoopContext
from jinja2.sandbox import ImmutableSandboxedEnvironment
from jinja2.utils import Namespace

compiled = chat._compile_jinja_template(template)
env = compiled.environment
if (
type(env) is not ImmutableSandboxedEnvironment
or env.undefined is not Undefined
or env.finalize is not None
or env.context_class is not Context
or env.concat is not Environment.concat
or env.autoescape is not False
or env.is_async
):
return False
tree = env.parse(template) # HF's environment understands {% generation %}.
parents = {
child: node
for node in (tree, *tree.find_all(nodes.Node))
for child in node.iter_child_nodes()
}
macros = {node.name for node in tree.find_all(nodes.Macro)}
filters = set("default length tojson trim items string safe".split())
tests = set("string iterable mapping none undefined true false defined".split())
methods = set(
"get items keys values startswith endswith strip lstrip rstrip split rsplit replace lower upper join".split()
)
# A method object can expose an address when printed or aliased. Permit
# only direct calls of the explicitly nonmutating methods above.
method_names = {
name
for cls in (str, dict, list, tuple, int, float, LoopContext)
for name in dir(cls)
if callable(getattr(cls, name))
}
structural = set(
"Template Output TemplateData Const Name Getattr Getitem Slice If For Assign AssignBlock NSRef Macro Call CallBlock Keyword Filter Test Compare Operand And Or Not Neg Pos Add Sub Mul Div FloorDiv Mod Pow Concat List Tuple Dict Pair CondExpr Break Continue ExtensionAttribute".split()
)
compiler = getattr(
chat,
"_cached_compile_jinja_template",
chat._compile_jinja_template.__wrapped__,
)
for node in (tree, *tree.find_all(nodes.Node)):
parent = parents.get(node)
direct_call = isinstance(parent, nodes.Call) and parent.node is node
if type(node).__name__ not in structural:
return False
if isinstance(node, nodes.Name):
if node.name in {"self", "super", "caller"}:
return False
if node.name in env.globals or node.name in macros:
if not direct_call:
return False
if isinstance(node, nodes.Getattr):
if node.attr.startswith("_") or (
node.attr in method_names and not direct_call
):
return False
if isinstance(node, nodes.ExtensionAttribute) and not direct_call:
return False
if isinstance(node, nodes.Getitem):
# Dynamic lookup can fetch a callable or renderer/context object.
parts = (
(node.arg.start, node.arg.stop, node.arg.step)
if isinstance(node.arg, nodes.Slice)
else (node.arg,)
)
parts = tuple(
part.node if isinstance(part, (nodes.Neg, nodes.Pos)) else part
for part in parts
)
if any(
part is not None
and not (isinstance(part, nodes.Const) and type(part.value) is int)
for part in parts
):
return False
if isinstance(node, nodes.Call):
target = node.node
if node.dyn_args is not None or node.dyn_kwargs is not None:
return False
if isinstance(target, nodes.Name):
if target.name in macros:
continue
helper = env.globals.get(target.name)
if target.name == "namespace" and helper is Namespace:
continue
if target.name == "raise_exception" and _code_contains(
compiler.__code__,
getattr(helper, "__code__", None),
"raise_exception",
):
continue
return False
if isinstance(target, nodes.Getattr) and target.attr in methods:
continue
if isinstance(target, nodes.ExtensionAttribute) and isinstance(
parent, nodes.CallBlock
):
extension = env.extensions.get(target.identifier)
helper = getattr(
getattr(extension, target.name, None), "__func__", None
)
if target.name == "_generation_support" and _code_contains(
compiler.__code__,
getattr(helper, "__code__", None),
"_generation_support",
):
continue
return False
if isinstance(node, (nodes.Filter, nodes.Test)):
if isinstance(node, nodes.Test):
if (
node.name not in tests
or env.tests.get(node.name)
is not defaults.DEFAULT_TESTS[node.name]
):
return False
elif node.name not in filters:
return False
elif node.name == "tojson":
if not _code_contains(
compiler.__code__,
getattr(env.filters.get(node.name), "__code__", None),
"tojson",
):
return False
elif (
env.filters.get(node.name)
is not defaults.DEFAULT_FILTERS[node.name]
):
return False
if node.name == "items" and not (
isinstance(parent, nodes.For) and parent.iter is node
):
return False # Otherwise a generator's repr can expose identity.
return True
except Exception:
return False
118 changes: 110 additions & 8 deletions src/art/trajectories/_tokenize.py
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,7 @@
)
from ._history import _model_matches
from ._protocols import Exchange
from ._render_cache import _render_context_key, cacheable_chat_template

_TOKEN_ID = re.compile(r"token_id:(\d+)$")
_WARNED_PREFIX_RETOKENIZATION = False
Expand Down Expand Up @@ -92,6 +93,73 @@ def __call__(
) -> str: ...


class _PrefixChatRenderCache:
"""Reuse exact baseline prefixes within one history's rendering context.

Probes still locate every assistant span. Equal completed text does not prove
equal generation prompts, so changed prefixes never reuse baseline renders.
Keep one baseline plus bounded prefix deltas, not quadratic rendered strings.
"""

_MAX_BYTES = 8 * 1024 * 1024
_MAX_ENTRIES = 1024

def __init__(self, render: _ChatRender) -> None:
self.render = render
self.context: tuple[object, ...] | None = None
self.settings: object = None
self.text = ""
self.prefixes: dict[tuple[int, bool], tuple[int, str]] = {}
self.bytes = 0

def for_messages(
self, messages: list[dict[str, Any]], text: str, *, settings: object = None
) -> _ChatRender:
try:
context = tuple(_render_context_key(message) for message in messages)
except (TypeError, RecursionError):
return self.render
if self.context is None or settings != self.settings:
self.context, self.text = context, text
self.settings = settings
self.prefixes.clear()
self.bytes = 0
common = 0
for original, current in zip(self.context, context):
if original != current:
break
common += 1

def render(
selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool
) -> str:
count = len(selected_messages)
if count > common or any(
original is not current
for original, current in zip(messages, selected_messages)
):
return self.render(
selected_messages, add_generation_prompt=add_generation_prompt
)
key = count, add_generation_prompt
if key in self.prefixes:
prefix, tail = self.prefixes[key]
return self.text[:prefix] + tail
value = self.render(
selected_messages, add_generation_prompt=add_generation_prompt
)
if len(self.prefixes) < self._MAX_ENTRIES:
prefix = _common_prefix_length(self.text, value)
tail = value[prefix:]
size = 256 + 4 * len(tail)
if self.bytes + size <= self._MAX_BYTES:
self.prefixes[key] = prefix, tail
self.bytes += size
return value

return render


class _TokenChatRender(Protocol):
def __call__(
self,
Expand Down Expand Up @@ -4530,13 +4598,11 @@ def raw_render(
)
)

def render_text(
def render_normalized_text(
selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool
) -> str:
value = resolved_tokenizer.apply_chat_template(
normalize_tool_call_arguments_for_chat_template(
selected_messages, template
),
selected_messages,
tools=history.tools,
tokenize=False,
add_generation_prompt=add_generation_prompt,
Expand All @@ -4547,17 +4613,53 @@ def render_text(
raise TypeError("Chat template did not render text")
return value

def render_text(
selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool
) -> str:
return render_normalized_text(
normalize_tool_call_arguments_for_chat_template(
selected_messages, template
),
add_generation_prompt=add_generation_prompt,
)

prefix_render_cache = _PrefixChatRenderCache(render_normalized_text)

def segmented_render(
selected_messages: list[dict[str, Any]], *, add_generation_prompt: bool
) -> tuple[list[int], list[bool]]:
try:
text = render_text(
selected_messages, add_generation_prompt=add_generation_prompt
)
if cacheable_chat_template(
resolved_tokenizer, template, history.tools, kwargs, selected_messages
):
# Normalization is message-local. Only the admitted nonmutating
# renderer may share its normalized messages between prefixes.
selected_messages = normalize_tool_call_arguments_for_chat_template(
selected_messages, template
)
text = render_normalized_text(
selected_messages, add_generation_prompt=add_generation_prompt
)
span_render = prefix_render_cache.for_messages(
selected_messages,
text,
settings=_render_context_key(
[
history.tools,
kwargs,
getattr(resolved_tokenizer, "special_tokens_map"),
]
),
)
else:
text = render_text(
selected_messages, add_generation_prompt=add_generation_prompt
)
span_render = render_text
spans = _assistant_char_spans(
selected_messages,
text,
render_text,
span_render,
add_generation_prompt=add_generation_prompt,
)
except (TypeError, KeyError):
Expand Down
Loading
Loading