diff --git a/dev/trainer_rank.py b/dev/trainer_rank.py index 0bfc45eb8..65aa2c6ff 100644 --- a/dev/trainer_rank.py +++ b/dev/trainer_rank.py @@ -5,9 +5,9 @@ import torch import torch.distributed as dist from trainer_rank_support import load_random_checkpoints -from transformers import AutoTokenizer import typer +from art import get_tokenizer from art.trainer_rank import AdamParams, ForwardInput, TrainerRank @@ -34,9 +34,7 @@ def main( from art.megatron import train as megatron_train - tokenizer = cast( - Any, AutoTokenizer.from_pretrained(model, trust_remote_code=True) - ) + tokenizer = cast(Any, get_tokenizer(model, trust_remote_code=True)) inputs: list[ForwardInput[torch.Tensor, None, None, None]] = [] rows = load_dataset("roneneldan/TinyStories", split="train", streaming=True) for row in islice(rows, samples): diff --git a/dev/trainer_rank_landing_acceptance.py b/dev/trainer_rank_landing_acceptance.py index 95513fd1d..dadd80c97 100644 --- a/dev/trainer_rank_landing_acceptance.py +++ b/dev/trainer_rank_landing_acceptance.py @@ -197,13 +197,11 @@ def _check_corpus_tokenizer(corpus: dict[str, Any], model: str) -> dict[str, Any compares vocabulary sizes; it fails the cell loudly on any difference. """ - from transformers import AutoTokenizer + from art import get_tokenizer corpus_model = str(corpus.get("tokenizer_model") or "") - model_tokenizer = AutoTokenizer.from_pretrained(model, trust_remote_code=True) - corpus_tokenizer = AutoTokenizer.from_pretrained( - corpus_model, trust_remote_code=True - ) + model_tokenizer = get_tokenizer(model, trust_remote_code=True) + corpus_tokenizer = get_tokenizer(corpus_model, trust_remote_code=True) sample = corpus["groups"][0]["histories"][0]["tokens"][:256] problems = [] if len(model_tokenizer) != len(corpus_tokenizer): diff --git a/examples/hn_title_generator/train.py b/examples/hn_title_generator/train.py index f2f0b98d4..2759834fa 100644 --- a/examples/hn_title_generator/train.py +++ b/examples/hn_title_generator/train.py @@ -8,10 +8,10 @@ import openai from openai.types.chat import ChatCompletionMessageParam from openpipe import AsyncOpenPipe -from transformers.models.auto.tokenization_auto import AutoTokenizer from utils import cache, prompt_for_title, pull_data, score_title import art +from art import get_tokenizer from art.local import LocalBackend from art.utils import iterate_dataset, limit_concurrency @@ -37,7 +37,7 @@ def filter_on_length(data: Dataset, max_length: int, tokenizer_name: str) -> Dat print( f"Filtering dataset for max prompt length: {max_length} using tokenizer: {tokenizer_name}" ) - tokenizer = AutoTokenizer.from_pretrained(tokenizer_name) + tokenizer = get_tokenizer(tokenizer_name) def check_length(x): # Ensure 'prompt' is a list of dicts diff --git a/pyproject.toml b/pyproject.toml index 96e66f7e6..82a527972 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -152,7 +152,7 @@ build-backend = "hatchling.build" allow-direct-references = true [tool.hatch.build.targets.wheel] -packages = ["src/art", "src/mp_actors"] +packages = ["src/art", "src/art_inference", "src/mp_actors"] [tool.hatch.build.targets.wheel.force-include] ".agents/skills" = "art/skills" diff --git a/src/art/__init__.py b/src/art/__init__.py index a5fdc3f19..fb8fe9638 100644 --- a/src/art/__init__.py +++ b/src/art/__init__.py @@ -71,6 +71,7 @@ from .model import Model, TrainableModel from .pipeline_tuner import PipelineAutotuneConfig, PipelineRuntimeConfig from .serverless import ServerlessBackend +from .tokenizer import get_tokenizer from .trajectories import ( Trajectory, TrajectoryGroup, @@ -116,6 +117,7 @@ "PIPELINE_RL_METRIC_DEFINITIONS", "PIPELINE_RL_SCORE_METRICS", "get_megatron_runtime_config", + "get_tokenizer", "init_megatron_runtime_config", "ServerlessBackend", "ServerlessTrainResult", diff --git a/src/art/local/backend.py b/src/art/local/backend.py index 51cb18358..cc5f72b8e 100644 --- a/src/art/local/backend.py +++ b/src/art/local/backend.py @@ -16,9 +16,9 @@ from typing import TYPE_CHECKING, Any, AsyncIterator, Iterable, Literal, cast import warnings +from art.tokenizer import get_tokenizer from art.utils.chat_template import ( chat_template_with_preserved_thinking, - configure_preserved_thinking_chat_template, ) from art.utils.lifecycle import ( PROCESS_SHUTDOWN_TIMEOUT_SECONDS, @@ -40,7 +40,6 @@ from pydantic import BaseModel, ConfigDict import torch from tqdm import auto as tqdm -from transformers import AutoTokenizer from transformers.tokenization_utils_base import PreTrainedTokenizerBase from typing_extensions import Self @@ -196,7 +195,9 @@ def _apply_configured_chat_template( ) -> None: chat_template = _configured_chat_template_value(internal_config) if chat_template is not None: - tokenizer.chat_template = chat_template + tokenizer.chat_template = cast( + str, chat_template_with_preserved_thinking(chat_template) + ) def _model_support_handler( @@ -241,9 +242,9 @@ def _apply_configured_chat_template_server_args( chat_template = _model_support_default_chat_template( base_model, internal_config ) - if chat_template is None and _should_probe_preserve_thinking_template(base_model): + if chat_template is None and base_model is not None: try: - tokenizer = AutoTokenizer.from_pretrained(base_model) + tokenizer = get_tokenizer(base_model) except (OSError, ValueError) as error: warnings.warn( f"Could not load {base_model!r} to configure prior-thinking " @@ -252,13 +253,14 @@ def _apply_configured_chat_template_server_args( stacklevel=2, ) else: - default = getattr(tokenizer, "chat_template", None) - preserved = chat_template_with_preserved_thinking(default) - if preserved != default: - chat_template = cast(str, preserved) + template = getattr(tokenizer, "chat_template", None) + if isinstance(template, str) and ("{{" in template or "{%" in template): + chat_template = template if chat_template is None: return - server_args.setdefault("chat_template", chat_template) + server_args.setdefault( + "chat_template", chat_template_with_preserved_thinking(chat_template) + ) if chat_template_content_format := internal_config.get( "chat_template_content_format" ): @@ -269,13 +271,6 @@ def _apply_configured_chat_template_server_args( config_dict["server_args"] = server_args -def _should_probe_preserve_thinking_template(base_model: str | None) -> bool: - if base_model is None: - return False - model_name = base_model.rstrip("/").rsplit("/", 1)[-1] - return model_name.startswith(("Qwen3-", "Qwen3.5-")) - - def _tokenizer_cache_key( base_model: str, internal_config: dev.InternalModelConfig, @@ -287,12 +282,7 @@ def _tokenizer_cache_key( def _load_training_tokenizer(base_model: str) -> PreTrainedTokenizerBase: - return cast( - PreTrainedTokenizerBase, - configure_preserved_thinking_chat_template( - AutoTokenizer.from_pretrained(base_model) - ), - ) + return get_tokenizer(base_model) class LocalBackend: diff --git a/src/art/megatron/dsv4/tokenizer.py b/src/art/megatron/dsv4/tokenizer.py index 03d44b643..1c59a19f5 100644 --- a/src/art/megatron/dsv4/tokenizer.py +++ b/src/art/megatron/dsv4/tokenizer.py @@ -1,11 +1,18 @@ from __future__ import annotations +from contextvars import ContextVar import copy from typing import Any from transformers.tokenization_utils_base import PreTrainedTokenizerBase -from art.megatron.dsv4.encoding import encode_messages +from art.megatron.dsv4 import encoding +from art_inference.append_only import patch_deepseek_renderer, preserves_history + +_PRESERVE_HISTORY: ContextVar[bool] = ContextVar( + "art_dsv4_preserve_history", default=True +) +patch_deepseek_renderer(encoding, _PRESERVE_HISTORY.get) DSV4_CHAT_TEMPLATE_MARKER = "deepseek_v4_python_encoder enable_thinking" @@ -35,12 +42,28 @@ def apply_chat_template( tools: list[dict[str, Any]] | None = None, **kwargs: Any, ) -> str | list[int]: + chat_template = kwargs.get("chat_template") + if chat_template is None: + chat_template = self.chat_template + if chat_template != DSV4_CHAT_TEMPLATE_MARKER: + return super().apply_chat_template(messages, tools=tools, **kwargs) thinking = bool(kwargs.get("thinking", False)) or bool( kwargs.get("enable_thinking", False) ) thinking_mode = "thinking" if thinking else "chat" conversation = kwargs.get("conversation", messages) - rendered_messages = list(conversation) + rendered_messages = [ + { + **message, + "reasoning": message.get( + "reasoning", + message.get("reasoning_content", message.get("thinking")), + ), + } + if message.get("role") == "assistant" + else message + for message in conversation + ] if tools: rendered_messages.insert(0, {"role": "system", "tools": tools}) @@ -55,12 +78,17 @@ def apply_chat_template( else: reasoning_effort = "high" - prompt = encode_messages( - rendered_messages, - thinking_mode=thinking_mode, - drop_thinking=kwargs.get("drop_thinking", True), - reasoning_effort=reasoning_effort, - ) + preserve = preserves_history(kwargs) + token = _PRESERVE_HISTORY.set(preserve) + try: + prompt = encoding.encode_messages( + rendered_messages, + thinking_mode=thinking_mode, + drop_thinking=not preserve, + reasoning_effort=reasoning_effort, + ) + finally: + _PRESERVE_HISTORY.reset(token) if not kwargs.get("tokenize", True): return prompt tokenizer_kwargs = { diff --git a/src/art/tinker/prefix_cache.py b/src/art/tinker/prefix_cache.py deleted file mode 100644 index 81c8f4620..000000000 --- a/src/art/tinker/prefix_cache.py +++ /dev/null @@ -1,208 +0,0 @@ -from __future__ import annotations - -from collections import OrderedDict -from dataclasses import dataclass -from typing import Sequence - - -@dataclass(frozen=True) -class PrefixEntry: - rendered_len: int - raw_prefix: tuple[int, ...] - - -@dataclass -class PrefixCacheStats: - max_entries: int - lookups: int = 0 - hits: int = 0 - misses: int = 0 - inserts: int = 0 - replaced_entries: int = 0 - evictions: int = 0 - splits: int = 0 - pruned_nodes: int = 0 - merged_nodes: int = 0 - lru_repairs: int = 0 - - -class _RadixEdge: - __slots__ = ("label", "child") - - def __init__(self, label: tuple[int, ...], child: _RadixNode) -> None: - self.label = label - self.child = child - - -class _RadixNode: - __slots__ = ("entry", "children", "parent", "parent_token") - - def __init__( - self, parent: _RadixNode | None = None, parent_token: int | None = None - ) -> None: - self.entry: PrefixEntry | None = None - self.children: dict[int, _RadixEdge] = {} - self.parent = parent - self.parent_token = parent_token - - -def _common_prefix_len( - tokens: Sequence[int], start: int, label: tuple[int, ...] -) -> int: - max_len = min(len(tokens) - start, len(label)) - i = 0 - while i < max_len and tokens[start + i] == label[i]: - i += 1 - return i - - -class LRUTrieCache: - """LRU-bounded radix trie for token sequence rewrites.""" - - def __init__(self, max_entries: int = 16_384) -> None: - if max_entries <= 0: - raise ValueError("max_entries must be positive") - self._root = _RadixNode() - self._lru: OrderedDict[_RadixNode, None] = OrderedDict() - self._max_entries = max_entries - self.stats = PrefixCacheStats(max_entries=max_entries) - - def lookup(self, rendered_tokens: Sequence[int]) -> PrefixEntry | None: - self.stats.lookups += 1 - node = self._root - idx = 0 - best_node = None - while idx < len(rendered_tokens): - edge = node.children.get(rendered_tokens[idx]) - if edge is None: - break - matched = _common_prefix_len(rendered_tokens, idx, edge.label) - if matched != len(edge.label): - break - idx += matched - node = edge.child - if node.entry is not None: - best_node = node - if best_node is None: - self.stats.misses += 1 - return None - self.stats.hits += 1 - try: - self._lru.move_to_end(best_node) - except KeyError: - self.stats.lru_repairs += 1 - self._lru[best_node] = None - self._lru.move_to_end(best_node) - self._evict() - return best_node.entry - - def insert(self, rendered_prefix: Sequence[int], raw_prefix: Sequence[int]) -> None: - self.stats.inserts += 1 - node = self._root - idx = 0 - while idx < len(rendered_prefix): - token = rendered_prefix[idx] - edge = node.children.get(token) - if edge is None: - child = _RadixNode(parent=node, parent_token=token) - node.children[token] = _RadixEdge(tuple(rendered_prefix[idx:]), child) - node = child - idx = len(rendered_prefix) - break - - matched = _common_prefix_len(rendered_prefix, idx, edge.label) - if matched == len(edge.label): - idx += matched - node = edge.child - continue - - mid = _RadixNode(parent=node, parent_token=token) - self.stats.splits += 1 - old_suffix = edge.label[matched:] - old_child = edge.child - old_child.parent = mid - old_child.parent_token = old_suffix[0] - mid.children[old_suffix[0]] = _RadixEdge(old_suffix, old_child) - edge.label = edge.label[:matched] - edge.child = mid - node = mid - idx += matched - if idx < len(rendered_prefix): - new_token = rendered_prefix[idx] - child = _RadixNode(parent=node, parent_token=new_token) - node.children[new_token] = _RadixEdge( - tuple(rendered_prefix[idx:]), child - ) - node = child - break - - if node.entry is not None: - self.stats.replaced_entries += 1 - node.entry = PrefixEntry( - rendered_len=len(rendered_prefix), raw_prefix=tuple(raw_prefix) - ) - self._lru[node] = None - self._lru.move_to_end(node) - self._evict() - - def _evict(self) -> None: - while len(self._lru) > self._max_entries: - old_node, _ = self._lru.popitem(last=False) - self.stats.evictions += 1 - old_node.entry = None - self._prune(old_node) - - def _prune(self, node: _RadixNode) -> None: - # Collapse empty branches after eviction so the bounded cache stays bounded. - while node.parent is not None: - parent = node.parent - parent_token = node.parent_token - assert parent_token is not None - - if node.entry is None and not node.children: - del parent.children[parent_token] - self.stats.pruned_nodes += 1 - node = parent - continue - - if node.entry is None and len(node.children) == 1: - _, child_edge = next(iter(node.children.items())) - parent_edge = parent.children[parent_token] - parent_edge.label = parent_edge.label + child_edge.label - parent_edge.child = child_edge.child - child_edge.child.parent = parent - child_edge.child.parent_token = parent_token - self.stats.merged_nodes += 1 - node = parent - continue - - break - - def snapshot_stats(self) -> dict[str, int | float]: - hit_rate = self.stats.hits / self.stats.lookups if self.stats.lookups else 0.0 - return { - "enabled": True, - "max_entries": self.stats.max_entries, - "current_entries": len(self._lru), - "node_count": self._node_count(), - "lookups": self.stats.lookups, - "hits": self.stats.hits, - "misses": self.stats.misses, - "hit_rate": hit_rate, - "inserts": self.stats.inserts, - "replaced_entries": self.stats.replaced_entries, - "evictions": self.stats.evictions, - "splits": self.stats.splits, - "pruned_nodes": self.stats.pruned_nodes, - "merged_nodes": self.stats.merged_nodes, - "lru_repairs": self.stats.lru_repairs, - } - - def _node_count(self) -> int: - count = 0 - stack = [self._root] - while stack: - node = stack.pop() - count += 1 - stack.extend(edge.child for edge in node.children.values()) - return count diff --git a/src/art/tinker/server.py b/src/art/tinker/server.py index 1767fa868..0909d58b9 100644 --- a/src/art/tinker/server.py +++ b/src/art/tinker/server.py @@ -29,14 +29,24 @@ from pydantic import BaseModel, Field, SkipValidation, TypeAdapter import tinker from tinker_cookbook import renderers -from tinker_cookbook.tokenizer_utils import get_tokenizer +from tinker_cookbook.tokenizer_utils import Tokenizer as CookbookTokenizer from transformers.tokenization_utils_base import BatchEncoding import uvicorn -from art.tinker.prefix_cache import LRUTrieCache from art.tinker.renderers import get_renderer_name, is_qwen3_dot_family_model +from art.token_prefix import TokenPrefixStore +from art.tokenizer import get_tokenizer from art.types import Message, Tools -from art.utils.chat_template import default_chat_template_kwargs_for_tokenizer +from art.utils.append_only import ( + chat_prefix_observations, + has_renderable_tool_arguments, + output_prefix_observations, + preserves_history, +) +from art.utils.chat_template import ( + default_chat_template_kwargs_for_tokenizer, + normalize_tool_call_arguments_for_chat_template, +) from mp_actors import close_proxy, move_to_child_process @@ -120,7 +130,7 @@ class OpenAICompatibleTinkerServer: port: int | None = None num_workers: int | None = None max_concurrent_sampling_clients: int | None = None - _prefix_cache: LRUTrieCache = field(default_factory=LRUTrieCache) + _prefix_cache: TokenPrefixStore = field(default_factory=TokenPrefixStore) _task: asyncio.Task[None] | None = None _tenants: dict[str, "OpenAICompatibleTinkerServerTenant"] = field( default_factory=dict @@ -366,19 +376,30 @@ async def chat_completions( worker = next(workers) tenant = self._get_request_tenant(request) samplable_model = await tenant.get_samplable_model(body["model"]) + template_kwargs = cast(dict[str, Any], body).get("chat_template_kwargs") + preserve = preserves_history(template_kwargs) + scope = json.dumps( + [id(tenant), samplable_model.base_model, template_kwargs], + sort_keys=True, + ) rendered_prompt_tokens = await worker.prompt_tokens( base_model=samplable_model.base_model, messages=list(body["messages"]), tools=list(body.get("tools", [])) if "tools" in body else None, + chat_template_kwargs=template_kwargs, ) prompt_tokens = rendered_prompt_tokens - prefix_entry = self._prefix_cache.lookup(rendered_prompt_tokens) - if prefix_entry is not None and prefix_entry.rendered_len <= len( + prefix_entry = ( + self._prefix_cache.lookup(scope, rendered_prompt_tokens, None) + if preserve + else None + ) + if prefix_entry is not None and prefix_entry.rendered_length <= len( rendered_prompt_tokens ): prompt_tokens = ( list(prefix_entry.raw_prefix) - + rendered_prompt_tokens[prefix_entry.rendered_len :] + + rendered_prompt_tokens[prefix_entry.rendered_length :] ) try: async with samplable_model.sampling_client() as sampling_client: @@ -407,18 +428,20 @@ async def chat_completions( raise HTTPException(status_code=e.status_code, detail=detail) from e ( chat_completion, - token_discrepancies, - ) = await worker.chat_completion_and_token_discrepancies( + prefixes, + ) = await worker.chat_completion_and_prefixes( base_model=samplable_model.base_model, sample_response=sample_response, model_name=body["model"], - prompt_tokens=len(prompt_tokens), + prompt_tokens=prompt_tokens, + rendered_prompt=rendered_prompt_tokens, + messages=list(body["messages"]), + tools=list(body.get("tools", [])) if "tools" in body else None, + chat_template_kwargs=template_kwargs, ) - for rendered_response_tokens, raw_response_tokens in token_discrepancies: - self._prefix_cache.insert( - rendered_prompt_tokens + rendered_response_tokens, - prompt_tokens + raw_response_tokens, - ) + if preserve: + for rendered, raw, edits in prefixes: + self._prefix_cache.insert(scope, rendered, raw, "content", edits) return chat_completion server_config = uvicorn.Config( @@ -550,14 +573,24 @@ async def prompt_tokens( base_model: str, messages: list[ChatCompletionMessageParam], tools: list[ChatCompletionToolUnionParam] | None, + *, + add_generation_prompt: bool = True, + chat_template_kwargs: dict[str, Any] | None = None, ) -> list[int]: normalized_messages = _normalize_qwen3_dot_messages(base_model, messages) tokenizer = self._get_renderer(base_model).tokenizer - chat_template_kwargs = default_chat_template_kwargs_for_tokenizer(tokenizer) + normalized_messages = normalize_tool_call_arguments_for_chat_template( + normalized_messages, getattr(tokenizer, "chat_template", None) + ) + chat_template_kwargs = { + **default_chat_template_kwargs_for_tokenizer(tokenizer), + **(chat_template_kwargs or {}), + } encoding = tokenizer.apply_chat_template( cast(Any, normalized_messages), tools=cast(Any, tools), - add_generation_prompt=True, + tokenize=True, + add_generation_prompt=add_generation_prompt, **chat_template_kwargs, ) if isinstance(encoding, BatchEncoding): @@ -590,28 +623,69 @@ async def messages_and_choices_prompt_tokens_and_choice_offsets( ) return (result.token_ids, result.choice_offsets) if result is not None else None - async def chat_completion_and_token_discrepancies( + async def chat_completion_and_prefixes( self, base_model: str, sample_response: tinker.SampleResponse, model_name: str, - prompt_tokens: int, - ) -> tuple[ChatCompletion, list[tuple[list[int], list[int]]]]: + prompt_tokens: list[int], + rendered_prompt: list[int], + messages: list[ChatCompletionMessageParam], + tools: list[ChatCompletionToolUnionParam] | None, + chat_template_kwargs: dict[str, Any] | None = None, + ) -> tuple[ChatCompletion, list[tuple[list[int], list[int], tuple[Any, ...]]]]: renderer = self._get_renderer(base_model) choices: list[Choice] = [] - token_discrepancies: list[tuple[list[int], list[int]]] = [] + prefixes = [] for i, sequence in enumerate(sample_response.sequences): assert sequence.logprobs is not None, "Logprobs are required" assert len(sequence.tokens) == len(sequence.logprobs), ( "Tokens and logprobs must have the same length" ) - rendered_response_tokens = renderer.tokenizer.encode( - renderer.tokenizer.decode(sequence.tokens) - ) - if rendered_response_tokens != sequence.tokens: - token_discrepancies.append((rendered_response_tokens, sequence.tokens)) message, _ = renderer.parse_response(sequence.tokens) openai_message = renderer.to_openai_message(message) + if preserves_history(chat_template_kwargs): + prefixes.extend( + output_prefix_observations( + renderer.tokenizer, + rendered_prompt, + prompt_tokens, + sequence.tokens, + ) + ) + if preserves_history( + chat_template_kwargs + ) and has_renderable_tool_arguments(openai_message): + + async def render(assistant: dict[str, Any]) -> list[int]: + return await self.prompt_tokens( + base_model, + [*messages, cast(Any, assistant)], + tools, + add_generation_prompt=False, + chat_template_kwargs=chat_template_kwargs, + ) + + reasoning = { + key: openai_message[key] + for key in ("reasoning", "reasoning_content", "thinking") + if openai_message.get(key) is not None + } + prefixes.extend( + chat_prefix_observations( + renderer.tokenizer, + rendered_prompt, + await render(openai_message), + prompt_tokens, + sequence.tokens, + reasoning_prompt=await render( + {"role": "assistant", "content": "", **reasoning} + ) + if reasoning + else None, + complete=sequence.stop_reason == "stop", + ) + ) tool_calls = ( [ ChatCompletionMessageFunctionToolCall( @@ -635,10 +709,12 @@ async def chat_completion_and_token_discrepancies( Choice( finish_reason=sequence.stop_reason, index=i, - message=ChatCompletionMessage( - content=openai_message.get("content") or None, - role="assistant", - tool_calls=tool_calls, # type: ignore + message=ChatCompletionMessage.model_validate( + { + **openai_message, + "role": "assistant", + "tool_calls": tool_calls, + } ), logprobs=ChoiceLogprobs( content=[ @@ -669,18 +745,19 @@ async def chat_completion_and_token_discrepancies( object="chat.completion", usage=CompletionUsage( completion_tokens=completion_tokens, - prompt_tokens=prompt_tokens, - total_tokens=completion_tokens + prompt_tokens, + prompt_tokens=len(prompt_tokens), + total_tokens=completion_tokens + len(prompt_tokens), ), ), - token_discrepancies, + prefixes, ) def _get_renderer(self, base_model: str) -> renderers.Renderer: if base_model not in self._renderers: self._renderers[base_model] = renderers.get_renderer( name=get_renderer_name(base_model), - tokenizer=get_tokenizer(base_model), + # Cookbook's annotation omits the fast HF tokenizers it accepts. + tokenizer=cast(CookbookTokenizer, get_tokenizer(base_model)), model_name=base_model, ) return self._renderers[base_model] diff --git a/src/art/tinker_native/backend.py b/src/art/tinker_native/backend.py index 8641c3318..879c113ac 100644 --- a/src/art/tinker_native/backend.py +++ b/src/art/tinker_native/backend.py @@ -25,10 +25,13 @@ from openai.types.chat.completion_create_params import CompletionCreateParams from openai.types.completion_usage import CompletionUsage import tinker -from tinker_cookbook import renderers, tokenizer_utils +from tinker_cookbook import renderers +from tinker_cookbook.tokenizer_utils import Tokenizer as CookbookTokenizer import torch import uvicorn +from art.tokenizer import get_tokenizer + from .. import dev from ..adapter_leases import pin_inference_step, pinned_inference_step from ..backend import Backend @@ -715,12 +718,15 @@ async def _build_model_state(self, model: TrainableModel) -> ModelState: service_client = tinker.ServiceClient() rest_client = service_client.create_rest_client() - tokenizer = tokenizer_utils.get_tokenizer(model.base_model) + tokenizer = get_tokenizer(model.base_model) renderer = renderers.get_renderer( name=config.renderer_name, - tokenizer=tokenizer, + # Cookbook's annotation omits the fast HF tokenizers it accepts. + tokenizer=cast(CookbookTokenizer, tokenizer), model_name=model.base_model, ) + if hasattr(renderer, "strip_thinking_from_history"): + setattr(renderer, "strip_thinking_from_history", False) saved_state = model.read_state() or {} tinker_run_ids = list(saved_state.get(STATE_KEY_RUN_IDS, [])) diff --git a/src/art/token_prefix.py b/src/art/token_prefix.py new file mode 100644 index 000000000..0af1a4ce0 --- /dev/null +++ b/src/art/token_prefix.py @@ -0,0 +1,8 @@ +"""Compatibility imports for ART's lightweight inference helpers.""" + +import sys + +from art_inference import token_prefix as _implementation +from art_inference.token_prefix import * # noqa: F403 + +sys.modules[__name__] = _implementation diff --git a/src/art/tokenizer.py b/src/art/tokenizer.py new file mode 100644 index 000000000..f3e26414e --- /dev/null +++ b/src/art/tokenizer.py @@ -0,0 +1,128 @@ +"""The default tokenizer used by ART's renderers and training backends.""" + +from __future__ import annotations + +from copy import copy +import os +import sys +from typing import TYPE_CHECKING, Any, cast + +from .utils.chat_template import configure_preserved_thinking_chat_template + +if TYPE_CHECKING: + from transformers import PreTrainedTokenizerBase + +_KIMI_TOKENIZER_REVISIONS = { + "moonshotai/Kimi-K2-Thinking": "a51ccc050d73dab088bf7b0e2dd9b30ae85a4e55", + "moonshotai/Kimi-K2.5": "2426b45b6af0da48d0dcce71bbce6225e5c73adc", + "moonshotai/Kimi-K2.6": "b5aabbfb20227ed42becbf5541dbffd213942c58", +} +_LLAMA_TEXT_MODELS = { + "meta-llama/Llama-3.1-8B-Instruct", + "meta-llama/Llama-3.1-8B", + "meta-llama/Llama-3.1-70B", + "meta-llama/Llama-3.2-1B", + "meta-llama/Llama-3.2-3B", + "meta-llama/Llama-3.3-70B-Instruct", +} +_LLAMA_CHAT_TOKENIZER = "thinkingmachineslabinc/meta-llama-3-instruct-tokenizer" + + +def get_tokenizer( + base_model: str, *, revision: str | None = None, **kwargs: Any +) -> PreTrainedTokenizerBase: + """Load a model's tokenizer with ART's history-preserving defaults. + + ``revision`` and other keyword arguments go to ``from_pretrained``. Hugging + Face tokenizers are not cached here: configuring one model's template must not + change another model's renderer. Hugging Face still caches downloaded files. + ART's inference and trajectory code cache these instances where appropriate. + Unpinned Llama text models retain Cookbook's public fallback and use its chat + template when the base model has none. Explicit revisions keep native files. + """ + # Tinker suffixes identify a model variant, not a tokenizer revision. Local + # paths may themselves contain colons. + local = os.path.isdir(base_model) + model = base_model if local else base_model.split(":", 1)[0] + registered = sys.modules.get("tinker_cookbook.tokenizer_utils") + is_registered = getattr(registered, "is_tokenizer_registered", None) + if callable(is_registered) and is_registered(base_model): + assert registered is not None + if revision is not None or kwargs: + raise ValueError( + "Registered tokenizer factories do not accept loader options" + ) + return cast( + "PreTrainedTokenizerBase", + configure_preserved_thinking_chat_template( + copy(registered.get_tokenizer(base_model)) + ), + ) + if model.startswith("thinkingmachines/Inkling"): + if revision is not None or kwargs: + raise ValueError("Inkling tokenizers do not accept Hugging Face options") + from tinker_cookbook.tokenizer_utils import get_tokenizer as get_tml_tokenizer + + return cast( + "PreTrainedTokenizerBase", + configure_preserved_thinking_chat_template( + copy(get_tml_tokenizer(base_model)) + ), + ) + + from transformers import AutoTokenizer, PreTrainedTokenizerFast + + loader: Any = AutoTokenizer + if model.startswith("deepseek-ai/DeepSeek-V4-"): + loader = PreTrainedTokenizerFast + elif model.startswith("moonshotai/Kimi-K2"): + # AutoTokenizer can select an incompatible fast backend for Kimi. Its + # custom tokenizer also renders tool declarations as TypeScript. + if kwargs.get("trust_remote_code") is False: + raise ValueError( + "Kimi tokenization requires its repository's custom tokenizer code" + ) + from transformers.dynamic_module_utils import get_class_from_dynamic_module + + revision = revision or _KIMI_TOKENIZER_REVISIONS.get(model) + loader = cast( + "type[PreTrainedTokenizerBase]", + get_class_from_dynamic_module( + "tokenization_kimi.TikTokenTokenizer", + model, + revision=revision, + **kwargs, + ), + ) + kwargs.setdefault( + "trust_remote_code", + os.path.isdir(model) + or os.environ.get("HF_TRUST_REMOTE_CODE", "").lower() in ("1", "true", "yes"), + ) + if revision is not None: + kwargs["revision"] = revision + try: + tokenizer = loader.from_pretrained(model, **kwargs) + except OSError: + if revision is not None or model not in _LLAMA_TEXT_MODELS: + raise + # Tinker users need not have access to Meta's gated repositories. Keep + # Cookbook's public text-tokenizer fallback, but prefer the actual model + # and never apply a model-specific commit to a different repository. + tokenizer = loader.from_pretrained(_LLAMA_CHAT_TOKENIZER, **kwargs) + if ( + revision is None + and model in _LLAMA_TEXT_MODELS + and not getattr(tokenizer, "chat_template", None) + ): + tokenizer.chat_template = loader.from_pretrained( + _LLAMA_CHAT_TOKENIZER, **kwargs + ).chat_template + if model.startswith("deepseek-ai/DeepSeek-V4-"): + from .megatron.dsv4.tokenizer import get_dsv4_tokenizer + + tokenizer = get_dsv4_tokenizer(tokenizer) + return cast( + "PreTrainedTokenizerBase", + configure_preserved_thinking_chat_template(tokenizer), + ) diff --git a/src/art/trajectories/_tokenize.py b/src/art/trajectories/_tokenize.py index 007380922..ccd498775 100644 --- a/src/art/trajectories/_tokenize.py +++ b/src/art/trajectories/_tokenize.py @@ -24,7 +24,6 @@ from ..utils.chat_template import ( chat_template_with_preserved_thinking, - configure_preserved_thinking_chat_template, default_chat_template_kwargs_for_template, normalize_tool_call_arguments_for_chat_template, ) @@ -57,9 +56,6 @@ from ._history import _model_matches from ._protocols import Exchange -if TYPE_CHECKING: - from transformers import PreTrainedTokenizerBase - _TOKEN_ID = re.compile(r"token_id:(\d+)$") _WARNED_PREFIX_RETOKENIZATION = False _TOKENIZER_LOAD_LOCK = threading.Lock() @@ -1607,22 +1603,10 @@ def _tokenizer_config(model: str, base_model: str | None) -> _TokenizerConfig: @lru_cache(maxsize=8) def _cached_tokenizer(base_model: str, revision: str | None) -> Tokenizer: - try: - from transformers import AutoTokenizer - except ImportError as exc: - raise RuntimeError( - "Tokenizer fallback requires ART's backend or tinker dependencies" - ) from exc - try: - tokenizer = AutoTokenizer.from_pretrained( - base_model, - revision=revision, - ) - if base_model.startswith("deepseek-ai/DeepSeek-V4-"): - from ..megatron.dsv4.tokenizer import get_dsv4_tokenizer + from ..tokenizer import get_tokenizer - tokenizer = get_dsv4_tokenizer(cast("PreTrainedTokenizerBase", tokenizer)) - return _as_tokenizer(configure_preserved_thinking_chat_template(tokenizer)) + try: + return _as_tokenizer(get_tokenizer(base_model, revision=revision)) except Exception as exc: raise ValueError( f"Could not load tokenizer for {base_model!r}; pass base_model explicitly" diff --git a/src/art/utils/append_only.py b/src/art/utils/append_only.py new file mode 100644 index 000000000..114528953 --- /dev/null +++ b/src/art/utils/append_only.py @@ -0,0 +1,8 @@ +"""Compatibility imports for ART's lightweight inference helpers.""" + +import sys + +from art_inference import append_only as _implementation +from art_inference.append_only import * # noqa: F403 + +sys.modules[__name__] = _implementation diff --git a/src/art/utils/chat_template.py b/src/art/utils/chat_template.py index 1be3db827..23e152361 100644 --- a/src/art/utils/chat_template.py +++ b/src/art/utils/chat_template.py @@ -1,160 +1,8 @@ -import json -import re -from typing import Any +"""Compatibility imports for ART's lightweight inference helpers.""" -THINKING_CHAT_TEMPLATE_KWARGS: dict[str, Any] = { - "enable_thinking": False, - "preserve_thinking": True, -} -TOOL_CALL_ARGUMENTS_AS_MAPPING_ATTR = "_art_tool_call_arguments_as_mapping" -_QWEN_DROP_PRIOR_THINKING = "{%- if loop.index0 > ns.last_query_index %}" -_QWEN_PRESERVE_PRIOR_THINKING = ( - "{%- if (preserve_thinking is defined and preserve_thinking is true) or " - "(loop.index0 > ns.last_query_index) %}" -) -_GEMMA_DROP_PRIOR_THINKING = ( - "thinking_text and loop.index0 > ns_turn.last_user_idx and " - "message.get('tool_calls')" -) -_GEMMA_PRESERVE_PRIOR_THINKING = ( - "thinking_text and ((preserve_thinking is defined and preserve_thinking is true) " - "or loop.index0 > ns_turn.last_user_idx) and message.get('tool_calls')" -) -_MINIMAX_DROP_PRIOR_THINKING = "reasoning_content and loop.index0 > ns.last_user_index" -_MINIMAX_PRESERVE_PRIOR_THINKING = ( - "reasoning_content and ((preserve_thinking is defined and preserve_thinking is " - "true) or loop.index0 > ns.last_user_index)" -) +import sys +from art_inference import chat_template as _implementation +from art_inference.chat_template import * # noqa: F403 -def chat_template_with_preserved_thinking(chat_template: object) -> object: - """Add opt-in prior-turn reasoning gates to supported templates.""" - if not isinstance(chat_template, str): - return chat_template - replacements = ( - ( - _QWEN_DROP_PRIOR_THINKING, - _QWEN_PRESERVE_PRIOR_THINKING, - "enable_thinking" in chat_template, - ), - ( - _GEMMA_DROP_PRIOR_THINKING, - _GEMMA_PRESERVE_PRIOR_THINKING, - True, - ), - ( - _MINIMAX_DROP_PRIOR_THINKING, - _MINIMAX_PRESERVE_PRIOR_THINKING, - True, - ), - ) - for old, new, supported in replacements: - if supported and chat_template.count(old) == 1: - chat_template = chat_template.replace(old, new) - return chat_template - - -def configure_preserved_thinking_chat_template(tokenizer: object) -> object: - chat_template = getattr(tokenizer, "chat_template", None) - configured = chat_template_with_preserved_thinking(chat_template) - if configured != chat_template: - setattr(tokenizer, "chat_template", configured) - return tokenizer - - -def default_chat_template_kwargs_for_template( - chat_template: object, -) -> dict[str, Any]: - kwargs: dict[str, Any] = {} - if not isinstance(chat_template, str): - return kwargs - if "enable_thinking" in chat_template: - kwargs["enable_thinking"] = False - if "preserve_thinking" in chat_template: - kwargs["preserve_thinking"] = True - if "clear_thinking" in chat_template: - kwargs["clear_thinking"] = False - if "deepseek_v4_python_encoder" in chat_template: - kwargs["drop_thinking"] = False - return kwargs - - -def default_chat_template_kwargs_for_tokenizer(tokenizer: object) -> dict[str, Any]: - return default_chat_template_kwargs_for_template( - getattr(tokenizer, "chat_template", None) - ) - - -def merge_chat_template_kwargs( - defaults: dict[str, Any] | None, - overrides: dict[str, Any] | None, -) -> dict[str, Any]: - return {**(defaults or {}), **(overrides or {})} - - -def _template_requires_structured_tool_arguments(chat_template: object) -> bool: - if not isinstance(chat_template, str): - return False - arguments_access = ( - r"(?:(? list[dict[str, Any]]: - """Give chat templates the structured tool arguments they require. - - Templates that interpolate the raw JSON string must keep string arguments, - so only templates that iterate structured arguments trigger normalization. - """ - if not require_mapping and not _template_requires_structured_tool_arguments( - chat_template - ): - return messages - normalized: list[dict[str, Any]] = [] - for message in messages: - calls = message.get("tool_calls") - if not isinstance(calls, list): - normalized.append(message) - continue - normalized_calls = [] - for call in calls: - function = call.get("function") if isinstance(call, dict) else None - arguments = ( - function.get("arguments") if isinstance(function, dict) else None - ) - if isinstance(arguments, str): - assert isinstance(function, dict) - try: - arguments = json.loads(arguments) if arguments.strip() else {} - except json.JSONDecodeError as error: - raise ValueError( - "tool-call arguments are not valid JSON" - ) from error - if not isinstance(arguments, dict): - raise ValueError("tool-call arguments must decode to a JSON object") - call = {**call, "function": {**function, "arguments": arguments}} - normalized_calls.append(call) - normalized.append({**message, "tool_calls": normalized_calls}) - return normalized +sys.modules[__name__] = _implementation diff --git a/src/art_inference/__init__.py b/src/art_inference/__init__.py new file mode 100644 index 000000000..94378d1a9 --- /dev/null +++ b/src/art_inference/__init__.py @@ -0,0 +1 @@ +"""Dependency-free inference helpers shipped with ART and its serving runtime.""" diff --git a/src/art_inference/append_only.py b/src/art_inference/append_only.py new file mode 100644 index 000000000..989c1c61a --- /dev/null +++ b/src/art_inference/append_only.py @@ -0,0 +1,464 @@ +"""Relate an engine's rendered assistant message to its original sampled IDs. + +Only inference may establish these mappings, using its own parsed response and +renderer. Training tokenization must continue to use the actual served tokens. +""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Mapping, Sequence +from functools import wraps +import json +import sys +from typing import Any + +from .token_prefix import PrefixEdit, prefix_edits + +PrefixObservation = tuple[list[int], list[int], tuple[PrefixEdit, ...]] + + +def _rendering_edits(tokenizer, rendered, raw): + # Stable protocol/multimodal markers separate independent text edits. Keep + # them outside replacements so native image offsets remain translatable. + special = set(getattr(tokenizer, "all_special_ids", ())) + source = [(i, token) for i, token in enumerate(rendered) if token in special] + target = [(i, token) for i, token in enumerate(raw) if token in special] + if not source or [t for _, t in source] != [t for _, t in target]: + return prefix_edits(rendered, raw) + edits = [] + start = raw_start = 0 + for (stop, _), (raw_stop, _) in zip( + [*source, (len(rendered), None)], [*target, (len(raw), None)], strict=True + ): + edits.extend( + PrefixEdit(start + edit.start, start + edit.stop, edit.replacement) + for edit in prefix_edits(rendered[start:stop], raw[raw_start:raw_stop]) + ) + start, raw_start = stop + 1, raw_stop + 1 + return tuple(edits) + + +def patch_deepseek_renderer( + module: Any, preserves: Callable[[], bool], *, prefix_only: bool = False +) -> None: + """Render historical DeepSeek turns in their original reasoning mode.""" + render = module.render_message + + @wraps(render) + def render_message(index, messages, thinking_mode, *args, **kwargs): + if preserves(): + # A generation setting controls the next answer. Parsed historical + # answers already specify whether that turn included reasoning. + for message in messages[index:]: + if message.get("role") == "assistant": + thinking_mode = ( + "thinking" + if any( + message.get(name) is not None + for name in ("reasoning", "reasoning_content", "thinking") + ) + else "chat" + ) + break + if message is not messages[index] and message.get("role") in ( + "user", + "developer", + ): + break + # V3.2 also gates reasoning and the preceding on the last + # user in the full chat, even with drop_thinking=False. + if prefix_only: + messages = messages[: index + 1] + message = messages[index] + if ( + message.get("role") == "assistant" + and thinking_mode == "thinking" + and not any( + message.get(name) + for name in ( + "reasoning", + "reasoning_content", + "thinking", + "tool_calls", + ) + ) + ): + # V3.2 rejects an empty reasoning field, although sampling + # immediately is valid. Preserve that exact block. + return module.thinking_end_token + render( + index, messages, "chat", *args, **kwargs + ) + elif not prefix_only: + # V4's encoder overrides drop_thinking when tools are present. An + # explicit request to remove history still controls each rendering. + kwargs["drop_thinking"] = True + return render(index, messages, thinking_mode, *args, **kwargs) + + module.render_message = render_message + + +def patch_harmony(module: Any, preserves: Callable[[], bool], namespace: str) -> None: + """Keep Harmony analysis and final-turn stops, including imported aliases.""" + original = module.render_for_completion + + @wraps(original) + def render(messages): + if not preserves(): + return original(messages) + # Harmony is provided by the inference engine's environment. + from openai_harmony import ( # ty: ignore[unresolved-import] + Conversation, + RenderConversationConfig, + Role, + ) + + encoding = module.get_encoding() + config = RenderConversationConfig(auto_drop_analysis=False) + # Harmony's whole-conversation renderer changes prior <|return|> to + # <|end|>. Rendering each completed message keeps its original framing. + tokens = [ + token + for message in messages + for token in encoding.render_conversation_for_training( + Conversation.from_messages([message]), config=config + ) + ] + return tokens + encoding.render_conversation_for_completion( + Conversation.from_messages([]), Role.ASSISTANT, config=config + ) + + replacements = {original: render} + drop = getattr(module, "auto_drop_analysis_messages", None) + if drop is not None: + replacements[drop] = lambda messages: ( + messages if preserves() else drop(messages) + ) + for name, loaded in tuple(sys.modules.items()): + if loaded is not None and name.startswith(namespace): + for attribute, value in tuple(vars(loaded).items()): + for old, new in replacements.items(): + if value is old: + setattr(loaded, attribute, new) + module.render_for_completion = render + if drop is not None: + module.auto_drop_analysis_messages = replacements[drop] + + +def shifted_span(start: int, length: int, edits: Sequence[PrefixEdit]) -> int: + """Move a multimodal span after text edits without changing its contents.""" + shift = 0 + for edit in edits: + if edit.stop <= start: + shift += len(edit.replacement) - (edit.stop - edit.start) + elif edit.start < start + length: + raise ValueError("History token edit intersects a multimodal placeholder") + return start + shift + + +def aligned_values( + values: Any, edits: Sequence[PrefixEdit], *, positions=False, fill=None +): + """Translate per-token metadata through edits confined to text tokens.""" + if positions: + import torch + + parts, start, shift = [], 0, 0 + for edit in edits: + parts.append(values[..., start : edit.start] + shift) + origin = ( + values[..., edit.start : edit.start + 1] + if edit.start < values.shape[-1] + else values[..., -1:] + 1 + ) + parts.append( + origin + shift + values.new_tensor(list(range(len(edit.replacement)))) + ) + shift += len(edit.replacement) - (edit.stop - edit.start) + start = edit.stop + return torch.cat([*parts, values[..., start:] + shift], dim=-1) + result = list(values) + for edit in reversed(edits): + replaced = result[edit.start : edit.stop] + value = ( + fill + if fill is not None + else (result[min(edit.start, len(result) - 1)] if result else 0) + ) + if any(item != value for item in replaced): + raise ValueError("History token edit crosses a per-token metadata boundary") + result[edit.start : edit.stop] = [value] * len(edit.replacement) + return values.new_tensor(result) if hasattr(values, "new_tensor") else result + + +def output_prefix_observations( + tokenizer: Any, + rendered_prompt: Sequence[int], + raw_prompt: Sequence[int], + raw_output: Sequence[int], + *, + prompt_edits: Sequence[PrefixEdit] | None = None, +) -> list[PrefixObservation]: + """Preserve native IDs for the response and its last reasoning boundary. + + In particular, reasoning can survive a later action edit, for streaming and + non-streaming generations in any protocol. A decode/encode pass must retain + the same special-token sequence before we use an intermediate boundary. + Keep at most two entries so tool-heavy outputs do not multiply prompt storage. + """ + output = list(raw_output) + if not output: + return [] + rendered = list( + tokenizer.encode(_decode(tokenizer, output), add_special_tokens=False) + ) + if not rendered: + return [] + special_ids = set(getattr(tokenizer, "all_special_ids", ())) + raw_boundaries = [ + (i + 1, token) for i, token in enumerate(output) if token in special_ids + ] + rendered_boundaries = [ + (i + 1, token) for i, token in enumerate(rendered) if token in special_ids + ] + boundaries = [] + if [token for _, token in raw_boundaries] == [ + token for _, token in rendered_boundaries + ]: + boundaries = [ + (r[0], n[0]) + for r, n in zip(rendered_boundaries, raw_boundaries, strict=True) + if _decode(tokenizer, [r[1]]) in {"", "", "<|end|>"} + ][-1:] + if not boundaries or boundaries[-1] != (len(rendered), len(output)): + boundaries.append((len(rendered), len(output))) + entries = [] + prompt_edits = ( + tuple(prompt_edits) + if prompt_edits is not None + else _rendering_edits(tokenizer, rendered_prompt, raw_prompt) + ) + offset = len(rendered_prompt) + for rendered_end, raw_end in boundaries: + edits = tuple( + PrefixEdit(offset + edit.start, offset + edit.stop, edit.replacement) + for edit in _rendering_edits( + tokenizer, rendered[:rendered_end], output[:raw_end] + ) + ) + entries.append( + ( + [*rendered_prompt, *rendered[:rendered_end]], + [*raw_prompt, *output[:raw_end]], + (*prompt_edits, *edits), + ) + ) + return entries + + +def merge_chat_delta(message: dict[str, Any], delta: dict[str, Any]) -> None: + for name, value in delta.items(): + if name == "tool_calls" and isinstance(value, list): + calls = message.setdefault("tool_calls", []) + for part in value: + index = part["index"] + while len(calls) <= index: + calls.append({}) + merge_chat_delta( + calls[index], {k: v for k, v in part.items() if k != "index"} + ) + elif isinstance(value, dict): + merge_chat_delta(message.setdefault(name, {}), value) + elif isinstance(value, str): + message[name] = ( + value + if name in {"role", "type", "id"} + else message.get(name, "") + value + ) + + +def preserves_history(kwargs: Mapping[str, Any] | None) -> bool: + options = kwargs or {} + return not ( + options.get("preserve_thinking") is False + or options.get("clear_thinking") is True + or options.get("drop_thinking") is True + ) + + +def has_renderable_tool_arguments(message: Mapping[str, Any]) -> bool: + functions = [call.get("function", {}) for call in message.get("tool_calls") or []] + if message.get("function_call"): + functions.append(message["function_call"]) + try: + return all( + not isinstance(f.get("arguments"), str) + or isinstance(json.loads(f["arguments"] or "{}"), dict) + for f in functions + ) + except json.JSONDecodeError: + return False + + +def _decode(tokenizer: Any, tokens: Sequence[int]) -> str: + return tokenizer.decode(list(tokens), skip_special_tokens=False) + + +def _whitespace_prefix(tokenizer: Any, rendered: list[int], raw: list[int]) -> int: + """Find an exact raw-token boundary for a whitespace-normalized prefix.""" + target = _decode(tokenizer, rendered) + source = _decode(tokenizer, raw) + i = j = 0 + while i < len(target): + if target[i].isspace(): + i += 1 + elif j < len(source) and source[j].isspace(): + j += 1 + elif j < len(source) and target[i] == source[j]: + i += 1 + j += 1 + else: + return 0 + if target and target[-1].isspace(): + while j < len(source) and source[j].isspace(): + j += 1 + # Decoding keeps the native IDs, including non-canonical tokenizations. A + # boundary inside a token cannot be preserved independently of the action. + low, high = 0, len(raw) + while low < high: + middle = (low + high) // 2 + if len(_decode(tokenizer, raw[:middle])) < j: + low = middle + 1 + else: + high = middle + return low if _decode(tokenizer, raw[:low]) == source[:j] else 0 + + +def chat_prefix_observations( + tokenizer: Any, + rendered_prompt: Sequence[int], + completed_prompt: Sequence[int], + raw_prompt: Sequence[int], + raw_output: Sequence[int], + *, + reasoning_prompt: Sequence[int] | None = None, + complete: bool = True, +) -> list[PrefixObservation]: + """Record a native response and, separately, its unchanged reasoning. + + ``completed_prompt`` must come from rendering the original request plus the + engine's parsed response. ``reasoning_prompt`` renders the same request with + only that response's reasoning. Keeping the latter mapping lets callers edit + the action without losing the exact reasoning prefix. Explicit requests to + drop thinking must bypass observation and lookup. + + Truncated responses may contribute a reasoning prefix, but cannot replace a + completed turn's framing. No sampled IDs or log probabilities are invented. + """ + prompt, completed = list(rendered_prompt), list(completed_prompt) + output = list(raw_output) + if not output or completed[: len(prompt)] != prompt: + return [] + entries: list[PrefixObservation] = [] + if reasoning_prompt is not None: + boundary = next( + (i for i, (a, b) in enumerate(zip(completed, reasoning_prompt)) if a != b), + min(len(completed), len(reasoning_prompt)), + ) + if boundary > len(prompt): + sampled_boundary = _whitespace_prefix( + tokenizer, completed[len(prompt) : boundary], output + ) + if sampled_boundary: + rendered = completed[:boundary] + raw = [*raw_prompt, *output[:sampled_boundary]] + entries.append( + (rendered, raw, _rendering_edits(tokenizer, rendered, raw)) + ) + if complete and len(completed) > len(prompt): + # A template's separator after the sampled stop belongs to the next + # turn. Leave it in the rendered suffix instead of swallowing it. + while ( + len(completed) > len(prompt) + and _decode(tokenizer, completed[-1:]).isspace() + ): + completed.pop() + # User stop strings (and APIs omitting the sampled stop ID) do not + # prove the template's terminal framing. Never delete that framing. + if completed[-1] != output[-1]: + return entries + raw = [*raw_prompt, *output] + entries.append((completed, raw, _rendering_edits(tokenizer, completed, raw))) + return entries + + +async def chat_response_prefixes( + tokenizer: Any, + request: Any, + raw_prompt: Sequence[int], + choices: Sequence[tuple[Mapping[str, Any], Sequence[int], bool]], + render: Callable[[Any], Awaitable[list[int] | None]], +) -> list[PrefixObservation]: + """Observe parsed native chat choices using that engine's request renderer.""" + payload = request.model_dump(mode="python") + if not preserves_history(payload.get("chat_template_kwargs")): + return [] + rendered = await render(request) + messages = payload.get("messages") + if rendered is None or not isinstance(messages, list): + return [] + fields = type(request).model_fields + payload.update( + (name, value) + for name, value in ( + ("add_generation_prompt", False), + ("continue_final_message", False), + ("input_ids", None), + ("n", 1), + ("stream", False), + ("stream_options", None), + # These turns have already completed: reserve no further output and + # retain the entire history when observing its rendered tokens. + ("max_tokens", None), + ("max_completion_tokens", None), + ("truncate_prompt_tokens", None), + ) + if name in fields + ) + + async def complete(message: Mapping[str, Any]) -> list[int] | None: + return await render( + type(request).model_validate({**payload, "messages": [*messages, message]}) + ) + + entries: list[PrefixObservation] = [] + for message, output, finished in choices: + # An invalid sampled call must reach the caller unchanged. It cannot be + # rendered as a completed tool turn; native special-token observations + # still preserve its reasoning prefix. + if not has_renderable_tool_arguments(message): + continue + completed = await complete(message) + if completed is None: + continue + reasoning = { + key: message[key] + for key in ("reasoning", "reasoning_content", "thinking") + if message.get(key) is not None + } + reasoning_prompt = ( + await complete({"role": "assistant", "content": "", **reasoning}) + if reasoning + else None + ) + entries.extend( + chat_prefix_observations( + tokenizer, + rendered, + completed, + raw_prompt, + output, + reasoning_prompt=reasoning_prompt, + complete=finished, + ) + ) + return entries diff --git a/src/art_inference/chat_template.py b/src/art_inference/chat_template.py new file mode 100644 index 000000000..20dd21c7e --- /dev/null +++ b/src/art_inference/chat_template.py @@ -0,0 +1,258 @@ +import json +import re +from typing import Any + +THINKING_CHAT_TEMPLATE_KWARGS: dict[str, Any] = { + "enable_thinking": False, + "preserve_thinking": True, +} +TOOL_CALL_ARGUMENTS_AS_MAPPING_ATTR = "_art_tool_call_arguments_as_mapping" +_QWEN_DROP_PRIOR_THINKING = "{%- if loop.index0 > ns.last_query_index %}" +_QWEN_PRESERVE_PRIOR_THINKING = ( + "{%- if (preserve_thinking is defined and preserve_thinking is true) or " + "(loop.index0 > ns.last_query_index) %}" +) +_GEMMA_DROP_PRIOR_THINKING = ( + "thinking_text and loop.index0 > ns_turn.last_user_idx and " + "message.get('tool_calls')" +) +_GEMMA_PRESERVE_PRIOR_THINKING = ( + "thinking_text and ((preserve_thinking is defined and preserve_thinking is true) " + "or (loop.index0 > ns_turn.last_user_idx and message.get('tool_calls')))" +) +_MINIMAX_DROP_PRIOR_THINKING = "reasoning_content and loop.index0 > ns.last_user_index" +_MINIMAX_PRESERVE_PRIOR_THINKING = ( + "reasoning_content and ((preserve_thinking is defined and preserve_thinking is " + "true) or loop.index0 > ns.last_user_index)" +) + + +def chat_template_with_preserved_thinking(chat_template: object) -> object: + """Preserve prior reasoning by default, while respecting explicit opt-outs.""" + if isinstance(chat_template, dict): + return { + name: chat_template_with_preserved_thinking(template) + for name, template in chat_template.items() + } + if not isinstance(chat_template, str): + return chat_template + replacements = ( + ( + _QWEN_DROP_PRIOR_THINKING, + _QWEN_PRESERVE_PRIOR_THINKING, + "enable_thinking" in chat_template, + ), + ( + _GEMMA_DROP_PRIOR_THINKING, + _GEMMA_PRESERVE_PRIOR_THINKING, + True, + ), + ( + _MINIMAX_DROP_PRIOR_THINKING, + _MINIMAX_PRESERVE_PRIOR_THINKING, + True, + ), + ( + "(loop.index0 > ns_turn.last_user_idx) or (preserve_thinking and message.get('tool_calls'))", + "preserve_thinking or (loop.index0 > ns_turn.last_user_idx)", + True, + ), + *( + ( + f"message.{field} and not future_final_message.found", + f"message.{field} and ((preserve_thinking is defined and preserve_thinking is true) or not future_final_message.found)", + True, + ) + for field in ("content", "thinking") + ), + ( + "{#- CoT is dropped during all previous turns, so we never render it for inference #}", + "{%- if preserve_thinking is defined and preserve_thinking is true and message.thinking is defined %}<|start|>assistant<|channel|>analysis<|message|>{{ message.thinking }}<|end|>{%- endif %}", + True, + ), + ( + '"<|start|>assistant<|channel|>final<|message|>" + message.content + "<|end|>"', + '"<|start|>assistant<|channel|>final<|message|>" + message.content + ("<|return|>" if preserve_thinking else "<|end|>")', + True, + ), + ) + for old, new, supported in replacements: + if supported and chat_template.count(old) == 1: + chat_template = chat_template.replace(old, new) + # Kimi 2.5 splits historical turns into a branch that blanks reasoning. + # Preserve them using its ordinary assistant-turn branch instead. + if ( + "set hist_msgs = messages[:ns.last_non_tool_call_assistant_msg+1]" + in chat_template + ): + chat_template = ( + chat_template.replace( + "set hist_msgs = messages[:ns.last_non_tool_call_assistant_msg+1]", + "set hist_msgs = [] if preserve_thinking else messages[:ns.last_non_tool_call_assistant_msg+1]", + ) + .replace( + "set suffix_msgs = messages[ns.last_non_tool_call_assistant_msg+1:]", + "set suffix_msgs = messages if preserve_thinking else messages[ns.last_non_tool_call_assistant_msg+1:]", + ) + .replace( + "{%- if thinking is defined and thinking is false -%}", + "{%- if thinking is defined and thinking is false and not preserve_thinking -%}", + 1, + ) + ) + # Qwen's native reasoning parser returns the sampled whitespace. The stock + # template trims it, then invents a newline before . Preserve the + # parsed field verbatim; keep the original non-thinking/legacy fallback. + if ( + "" in chat_template + and "reasoning_content|trim" in chat_template + and "if not preserve_thinking or message.reasoning_content" not in chat_template + ): + chat_template = chat_template.replace( + "{%- set reasoning_content = reasoning_content|trim %}", + "{%- if not preserve_thinking or message.reasoning_content is not string %}" + "{%- set reasoning_content = reasoning_content|trim %}{%- endif %}", + ) + chat_template = chat_template.replace( + "reasoning_content + '\\n\\n\\n'", + "reasoning_content + ('\\n\\n' if preserve_thinking and message.reasoning_content is string and reasoning_content else '\\n\\n\\n')", + ) + chat_template = chat_template.replace( + "set content = render_content(message.content, true)|trim", + "set content = (render_content(message.content, true) if preserve_thinking and message.role == 'assistant' else render_content(message.content, true)|trim)", + ) + if "clear_thinking" in chat_template: + chat_template = chat_template.replace( + "{{ content.strip() }}", + "{{ content if not clear_thinking else content.strip() }}", + ).replace( + "{%- if content.strip() -%}", + "{%- if (content if not clear_thinking else content.strip()) -%}", + ) + if "ns_turn.last_user_idx" in chat_template: + for value in ("message['content']", "item['text']"): + chat_template = chat_template.replace( + f"{{{{- {value} | trim -}}}}", + f"{{{{- {value} if preserve_thinking and message.role == 'assistant' else {value} | trim -}}}}", + ) + # Embed preservation defaults so direct tokenizer use and inference engines + # agree with ART. Passing a value explicitly still takes precedence. + for name, value in (("preserve_thinking", "true"), ("clear_thinking", "false")): + default = f"{{%- set {name} = {name} | default({value}) -%}}" + if name in chat_template and default not in chat_template: + chat_template = default + chat_template + return chat_template + + +def configure_preserved_thinking_chat_template(tokenizer: object) -> object: + chat_template = getattr(tokenizer, "chat_template", None) + configured = chat_template_with_preserved_thinking(chat_template) + if configured != chat_template: + setattr(tokenizer, "chat_template", configured) + return tokenizer + + +def default_chat_template_kwargs_for_template( + chat_template: object, +) -> dict[str, Any]: + kwargs = default_preservation_kwargs_for_template(chat_template) + if not isinstance(chat_template, str): + return kwargs + if "enable_thinking" in chat_template: + kwargs["enable_thinking"] = False + return kwargs + + +def default_preservation_kwargs_for_template(chat_template: object) -> dict[str, Any]: + """History defaults independent of whether the next turn should think.""" + kwargs: dict[str, Any] = {} + if not isinstance(chat_template, str): + return kwargs + if "preserve_thinking" in chat_template: + kwargs["preserve_thinking"] = True + if "clear_thinking" in chat_template: + kwargs["clear_thinking"] = False + if "deepseek_v4_python_encoder" in chat_template: + kwargs["drop_thinking"] = False + return kwargs + + +def default_chat_template_kwargs_for_tokenizer(tokenizer: object) -> dict[str, Any]: + return default_chat_template_kwargs_for_template( + getattr(tokenizer, "chat_template", None) + ) + + +def merge_chat_template_kwargs( + defaults: dict[str, Any] | None, + overrides: dict[str, Any] | None, +) -> dict[str, Any]: + return {**(defaults or {}), **(overrides or {})} + + +def _template_requires_structured_tool_arguments(chat_template: object) -> bool: + if not isinstance(chat_template, str): + return False + arguments_access = ( + r"(?:(? list[dict[str, Any]]: + """Give chat templates the structured tool arguments they require. + + Templates that interpolate the raw JSON string must keep string arguments, + so only templates that iterate structured arguments trigger normalization. + """ + if not require_mapping and not _template_requires_structured_tool_arguments( + chat_template + ): + return messages + normalized: list[dict[str, Any]] = [] + for message in messages: + calls = message.get("tool_calls") + if not isinstance(calls, list): + normalized.append(message) + continue + normalized_calls = [] + for call in calls: + function = call.get("function") if isinstance(call, dict) else None + arguments = ( + function.get("arguments") if isinstance(function, dict) else None + ) + if isinstance(arguments, str): + assert isinstance(function, dict) + try: + arguments = json.loads(arguments) if arguments.strip() else {} + except json.JSONDecodeError as error: + raise ValueError( + "tool-call arguments are not valid JSON" + ) from error + if not isinstance(arguments, dict): + raise ValueError("tool-call arguments must decode to a JSON object") + call = {**call, "function": {**function, "arguments": arguments}} + normalized_calls.append(call) + normalized.append({**message, "tool_calls": normalized_calls}) + return normalized diff --git a/src/art_inference/sglang.py b/src/art_inference/sglang.py new file mode 100644 index 000000000..680a9bd39 --- /dev/null +++ b/src/art_inference/sglang.py @@ -0,0 +1,382 @@ +"""History preservation for SGLang's native chat, Responses, and Messages APIs.""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable +from contextvars import ContextVar +import copy +from functools import wraps +import importlib +import json +from typing import Any + +from .append_only import ( + aligned_values, + chat_response_prefixes, + merge_chat_delta, + patch_deepseek_renderer, + patch_harmony, + preserves_history, + shifted_span, +) +from .chat_template import ( + chat_template_with_preserved_thinking, + configure_preserved_thinking_chat_template, +) +from .token_prefix import apply_prefix_edits + +_DROP_THINKING: ContextVar[bool] = ContextVar("art_sglang_drop_thinking", default=False) + + +def replace_prompt_tokens(tokenized, edits): + """Translate expanded multimodal positions with the preserved text prefix.""" + if not edits: + return + mm = getattr(tokenized, "mm_inputs", None) + if mm is not None: + mm = copy.copy(mm) + mm.mm_items = [copy.copy(item) for item in mm.mm_items] + for item in mm.mm_items: + if item.offsets is not None: + item.offsets = [ + ( + shifted_span(start, end - start + 1, edits), + shifted_span(start, end - start + 1, edits) + end - start, + ) + for start, end in item.offsets + ] + for name in ("input_ids", "padded_input_ids"): + if getattr(mm, name, None) is not None: + setattr(mm, name, apply_prefix_edits(list(getattr(mm, name)), edits)) + for name in ("mrope_positions", "token_type_ids", "visible_frame_counts"): + if getattr(mm, name, None) is not None: + setattr( + mm, + name, + aligned_values( + getattr(mm, name), + edits, + positions=name == "mrope_positions", + fill=0 if name == "token_type_ids" else None, + ), + ) + tokenized.mm_inputs = mm + if getattr(tokenized, "token_type_ids", None) is not None: + tokenized.token_type_ids = aligned_values( + tokenized.token_type_ids, edits, fill=0 + ) + + +class _TokenizerView: + def __init__(self, tokenizer: Any, generation: bool): + self._tokenizer = tokenizer + self._generation = generation + + def __getattr__(self, name): + return getattr(self._tokenizer, name) + + def apply_chat_template(self, *args, **kwargs): + kwargs["add_generation_prompt"] = self._generation + return self._tokenizer.apply_chat_template(*args, **kwargs) + + +def patch_history( + observe: Callable[[Any, list[Any]], Awaitable[None]], + importer=importlib.import_module, +) -> None: + chat = importer("sglang.srt.entrypoints.openai.serving_chat").OpenAIServingChat + responses = importer( + "sglang.srt.entrypoints.openai.serving_responses" + ).OpenAIServingResponses + if getattr(chat, "_art_append_only", False): + return + patch_harmony( + importer("sglang.srt.entrypoints.harmony_utils"), + lambda: not _DROP_THINKING.get(), + "sglang.", + ) + process = chat._process_messages + + def patch_encoder(encoding): + encode = encoding.encode_messages + + @wraps(encode) + def encode_messages(*args, **kwargs): + kwargs.setdefault("drop_thinking", _DROP_THINKING.get()) + return encode(*args, **kwargs) + + encoding.encode_messages = encode_messages + + for name in ("dsv4", "dsv32"): + encoding = importer(f"sglang.srt.entrypoints.openai.encoding_{name}") + patch_encoder(encoding) + patch_deepseek_renderer( + encoding, lambda: not _DROP_THINKING.get(), prefix_only=name == "dsv32" + ) + + @wraps(process) + def process_messages(self, request, is_multimodal): + configure_preserved_thinking_chat_template(self.tokenizer_manager.tokenizer) + if getattr(request, "chat_template", None) is not None: + request.chat_template = chat_template_with_preserved_thinking( + request.chat_template + ) + # Templates already embed their defaults. Adding them to this request + # would shadow a Responses adapter's outer rendering options. Likewise, + # retain its encoder context unless this chat request overrides it. + options = getattr(request, "chat_template_kwargs", None) or {} + token = _DROP_THINKING.set( + not preserves_history(options) + if any( + k in options + for k in ("preserve_thinking", "clear_thinking", "drop_thinking") + ) + else _DROP_THINKING.get() + ) + try: + return process(self, request, is_multimodal) + finally: + _DROP_THINKING.reset(token) + + chat._process_messages = process_messages + + async def record(self, request, raw_request, prompt, choices): + if prompt is None or not preserves_history( + getattr(request, "chat_template_kwargs", None) + ): + return + + async def render(value): + view = copy.copy(self) + generation = len(value.messages) == len(request.messages) + if not generation: + # Native request rendering converts a trailing assistant into + # a user turn. This view renders an already completed response. + view._handle_last_assistant_message = lambda messages, _: ( + messages, + None, + ) + view.tokenizer_manager = copy.copy(self.tokenizer_manager) + view.tokenizer_manager.tokenizer = _TokenizerView( + self.tokenizer_manager.tokenizer, + generation, + ) + result = view._process_messages( + value, + is_multimodal=self.tokenizer_manager.model_config.is_multimodal, + ).prompt_ids + return result if isinstance(result, list) and result else None + + entries = await chat_response_prefixes( + self.tokenizer_manager.tokenizer, request, prompt, choices, render + ) + if entries: + await observe(raw_request, entries) + + full = chat._handle_non_streaming_request + + @wraps(full) + async def complete(self, adapted_request, request, raw_request): + result = await full(self, adapted_request, request, raw_request) + payload = ( + result.model_dump(exclude_none=True) + if hasattr(result, "model_dump") + else json.loads(result.body) + ) + prompt = payload.get("prompt_token_ids") + choices = [] + for choice in payload.get("choices", []): + prompt = prompt or choice.get("prompt_token_ids") + ids = choice.get("token_ids") or (choice.get("logprobs") or {}).get( + "token_ids" + ) + if ids is not None: + choices.append( + ( + choice["message"], + ids, + choice.get("finish_reason") not in ("length", "abort"), + ) + ) + if choices: + await record(self, request, raw_request, prompt, choices) + return result + + chat._handle_non_streaming_request = complete + generate = chat._generate_chat_stream + + @wraps(generate) + async def stream(self, adapted_request, request, raw_request): + messages, tokens, finished = {}, {}, {} + prompt = None + async for event in generate(self, adapted_request, request, raw_request): + if isinstance(event, str) and event.startswith("data: "): + data = event[6:].strip() + if data == "[DONE]": + await record( + self, + request, + raw_request, + prompt, + [ + (message, tokens[index], finished.get(index, False)) + for index, message in messages.items() + if index in tokens + ], + ) + else: + payload = json.loads(data) + prompt = prompt or payload.get("prompt_token_ids") + for choice in payload.get("choices", []): + index = choice["index"] + merge_chat_delta( + messages.setdefault(index, {"role": "assistant"}), + choice.get("delta") or {}, + ) + if "token_ids" in choice: + tokens.setdefault(index, []).extend(choice["token_ids"]) + if choice.get("finish_reason"): + finished[index] = choice["finish_reason"] not in ( + "length", + "abort", + ) + yield event + + chat._generate_chat_stream = stream + + normalize = responses._normalize_response_message_for_chat + + @classmethod + def normalize_message(cls, message): + if hasattr(message, "model_dump"): + message = message.model_dump(exclude_none=True) + if isinstance(message, dict) and message.get("type") == "reasoning": + # A summary may differ from the sampled reasoning. Joining content + # parts with newlines would also alter the original response. + parts = message.get("content") or message.get("summary") or [] + text = "".join(part.get("text", "") for part in parts) + return {"role": "assistant", "reasoning_content": text} if text else None + return normalize(message) + + responses._normalize_response_message_for_chat = normalize_message + construct = responses._construct_input_messages + + @wraps(construct) + def input_messages(self, request, prev_response=None): + if prev_response is not None: + inputs = ( + [{"role": "user", "content": request.input}] + if isinstance(request.input, str) + else request.input + ) + request = request.model_copy( + update={"input": [*prev_response.output, *inputs]} + ) + prev_response = prev_response.model_copy(update={"output": []}) + return construct(self, request, prev_response) + + responses._construct_input_messages = input_messages + harmony = responses._construct_input_messages_with_harmony + + @wraps(harmony) + def harmony_messages(self, request, prev_response): + if prev_response is None: + return harmony(self, request, prev_response) + # SGLang deletes reasoning from the stored list in place. Give it a + # private list, and retain the original prefix unless rewriting was + # explicitly requested for this call. + previous = list(self.msg_store[prev_response.id]) + view = copy.copy(self) + view.msg_store = {**self.msg_store, prev_response.id: previous.copy()} + result = harmony(view, request, prev_response) + retained = view.msg_store[prev_response.id] + if preserves_history(getattr(request, "chat_template_kwargs", None)): + if result[: len(retained)] != retained: + raise RuntimeError("SGLang's stored-history rendering contract changed") + return [*previous, *result[len(retained) :]] + return result + + responses._construct_input_messages_with_harmony = harmony_messages + + create = responses.create_responses + + @wraps(create) + async def create_responses(self, request, raw_request=None): + async def record_response(payload): + generations = payload.get("token_generations") or [] + if getattr(self, "use_harmony", False) or len(generations) != 1: + return + generation = generations[0] + messages = self._merge_consecutive_assistant_messages( + [ + message + for item in payload.get("output", []) + if (message := self._normalize_response_message_for_chat(item)) + is not None + ] + ) + if len(messages) != 1 or messages[0].get("role") != "assistant": + return + previous = ( + self.response_store.get(request.previous_response_id) + if request.previous_response_id + else None + ) + protocol = importer("sglang.srt.entrypoints.openai.protocol") + view = protocol.ChatCompletionRequest( + model=request.model, + messages=self._construct_input_messages(request, previous), + tools=self._response_tools_to_chat_tools(request) or None, + chat_template=getattr(request, "chat_template", None), + chat_template_kwargs=getattr(request, "chat_template_kwargs", None), + ) + await record( + self, + view, + raw_request, + generation["prompt_token_ids"], + [ + ( + messages[0], + [token["token_id"] for token in generation["output_tokens"]], + payload.get("status") == "completed", + ) + ], + ) + + token = _DROP_THINKING.set( + not preserves_history(getattr(request, "chat_template_kwargs", None)) + ) + try: + result = await create(self, request, raw_request) + finally: + _DROP_THINKING.reset(token) + if not isinstance(result, AsyncIterator): + payload = ( + result.model_dump(exclude_none=True) + if hasattr(result, "model_dump") + else json.loads(result.body) + ) + await record_response(payload) + return result + + async def stream_response(): + token = _DROP_THINKING.set( + not preserves_history(getattr(request, "chat_template_kwargs", None)) + ) + try: + async for event in result: + if isinstance(event, str) and event.startswith( + ("event: response.completed\n", "event: response.incomplete\n") + ): + await record_response( + json.loads(event.split("data: ", 1)[1])["response"] + ) + yield event + finally: + _DROP_THINKING.reset(token) + + return stream_response() + + responses.create_responses = create_responses + chat._art_append_only = True diff --git a/src/art_inference/token_prefix.py b/src/art_inference/token_prefix.py new file mode 100644 index 000000000..9fdd8aff4 --- /dev/null +++ b/src/art_inference/token_prefix.py @@ -0,0 +1,1054 @@ +"""Bounded mappings from rendered history to previously served token IDs.""" + +from __future__ import annotations + +from array import array +import asyncio +import base64 +import binascii +from collections import OrderedDict +from collections.abc import Awaitable, Callable +from contextlib import suppress +from dataclasses import dataclass, field +import gzip +import hashlib +import json +import logging +import string +import sys +import time +from typing import Sequence +import zlib + +logger = logging.getLogger(__name__) + +COMPACT_PREFIX_VERSION = "sha256-u32be-v1" +COMPACT_PREFIX_MAX_CANDIDATES = 1024 +COMPACT_PREFIX_MAX_REPLACEMENT_IDS = 262_144 +COMPACT_PREFIX_RECORD_LIMIT = 4 * 1024 * 1024 +COMPACT_PREFIX_RESPONSE_LIMIT = 8 * 1024 * 1024 +COMPACT_PREFIX_HEADER_LIMIT = 6 * 1024 +COMPACT_PREFIX_HEADER_DECODE_LIMIT = 64 * 1024 +# Keep the original wire domain so deployed prefix stores remain compatible. +_DIGEST_DOMAIN = b"caladan-token-prefix\0" + COMPACT_PREFIX_VERSION.encode() + b"\0" +_MAX_PREFIX_TOKENS = 262_144 +_SharedPrefixEntry = tuple[str, list[int], list[int], str, tuple["PrefixEdit", ...]] + + +@dataclass(frozen=True, slots=True) +class PrefixEdit: + start: int + stop: int + replacement: tuple[int, ...] + + +@dataclass(frozen=True, slots=True) +class CompactPrefixCandidate: + rendered_length: int + rendered_digest: str + raw_digest: str + edits: tuple[PrefixEdit, ...] + + +@dataclass(frozen=True, slots=True) +class PrefixMatch: + rendered_length: int + raw_prefix: tuple[int, ...] + edits: tuple[PrefixEdit, ...] = field(default=(), compare=False) + + +def _token_bytes(token_ids: Sequence[int]) -> bytes: + # Engine and wire boundaries validate bool/type explicitly. Let the C array + # conversion enforce the numeric/range constraint without another O(n) + # Python scan on every 128K lookup. + try: + values = array("I", token_ids) + except (OverflowError, TypeError, ValueError) as error: + raise ValueError("token IDs must be unsigned 32-bit integers") from error + if values.itemsize != 4: + raise RuntimeError("this platform does not provide 32-bit unsigned integers") + if sys.byteorder == "little": + values.byteswap() + return values.tobytes() + + +def token_prefix_digests( + token_ids: Sequence[int], lengths: Sequence[int] +) -> dict[int, str]: + """Hash several token prefixes in one pass over a fixed-width encoding.""" + requested = sorted(set(lengths)) + if any( + not isinstance(length, int) + or isinstance(length, bool) + or not 0 <= length <= len(token_ids) + for length in requested + ): + raise ValueError("invalid token prefix length") + if not requested: + return {} + encoded = _token_bytes(token_ids[: requested[-1]]) + digest = hashlib.sha256(_DIGEST_DOMAIN) + previous = 0 + result: dict[int, str] = {} + for length in requested: + digest.update(encoded[previous * 4 : length * 4]) + result[length] = digest.hexdigest() + previous = length + return result + + +def token_digest(token_ids: Sequence[int]) -> str: + return token_prefix_digests(token_ids, [len(token_ids)])[len(token_ids)] + + +def apply_prefix_edits( + rendered_prefix: Sequence[int], edits: Sequence[PrefixEdit] +) -> list[int]: + result: list[int] = [] + cursor = 0 + for edit in edits: + if ( + not 0 <= edit.start <= edit.stop <= len(rendered_prefix) + or edit.start < cursor + ): + raise ValueError("token-prefix edits must be ordered and non-overlapping") + result.extend(rendered_prefix[cursor : edit.start]) + result.extend(edit.replacement) + cursor = edit.stop + result.extend(rendered_prefix[cursor:]) + return result + + +def prefix_edits( + rendered_prefix: Sequence[int], raw_prefix: Sequence[int] +) -> tuple[PrefixEdit, ...]: + """Return an exact, deterministic edit script without nonlinear alignment.""" + if rendered_prefix == raw_prefix: + return () + if len(rendered_prefix) == len(raw_prefix): + edits: list[PrefixEdit] = [] + start: int | None = None + for index, (rendered, raw) in enumerate( + zip(rendered_prefix, raw_prefix, strict=True) + ): + if rendered != raw and start is None: + start = index + elif rendered == raw and start is not None: + edits.append(PrefixEdit(start, index, tuple(raw_prefix[start:index]))) + start = None + if start is not None: + edits.append(PrefixEdit(start, len(raw_prefix), tuple(raw_prefix[start:]))) + return tuple(edits) + + start = 0 + limit = min(len(rendered_prefix), len(raw_prefix)) + while start < limit and rendered_prefix[start] == raw_prefix[start]: + start += 1 + suffix = 0 + while ( + suffix < len(rendered_prefix) - start + and suffix < len(raw_prefix) - start + and rendered_prefix[-1 - suffix] == raw_prefix[-1 - suffix] + ): + suffix += 1 + raw_stop = len(raw_prefix) - suffix + return ( + PrefixEdit( + start, + len(rendered_prefix) - suffix, + tuple(raw_prefix[start:raw_stop]), + ), + ) + + +def compact_prefix_candidate( + rendered_prefix: Sequence[int], + raw_prefix: Sequence[int], + edits: Sequence[PrefixEdit] | None = None, +) -> CompactPrefixCandidate: + exact_edits = tuple( + prefix_edits(rendered_prefix, raw_prefix) if edits is None else edits + ) + if apply_prefix_edits(rendered_prefix, exact_edits) != list(raw_prefix): + raise ValueError("token-prefix edits do not reconstruct the raw prefix") + return CompactPrefixCandidate( + rendered_length=len(rendered_prefix), + rendered_digest=token_digest(rendered_prefix), + raw_digest=token_digest(raw_prefix), + edits=exact_edits, + ) + + +def resolve_compact_prefix( + rendered: Sequence[int], candidates: Sequence[CompactPrefixCandidate] +) -> PrefixMatch | None: + eligible = sorted( + ( + candidate + for candidate in candidates + if 0 < candidate.rendered_length <= len(rendered) + ), + key=lambda candidate: candidate.rendered_length, + reverse=True, + ) + if not eligible: + return None + digests = token_prefix_digests( + rendered, [candidate.rendered_length for candidate in eligible] + ) + for candidate in eligible: + if digests[candidate.rendered_length] != candidate.rendered_digest: + continue + prefix = rendered[: candidate.rendered_length] + raw = apply_prefix_edits(prefix, candidate.edits) + if token_digest(raw) != candidate.raw_digest: + continue + return PrefixMatch( + rendered_length=candidate.rendered_length, + raw_prefix=tuple(raw), + edits=candidate.edits, + ) + return None + + +def compact_candidate_payload( + candidate: CompactPrefixCandidate, +) -> dict[str, object]: + return { + "rendered_length": candidate.rendered_length, + "rendered_digest": candidate.rendered_digest, + "raw_digest": candidate.raw_digest, + "edits": [ + [edit.start, edit.stop, list(edit.replacement)] for edit in candidate.edits + ], + } + + +def compact_candidate_from_payload(value: object) -> CompactPrefixCandidate: + if not isinstance(value, dict): + raise ValueError("token-prefix candidate must be an object") + payload: dict[str, object] = { + key: item for key, item in value.items() if isinstance(key, str) + } + if len(payload) != len(value): + raise ValueError("token-prefix candidate keys must be strings") + rendered_length = payload.get("rendered_length") + rendered_digest = payload.get("rendered_digest") + raw_digest = payload.get("raw_digest") + raw_edits = payload.get("edits") + if ( + not isinstance(rendered_length, int) + or isinstance(rendered_length, bool) + or not 0 < rendered_length <= _MAX_PREFIX_TOKENS + or not isinstance(rendered_digest, str) + or not isinstance(raw_digest, str) + or any( + len(digest) != 64 + or any(character not in string.hexdigits for character in digest) + for digest in (rendered_digest, raw_digest) + ) + or not isinstance(raw_edits, list) + or len(raw_edits) > 256 + ): + raise ValueError("invalid token-prefix candidate") + edits: list[PrefixEdit] = [] + cursor = 0 + replacement_count = 0 + for raw_edit in raw_edits: + if not isinstance(raw_edit, list) or len(raw_edit) != 3: + raise ValueError("invalid token-prefix edit") + start = raw_edit[0] + stop = raw_edit[1] + replacement = raw_edit[2] + if ( + not isinstance(start, int) + or isinstance(start, bool) + or not isinstance(stop, int) + or isinstance(stop, bool) + or not isinstance(replacement, list) + ): + raise ValueError("invalid token-prefix edit") + replacement_ids: list[int] = [] + for token_id in replacement: + if not isinstance(token_id, int) or isinstance(token_id, bool): + raise ValueError("invalid token-prefix edit") + replacement_ids.append(token_id) + if not cursor <= start <= stop <= rendered_length: + raise ValueError("token-prefix edits must be ordered and non-overlapping") + _token_bytes(replacement_ids) + replacement_count += len(replacement_ids) + if replacement_count > COMPACT_PREFIX_MAX_REPLACEMENT_IDS: + raise ValueError("token-prefix candidate has too many replacement tokens") + edits.append(PrefixEdit(start, stop, tuple(replacement_ids))) + cursor = stop + return CompactPrefixCandidate( + rendered_length=rendered_length, + rendered_digest=rendered_digest.lower(), + raw_digest=raw_digest.lower(), + edits=tuple(edits), + ) + + +def compact_candidates_header( + candidates: Sequence[CompactPrefixCandidate], +) -> str | None: + """Encode a complete candidate snapshot, or decline unsafe header growth.""" + prefix = f'{{"version":"{COMPACT_PREFIX_VERSION}","candidates":[' + suffix = "]}" + parts = [prefix] + size = len(prefix) + len(suffix) + for candidate in candidates: + encoded = json.dumps( + compact_candidate_payload(candidate), separators=(",", ":") + ) + separator = int(len(parts) > 1) + size += separator + len(encoded) + if size > COMPACT_PREFIX_HEADER_DECODE_LIMIT: + return None + if separator: + parts.append(",") + parts.append(encoded) + parts.append(suffix) + encoded = base64.urlsafe_b64encode( + gzip.compress("".join(parts).encode(), compresslevel=1) + ).decode() + return encoded if len(encoded) <= COMPACT_PREFIX_HEADER_LIMIT else None + + +def compact_candidates_from_header( + value: str, +) -> tuple[CompactPrefixCandidate, ...]: + if len(value) > COMPACT_PREFIX_HEADER_LIMIT: + raise ValueError("token-prefix candidate header is too large") + try: + compressed = base64.b64decode(value, altchars=b"-_", validate=True) + decompressor = zlib.decompressobj(16 + zlib.MAX_WBITS) + decoded = decompressor.decompress( + compressed, COMPACT_PREFIX_HEADER_DECODE_LIMIT + 1 + ) + if ( + len(decoded) > COMPACT_PREFIX_HEADER_DECODE_LIMIT + or decompressor.unconsumed_tail + ): + raise ValueError("token-prefix candidate header expands beyond its limit") + decoded += decompressor.flush( + COMPACT_PREFIX_HEADER_DECODE_LIMIT + 1 - len(decoded) + ) + if not decompressor.eof or decompressor.unused_data: + raise ValueError("token-prefix candidate header is malformed") + payload = json.loads(decoded) + except ( + binascii.Error, + json.JSONDecodeError, + UnicodeDecodeError, + zlib.error, + ) as error: + raise ValueError("token-prefix candidate header is malformed") from error + if ( + not isinstance(payload, dict) + or payload.get("version") != COMPACT_PREFIX_VERSION + or not isinstance(payload.get("candidates"), list) + or len(payload["candidates"]) > COMPACT_PREFIX_MAX_CANDIDATES + ): + raise ValueError("token-prefix candidate header is malformed") + return tuple( + compact_candidate_from_payload(candidate) for candidate in payload["candidates"] + ) + + +class _Edge: + __slots__ = ("label", "child") + + def __init__(self, label: tuple[int, ...], child: _Node) -> None: + self.label = label + self.child = child + + +@dataclass(slots=True) +class _Variant: + lineages: OrderedDict[str, int] + edits: tuple[PrefixEdit, ...] + + +class _Node: + __slots__ = ("children", "parent", "parent_token", "rendered_length", "variants") + + def __init__( + self, + parent: _Node | None = None, + parent_token: int | None = None, + rendered_length: int = 0, + ) -> None: + self.children: dict[int, _Edge] = {} + self.parent = parent + self.parent_token = parent_token + self.rendered_length = rendered_length + self.variants: dict[tuple[int, ...], _Variant] = {} + + +def _common_prefix_length( + values: Sequence[int], start: int, label: tuple[int, ...] +) -> int: + limit = min(len(values) - start, len(label)) + index = 0 + while index < limit and values[start + index] == label[index]: + index += 1 + return index + + +class TokenPrefixCache: + """A bounded radix trie retaining distinct observed raw tokenizations.""" + + def __init__( + self, + *, + max_variants: int = 4096, + max_token_ids: int = 2_000_000, + max_lineages_per_variant: int = 64, + ) -> None: + if min(max_variants, max_token_ids, max_lineages_per_variant) <= 0: + raise ValueError("token-prefix cache bounds must be positive") + self._root = _Node() + self._lru: OrderedDict[_Node, int] = OrderedDict() + self._variants = 0 + self._max_variants = max_variants + self._max_token_ids = max_token_ids + self._max_lineages_per_variant = max_lineages_per_variant + self._token_ids = 0 + self._observation = 0 + + def _matching_nodes(self, rendered_tokens: Sequence[int]) -> list[_Node]: + node = self._root + index = 0 + candidates: list[_Node] = [] + while index < len(rendered_tokens): + edge = node.children.get(rendered_tokens[index]) + if edge is None: + break + matched = _common_prefix_length(rendered_tokens, index, edge.label) + if matched != len(edge.label): + break + index += matched + node = edge.child + if node.variants: + candidates.append(node) + return candidates + + def lookup( + self, rendered_tokens: Sequence[int], lineage: str | None + ) -> PrefixMatch | None: + for candidate in reversed(self._matching_nodes(rendered_tokens)): + variants = candidate.variants + # Without a rollout identifier, never choose between distinct raw + # tokenizations of the same rendered history. + if lineage is None: + if len(variants) == 1: + raw, variant = next(iter(variants.items())) + self._lru.move_to_end(candidate) + return PrefixMatch(candidate.rendered_length, raw, variant.edits) + continue + matching = [ + (variant.lineages[lineage], raw, variant.edits) + for raw, variant in variants.items() + if lineage in variant.lineages + ] + if not matching: + continue + _, raw_prefix, edits = max(matching, key=lambda value: value[0]) + self._lru.move_to_end(candidate) + lineages = variants[raw_prefix].lineages + lineages.move_to_end(lineage) + return PrefixMatch(candidate.rendered_length, raw_prefix, edits) + return None + + def insert( + self, + rendered_prefix: Sequence[int], + raw_prefix: Sequence[int], + lineage: str, + edits: Sequence[PrefixEdit] | None = None, + ) -> bool: + if not rendered_prefix or not raw_prefix or not lineage: + return False + raw = tuple(raw_prefix) + cost = len(rendered_prefix) + len(raw) + if cost > self._max_token_ids: + return False + node = self._root + index = 0 + while index < len(rendered_prefix): + token = rendered_prefix[index] + edge = node.children.get(token) + if edge is None: + child = _Node( + parent=node, + parent_token=token, + rendered_length=len(rendered_prefix), + ) + node.children[token] = _Edge(tuple(rendered_prefix[index:]), child) + node = child + break + + matched = _common_prefix_length(rendered_prefix, index, edge.label) + if matched == len(edge.label): + index += matched + node = edge.child + continue + + middle = _Node( + parent=node, + parent_token=token, + rendered_length=index + matched, + ) + old_suffix = edge.label[matched:] + old_child = edge.child + old_child.parent = middle + old_child.parent_token = old_suffix[0] + middle.children[old_suffix[0]] = _Edge(old_suffix, old_child) + edge.label = edge.label[:matched] + edge.child = middle + node = middle + index += matched + if index < len(rendered_prefix): + new_token = rendered_prefix[index] + child = _Node( + parent=node, + parent_token=new_token, + rendered_length=len(rendered_prefix), + ) + node.children[new_token] = _Edge(tuple(rendered_prefix[index:]), child) + node = child + break + + exact_edits = tuple( + prefix_edits(rendered_prefix, raw_prefix) if edits is None else edits + ) + if apply_prefix_edits(rendered_prefix, exact_edits) != list(raw_prefix): + raise ValueError("token-prefix edits do not reconstruct the raw prefix") + variant = node.variants.get(raw) + if variant is None: + variant = node.variants[raw] = _Variant(OrderedDict(), exact_edits) + self._lru[node] = self._lru.get(node, 0) + cost + self._variants += 1 + self._token_ids += cost + else: + variant.edits = exact_edits + self._lru.move_to_end(node) + lineages = variant.lineages + self._observation += 1 + lineages[lineage] = self._observation + lineages.move_to_end(lineage) + while len(lineages) > self._max_lineages_per_variant: + lineages.popitem(last=False) + if node in self._lru: + self._lru.move_to_end(node) + self._evict() + return True + + def _evict(self) -> None: + while self._lru and ( + self._variants > self._max_variants or self._token_ids > self._max_token_ids + ): + self.evict_oldest() + + @property + def token_ids(self) -> int: + return self._token_ids + + def evict_oldest(self) -> bool: + if not self._lru: + return False + node, cost = self._lru.popitem(last=False) + self._token_ids -= cost + self._variants -= len(node.variants) + node.variants.clear() + self._prune(node) + return True + + def _prune(self, node: _Node) -> None: + while node.parent is not None: + parent = node.parent + token = node.parent_token + assert token is not None + if not node.variants and not node.children: + del parent.children[token] + node = parent + continue + if not node.variants and len(node.children) == 1: + child_edge = next(iter(node.children.values())) + parent_edge = parent.children[token] + parent_edge.label += child_edge.label + parent_edge.child = child_edge.child + child_edge.child.parent = parent + child_edge.child.parent_token = token + node = parent + continue + break + + +class TokenPrefixStore: + """LRU-bounded collection of model/tokenizer-scoped prefix tries.""" + + def __init__( + self, + *, + max_scopes: int = 64, + max_variants_per_scope: int = 4096, + max_token_ids_per_scope: int = 2_000_000, + max_token_ids: int = 4_000_000, + ) -> None: + if min(max_scopes, max_token_ids) <= 0: + raise ValueError("token-prefix store bounds must be positive") + self._max_scopes = max_scopes + self._max_variants = max_variants_per_scope + self._max_token_ids = max_token_ids_per_scope + self._max_total_token_ids = max_token_ids + self._token_ids = 0 + self._scopes: OrderedDict[str, TokenPrefixCache] = OrderedDict() + + def _cache(self, scope: str) -> TokenPrefixCache: + cache = self._scopes.get(scope) + if cache is None: + cache = TokenPrefixCache( + max_variants=self._max_variants, + max_token_ids=self._max_token_ids, + ) + self._scopes[scope] = cache + while len(self._scopes) > self._max_scopes: + _, removed = self._scopes.popitem(last=False) + self._token_ids -= removed.token_ids + else: + self._scopes.move_to_end(scope) + return cache + + def lookup( + self, scope: str, rendered_tokens: Sequence[int], lineage: str | None + ) -> PrefixMatch | None: + cache = self._scopes.get(scope) + if cache is None: + return None + self._scopes.move_to_end(scope) + return cache.lookup(rendered_tokens, lineage) + + def insert( + self, + scope: str, + rendered_prefix: Sequence[int], + raw_prefix: Sequence[int], + lineage: str, + edits: Sequence[PrefixEdit] | None = None, + ) -> bool: + cache = self._cache(scope) + before = cache.token_ids + inserted = cache.insert(rendered_prefix, raw_prefix, lineage, edits) + self._token_ids += cache.token_ids - before + while self._token_ids > self._max_total_token_ids and self._scopes: + oldest_scope, oldest = next(iter(self._scopes.items())) + before = oldest.token_ids + if not oldest.evict_oldest(): + self._scopes.pop(oldest_scope) + continue + self._token_ids -= before - oldest.token_ids + if oldest.token_ids == 0: + self._scopes.pop(oldest_scope) + return inserted + + +class CompactPrefixStore: + """Bounded shared snapshots keyed by scope and logical rollout lineage.""" + + def __init__( + self, + *, + max_candidates: int = 262_144, + max_candidates_per_lineage: int = COMPACT_PREFIX_MAX_CANDIDATES, + max_edit_token_ids: int = 1_000_000, + max_serialized_bytes_per_lineage: int = COMPACT_PREFIX_RESPONSE_LIMIT - 1024, + max_staged_attempts: int = 262_144, + staged_ttl: float = 300, + ) -> None: + if ( + min( + max_candidates, + max_candidates_per_lineage, + max_edit_token_ids, + max_serialized_bytes_per_lineage, + max_staged_attempts, + staged_ttl, + ) + <= 0 + ): + raise ValueError("compact token-prefix store bounds must be positive") + self._max_candidates = max_candidates + self._max_per_lineage = max_candidates_per_lineage + self._max_edit_token_ids = max_edit_token_ids + self._max_serialized_bytes_per_lineage = max_serialized_bytes_per_lineage + self._edit_token_ids = 0 + self._entries: OrderedDict[ + tuple[str, str, int, str], tuple[CompactPrefixCandidate, int, int] + ] = OrderedDict() + self._lineages: dict[ + tuple[str, str], OrderedDict[tuple[str, str, int, str], None] + ] = {} + self._lineage_bytes: dict[tuple[str, str], int] = {} + self._max_staged_attempts = max_staged_attempts + self._staged_ttl = staged_ttl + self._staged: OrderedDict[ + str, + tuple[ + float, + list[tuple[str, str, CompactPrefixCandidate]], + int, + ], + ] = OrderedDict() + self._staged_candidates = 0 + self._staged_edit_token_ids = 0 + self._staged_evictions = 0 + self._missing_commits = 0 + + @property + def staged_evictions(self) -> int: + return self._staged_evictions + + @property + def missing_commits(self) -> int: + return self._missing_commits + + def insert( + self, + scope: str, + lineage: str, + candidate: CompactPrefixCandidate, + ) -> bool: + cost = sum(len(edit.replacement) for edit in candidate.edits) + serialized = ( + len( + json.dumps( + compact_candidate_payload(candidate), separators=(",", ":") + ).encode() + ) + + 1 + ) + if ( + cost > self._max_edit_token_ids + or serialized > self._max_serialized_bytes_per_lineage + ): + return False + key = ( + scope, + lineage, + candidate.rendered_length, + candidate.rendered_digest, + ) + previous = self._entries.pop(key, None) + if previous is not None: + self._edit_token_ids -= previous[1] + self._lineage_bytes[(scope, lineage)] -= previous[2] + self._entries[key] = (candidate, cost, serialized) + self._edit_token_ids += cost + lineage_key = (scope, lineage) + self._lineage_bytes[lineage_key] = ( + self._lineage_bytes.get(lineage_key, 0) + serialized + ) + entries = self._lineages.setdefault(lineage_key, OrderedDict()) + entries.pop(key, None) + entries[key] = None + while ( + len(entries) > self._max_per_lineage + or self._lineage_bytes[lineage_key] > self._max_serialized_bytes_per_lineage + ): + self._remove(next(iter(entries))) + while ( + len(self._entries) > self._max_candidates + or self._edit_token_ids > self._max_edit_token_ids + ): + self._remove(next(iter(self._entries))) + return key in self._entries + + def candidates( + self, scope: str, lineage: str, max_rendered_length: int + ) -> list[CompactPrefixCandidate]: + entries = self._lineages.get((scope, lineage)) + if entries is None: + return [] + ordered = [ + self._entries[key][0] + for key in reversed(entries) + if key in self._entries + and self._entries[key][0].rendered_length <= max_rendered_length + ] + # Python's sort is stable, so equal-length variants remain newest first. + ordered.sort(key=lambda candidate: candidate.rendered_length, reverse=True) + return ordered + + def stage( + self, + attempt: str, + observations: Sequence[tuple[str, str, CompactPrefixCandidate]], + ) -> bool: + """Retain observations invisibly until their upstream attempt commits.""" + now = time.monotonic() + self._expire_staged(now) + existing = self._staged.pop(attempt, None) + entries = [] if existing is None else existing[1] + cost = 0 if existing is None else existing[2] + if existing is not None: + self._staged_candidates -= len(entries) + self._staged_edit_token_ids -= cost + combined = OrderedDict( + ( + ( + scope, + lineage, + candidate.rendered_length, + candidate.rendered_digest, + ), + (scope, lineage, candidate), + ) + for scope, lineage, candidate in entries + ) + for scope, lineage, candidate in observations: + candidate_cost = sum(len(edit.replacement) for edit in candidate.edits) + if candidate_cost <= self._max_edit_token_ids: + key = ( + scope, + lineage, + candidate.rendered_length, + candidate.rendered_digest, + ) + combined.pop(key, None) + combined[key] = (scope, lineage, candidate) + entries = list(combined.values()) + cost = sum( + len(edit.replacement) + for _, _, candidate in entries + for edit in candidate.edits + ) + if entries: + self._staged[attempt] = (now, entries, cost) + self._staged_candidates += len(entries) + self._staged_edit_token_ids += cost + while self._staged and ( + len(self._staged) > self._max_staged_attempts + or self._staged_candidates > self._max_candidates + or self._staged_edit_token_ids > self._max_edit_token_ids + ): + self._remove_staged(next(iter(self._staged))) + self._staged_evictions += 1 + return attempt in self._staged + + def commit(self, attempt: str) -> bool: + """Publish only observations produced by the selected upstream attempt.""" + self._expire_staged(time.monotonic()) + staged = self._staged.get(attempt) + if staged is None: + self._missing_commits += 1 + return False + observations = list(staged[1]) + self._remove_staged(attempt) + inserted = [ + self.insert(scope, lineage, candidate) + for scope, lineage, candidate in observations + ] + return all(inserted) + + def _expire_staged(self, now: float) -> None: + while self._staged: + attempt, (created, _, _) = next(iter(self._staged.items())) + if now - created < self._staged_ttl: + break + self._remove_staged(attempt) + + def _remove_staged(self, attempt: str) -> None: + removed = self._staged.pop(attempt, None) + if removed is None: + return + self._staged_candidates -= len(removed[1]) + self._staged_edit_token_ids -= removed[2] + + def _remove(self, key: tuple[str, str, int, str]) -> None: + removed = self._entries.pop(key, None) + if removed is None: + return + self._edit_token_ids -= removed[1] + lineage_key = key[:2] + self._lineage_bytes[lineage_key] -= removed[2] + entries = self._lineages[lineage_key] + entries.pop(key, None) + if not entries: + del self._lineages[lineage_key] + del self._lineage_bytes[lineage_key] + + +class TokenPrefixRuntime: + """Local hot cache with shared lookup and batched replication.""" + + def __init__( + self, + *, + lookup_shared: Callable[[str, list[int], str], Awaitable[PrefixMatch | None]], + insert_shared_many: Callable[[list[_SharedPrefixEntry]], Awaitable[None]], + max_pending_batches: int = 32, + ) -> None: + if max_pending_batches <= 0: + raise ValueError("replication bound must be positive") + self._local = TokenPrefixStore() + self._lookup_shared = lookup_shared + self._insert_shared_many = insert_shared_many + self._pending: asyncio.Queue[list[_SharedPrefixEntry]] = asyncio.Queue( + max_pending_batches + ) + self._replicator: asyncio.Task[None] | None = None + self._closed = False + + async def rewrite_with_edits( + self, + scope: str, + rendered_tokens: Sequence[int], + lineage: str, + *, + fallback_lineage: str | None = None, + shared_candidate: bool, + shared_match: PrefixMatch | None = None, + shared_resolved: bool = False, + ) -> tuple[list[int], list[int], tuple[PrefixEdit, ...]]: + """Return canonical/input tokens plus their exact sparse substitution.""" + canonical = list(rendered_tokens) + local_match = self._local.lookup(scope, canonical, lineage) + if local_match is None and fallback_lineage is not None: + local_match = self._local.lookup(scope, canonical, fallback_lineage) + if local_match is not None: + self._local.insert( + scope, + canonical[: local_match.rendered_length], + local_match.raw_prefix, + lineage, + local_match.edits, + ) + if shared_candidate and not shared_resolved: + shared_match = await self._lookup_shared(scope, canonical, lineage) + if shared_match is None and fallback_lineage is not None: + shared_match = await self._lookup_shared( + scope, canonical, fallback_lineage + ) + if shared_match is not None and ( + local_match is None + or shared_match.rendered_length > local_match.rendered_length + ): + exact_edits = shared_match.edits + rendered_prefix = canonical[: shared_match.rendered_length] + if not exact_edits and list(shared_match.raw_prefix) != rendered_prefix: + exact_edits = prefix_edits(rendered_prefix, shared_match.raw_prefix) + self._local.insert( + scope, + rendered_prefix, + shared_match.raw_prefix, + lineage, + exact_edits, + ) + match = PrefixMatch( + shared_match.rendered_length, + shared_match.raw_prefix, + exact_edits, + ) + else: + match = local_match + if match is None: + return canonical, canonical.copy(), () + return ( + canonical, + [*match.raw_prefix, *canonical[match.rendered_length :]], + match.edits, + ) + + async def insert_many( + self, + entries: Sequence[_SharedPrefixEntry], + ) -> None: + normalized: list[_SharedPrefixEntry] = [] + for scope, rendered, raw, lineage, edits in entries: + exact_edits = tuple(edits) + self._local.insert(scope, rendered, raw, lineage, exact_edits) + normalized.append((scope, rendered, raw, lineage, exact_edits)) + if not normalized or self._closed: + return + try: + self._pending.put_nowait(normalized) + except asyncio.QueueFull: + # Let the already-scheduled replicator drain a simultaneous burst + # once before dropping cross-replica work. + self._ensure_replicator() + await asyncio.sleep(0) + if self._closed: + return + try: + self._pending.put_nowait(normalized) + except asyncio.QueueFull: + logger.warning( + "shared token-prefix replication queue is full; " + "retaining local mappings only" + ) + return + self._ensure_replicator() + + def _ensure_replicator(self) -> None: + if self._replicator is None or self._replicator.done(): + self._replicator = asyncio.create_task( + self._replicate(), name="caladan-token-prefix-replication" + ) + + async def _replicate(self) -> None: + try: + while not self._pending.empty(): + batches = [self._pending.get_nowait()] + while True: + try: + batches.append(self._pending.get_nowait()) + except asyncio.QueueEmpty: + break + try: + entries = [entry for batch in batches for entry in batch] + await self._insert_shared_many(entries) + except asyncio.CancelledError: + raise + except Exception: + logger.warning( + "shared token-prefix insertion failed; " + "retaining local mappings", + exc_info=True, + ) + finally: + for _ in batches: + self._pending.task_done() + finally: + if self._replicator is asyncio.current_task(): + self._replicator = None + if not self._closed and not self._pending.empty(): + self._ensure_replicator() + + async def close(self, *, timeout: float = 2.0) -> bool: + """Drain pending replication for at most timeout seconds, then stop.""" + self._closed = True + drained = await self.flush(timeout=timeout) + replicator = self._replicator + if replicator is not None: + replicator.cancel() + with suppress(asyncio.CancelledError): + await replicator + if self._replicator is replicator: + self._replicator = None + while not self._pending.empty(): + self._pending.get_nowait() + self._pending.task_done() + return drained + + async def flush(self, *, timeout: float | None = None) -> bool: + """Wait until queued shared replication has completed.""" + if self._replicator is None: + return True + try: + await asyncio.wait_for(self._pending.join(), timeout) + except TimeoutError: + logger.warning("token-prefix replication did not drain before timeout") + return False + return True diff --git a/src/art_inference/vllm.py b/src/art_inference/vllm.py new file mode 100644 index 000000000..ec007a576 --- /dev/null +++ b/src/art_inference/vllm.py @@ -0,0 +1,445 @@ +"""Preserve sampled history before an ART-owned vLLM engine serves the next turn.""" + +from __future__ import annotations + +from collections.abc import AsyncIterator, Awaitable, Callable +from contextvars import ContextVar +import copy +from dataclasses import dataclass, field +from functools import wraps +import hashlib +import importlib +import inspect +import json +from types import MethodType +from typing import Any + +from .append_only import ( + aligned_values, + chat_response_prefixes, + merge_chat_delta, + output_prefix_observations, + patch_deepseek_renderer, + patch_harmony, + preserves_history, + shifted_span, +) +from .chat_template import ( + chat_template_with_preserved_thinking, + configure_preserved_thinking_chat_template, +) +from .token_prefix import TokenPrefixStore, prefix_edits + + +@dataclass +class _Turn: + scope: str + tokenizer: Any + rendered: list[int] = field(default_factory=list) + prompt: list[int] = field(default_factory=list) + outputs: dict[int, tuple[list[int], bool]] = field(default_factory=dict) + external_request: Any = None + external_observer: Callable[[Any, list[Any]], Awaitable[None]] | None = None + enabled: bool = True + + +_CURRENT: ContextVar[_Turn | None] = ContextVar("art_history", default=None) +_PREFIXES = TokenPrefixStore() +_DROP_THINKING: ContextVar[bool] = ContextVar("art_vllm_drop_thinking", default=False) + + +def _configure_native_tokenizer(tokenizer, importer): + apply = getattr(tokenizer, "apply_chat_template", None) + if ( + not callable(apply) + or getattr(apply, "__module__", None) != "vllm.tokenizers.deepseek_v32" + ): + return + module = importer("vllm.tokenizers.deepseek_v32") + if not getattr(module, "_art_append_only", False): + encode = module.encode_messages + patch_deepseek_renderer( + importer("vllm.tokenizers.deepseek_v32_encoding"), + lambda: not _DROP_THINKING.get(), + prefix_only=True, + ) + + @wraps(encode) + def encode_messages(*args, **kwargs): + kwargs["drop_thinking"] = _DROP_THINKING.get() + return encode(*args, **kwargs) + + module.encode_messages = encode_messages + module._art_append_only = True + + @wraps(apply) + def apply_template(self, *args, **kwargs): + token = _DROP_THINKING.set(not preserves_history(kwargs)) + try: + return apply(*args, **kwargs) + finally: + _DROP_THINKING.reset(token) + + tokenizer.apply_chat_template = MethodType(apply_template, tokenizer) + + +def set_external_history_observer( + observer: Callable[[Any, list[Any]], Awaitable[None]], + importer=importlib.import_module, +) -> None: + """Use a hosting service's transactional/distributed store for its requests.""" + base = importer("vllm.renderers.base").BaseRenderer + base._art_history_external_observer = staticmethod(observer) + + +async def _observe(turn: _Turn, entries) -> None: + if turn.external_request is not None: + assert turn.external_observer is not None + await turn.external_observer(turn.external_request, entries) + else: + for rendered, raw, edits in entries: + _PREFIXES.insert(turn.scope, rendered, raw, "content", edits) + + +def _tokens(prompt: Any) -> list[int] | None: + if isinstance(prompt, dict): + tokens = prompt.get("prompt_token_ids") + if isinstance(tokens, (list, tuple)): + return list(tokens) + return None + + +def replace_prompt_tokens(prompt, tokens, edits=None): + """Keep native multimodal offsets and per-token metadata aligned.""" + result = {**prompt, "prompt_token_ids": tokens} + if tokens == prompt["prompt_token_ids"]: + return result + edits = prefix_edits(prompt["prompt_token_ids"], tokens) if edits is None else edits + if prompt.get("mm_placeholders"): + result["mm_placeholders"] = {} + for modality, spans in prompt["mm_placeholders"].items(): + result["mm_placeholders"][modality] = [] + for span in spans: + shifted = copy.copy(span) + shifted.offset = shifted_span(span.offset, span.length, edits) + result["mm_placeholders"][modality].append(shifted) + for name in ( + "is_token_ids", + "prompt_is_token_ids", + "token_type_ids", + "assistant_tokens_mask", + ): + if prompt.get(name) is not None: + fill = ( + True + if name in ("is_token_ids", "prompt_is_token_ids") + else (0 if name == "token_type_ids" else None) + ) + result[name] = aligned_values(prompt[name], edits, fill=fill) + # Text offsets refer to the normalized string, which is no longer the + # source of these exact IDs. They are optional engine metadata. + result.pop("prompt_token_offsets", None) + return result + + +def patch_history(importer=importlib.import_module) -> None: + base = importer("vllm.renderers.base").BaseRenderer + if getattr(base, "_art_append_only", False): + return + init = base.__init__ + + @wraps(init) + def initialize(self, config, tokenizer): + if tokenizer is not None: + configure_preserved_thinking_chat_template(tokenizer) + _configure_native_tokenizer(tokenizer, importer) + init(self, config, tokenizer) + + base.__init__ = initialize + base._art_append_only = True + patch_harmony( + importer("vllm.entrypoints.openai.parser.harmony_utils"), + lambda: (turn := _CURRENT.get()) is None or turn.enabled, + "vllm.", + ) + + params = importer("vllm.renderers.params").ChatParams + apply_kwargs = params.get_apply_chat_template_kwargs + + @wraps(apply_kwargs) + def template_kwargs(self): + kwargs = { + "preserve_thinking": True, + "clear_thinking": False, + "drop_thinking": False, + **apply_kwargs(self), + } + if kwargs.get("chat_template") is not None: + kwargs["chat_template"] = chat_template_with_preserved_thinking( + kwargs["chat_template"] + ) + return kwargs + + params.get_apply_chat_template_kwargs = template_kwargs + + responses_utils = importer("vllm.entrypoints.openai.responses.utils") + responses_serving = importer("vllm.entrypoints.openai.responses.serving") + construct = responses_utils.construct_input_messages + + @wraps(construct) + def input_messages(*, prev_response_output=None, request_input, **kwargs): + # vLLM's previous_response_id path otherwise drops reasoning and tool + # calls, although its ordinary input path already knows how to retain + # both. Keep the protocol's existing instructions replacement semantics. + if prev_response_output: + request_input = [ + *prev_response_output, + *( + [{"role": "user", "content": request_input}] + if isinstance(request_input, str) + else request_input + ), + ] + return construct(request_input=request_input, **kwargs) + + setattr(responses_utils, "construct_input_messages", input_messages) + setattr(responses_serving, "construct_input_messages", input_messages) + + engine = importer("vllm.v1.engine.async_llm").AsyncLLM + generate = engine.generate + signature = inspect.signature(generate) + + @wraps(generate) + def serve(self, prompt, *args, **kwargs): + turn = _CURRENT.get() + tokens = _tokens(prompt) + if turn is None or not turn.enabled or tokens is None: + return generate(self, prompt, *args, **kwargs) + turn.rendered = tokens + # Batched completions run several generators concurrently in one API + # request. Each generator must publish only its own prompt and outputs. + outputs = {} + turn.outputs = outputs + sampling = signature.bind(self, prompt, *args, **kwargs).arguments.get( + "sampling_params" + ) + delta = getattr(getattr(sampling, "output_kind", None), "name", None) == "DELTA" + match = ( + _PREFIXES.lookup(turn.scope, tokens, None) + if turn.external_request is None + else None + ) + raw_prompt = ( + list(match.raw_prefix) + tokens[match.rendered_length :] + if match + else tokens + ) + turn.prompt = raw_prompt + prompt = replace_prompt_tokens(prompt, raw_prompt, match.edits if match else ()) + + async def stream(): + async for output in generate(self, prompt, *args, **kwargs): + for choice in output.outputs: + generated_ids = list(choice.token_ids) + if delta: + generated_ids = ( + outputs.get(choice.index, ([], False))[0] + generated_ids + ) + outputs[choice.index] = ( + generated_ids, + choice.finish_reason not in (None, "length", "abort"), + ) + yield output + # Decode/re-encode differences matter for completions and for chat + # renderers that already preserve the sampled text exactly. + for raw, _ in outputs.values(): + if turn.external_request is None: + await _observe( + turn, + output_prefix_observations( + turn.tokenizer, tokens, raw_prompt, raw + ), + ) + + return stream() + + engine.generate = serve + chat = importer("vllm.entrypoints.openai.chat_completion.serving").OpenAIServingChat + completion = importer( + "vllm.entrypoints.openai.completion.serving" + ).OpenAIServingCompletion + + def install(owner, method, *, observe_chat=False, observe_responses=False): + original = getattr(owner, method) + + @wraps(original) + async def create(self, request, raw_request=None): + headers = getattr(raw_request, "headers", {}) or {} + # Caladan supplies a distributed store and its own engine adapter. + # Its observer uses the same shared ART implementation. + external = bool(headers.get("x-caladan-prefix-scope")) + external_observer = getattr(base, "_art_history_external_observer", None) + if external and external_observer is None: + return await original(self, request, raw_request) + tokenizer = getattr(self.renderer, "tokenizer", None) + if tokenizer is None: + return await original(self, request, raw_request) + material = [ + request.model, + getattr(request, "chat_template", None), + getattr(request, "chat_template_kwargs", None), + headers.get("authorization", ""), + ] + scope = hashlib.sha256( + json.dumps(material, sort_keys=True).encode() + ).hexdigest() + turn = _Turn( + scope, + tokenizer, + external_request=raw_request if external else None, + external_observer=external_observer, + enabled=preserves_history( + getattr(request, "chat_template_kwargs", None) + ), + ) + + async def observe(messages): + if not turn.enabled or not observe_chat or not turn.prompt: + return + + async def render(value): + result = await self.render_chat_request(value) + if not isinstance(result, tuple) or len(result[1]) != 1: + return None + return _tokens(result[1][0]) + + choices = [ + (message, *turn.outputs[index]) + for index, message in messages.items() + if index in turn.outputs + ] + entries = await chat_response_prefixes( + tokenizer, request, turn.prompt, choices, render + ) + await _observe(turn, entries) + + async def observe_response(response): + if ( + not observe_responses + or not turn.enabled + or getattr(self, "use_harmony", False) + or not getattr(response, "output", None) + or not turn.prompt + ): + return + messages = responses_utils.construct_chat_messages_with_tool_call( + response.output + ) + # Built-in tool execution can contain several generations. The + # engine-level observations cover those individual turns. + if len(messages) != 1 or messages[0].get("role") != "assistant": + return + previous = None + if request.previous_response_id: + async with self.response_store_lock: + previous = self.response_store.get(request.previous_response_id) + conversation, inputs = await self._make_request(request, previous) + if len(inputs) != 1: + return + protocol = importer("vllm.entrypoints.openai.chat_completion.protocol") + tools = responses_utils.construct_tool_dicts( + request.tools, request.tool_choice + ) + view = protocol.ChatCompletionRequest( + model=request.model, + messages=conversation, + chat_template_kwargs=self._effective_chat_template_kwargs(request), + ) + + async def render(value): + online = getattr(self, "online_renderer", None) + if online is None: # vLLM 0.24 + online = self.openai_serving_render + _, inputs = await online.preprocess_chat( + value, + value.messages, + default_template=self.chat_template, + default_template_content_format=self.chat_template_content_format, + default_template_kwargs=self._effective_chat_template_kwargs( + request + ), + tool_dicts=tools, + parser=self.parser, + ) + return _tokens(inputs[0]) if len(inputs) == 1 else None + + if await render(view) != _tokens(inputs[0]): + return + if 0 in turn.outputs: + await _observe( + turn, + await chat_response_prefixes( + tokenizer, + view, + turn.prompt, + [(messages[0], *turn.outputs[0])], + render, + ), + ) + + token = _CURRENT.set(turn) + try: + result = await original(self, request, raw_request) + if not isinstance(result, AsyncIterator): + choices = getattr(result, "choices", ()) + await observe( + { + c.index: c.message.model_dump(exclude_none=True) + for c in choices + if hasattr(c, "message") + } + ) + await observe_response(result) + return result + finally: + _CURRENT.reset(token) + + async def stream(): + token = _CURRENT.set(turn) + messages = {} + try: + async for event in result: + if ( + observe_chat + and isinstance(event, str) + and event.startswith("data: ") + ): + data = event[6:].strip() + if data == "[DONE]": + await observe(messages) + else: + for choice in json.loads(data).get("choices", []): + message = messages.setdefault( + choice["index"], {"role": "assistant"} + ) + merge_chat_delta(message, choice.get("delta") or {}) + if observe_responses and getattr(event, "type", None) in ( + "response.completed", + "response.incomplete", + ): + await observe_response(event.response) + yield event + finally: + _CURRENT.reset(token) + + return stream() + + setattr(owner, method, create) + + install(chat, "create_chat_completion", observe_chat=True) + install(completion, "create_completion", observe_chat=False) + install( + responses_serving.OpenAIServingResponses, + "create_responses", + observe_responses=True, + ) diff --git a/tests/unit/test_append_only.py b/tests/unit/test_append_only.py new file mode 100644 index 000000000..99bff1ad5 --- /dev/null +++ b/tests/unit/test_append_only.py @@ -0,0 +1,195 @@ +import asyncio +from types import SimpleNamespace + +from pydantic import BaseModel +import pytest + +from art.token_prefix import TokenPrefixCache, apply_prefix_edits +from art.utils.append_only import chat_prefix_observations, chat_response_prefixes +from art_inference.append_only import ( + output_prefix_observations, + patch_deepseek_renderer, +) + + +@pytest.mark.parametrize("reasoning", [None, "", "\nthought\n"]) +def test_dsv32_earlier_turn_framing_survives_a_new_user_turn(reasoning): + def render(index, messages, thinking_mode): + last_user = max(i for i, m in enumerate(messages) if m["role"] == "user") + thinking = thinking_mode == "thinking" + if messages[index]["role"] == "user": + return "USER" + ( + "" if thinking and index == last_user else "" + ) + if thinking and index > last_user: + assert messages[index]["reasoning_content"], "Missing reasoning" + return messages[index]["reasoning_content"] + "END" + return "END" + + encoding = SimpleNamespace(render_message=render, thinking_end_token="") + preserve = True + patch_deepseek_renderer(encoding, lambda: preserve, prefix_only=True) + messages = [ + {"role": "user"}, + {"role": "assistant", "reasoning_content": reasoning}, + ] + + def encode(messages): + return "".join( + encoding.render_message(i, messages, "thinking") + for i in range(len(messages)) + ) + + completed = encode(messages) + assert encode([*messages, {"role": "user"}]).startswith(completed) + if reasoning: + preserve = False + assert not encode([*messages, {"role": "user"}]).startswith(completed) + + +class Tokenizer: + def encode(self, text): + return list(text.encode()) + + def decode(self, tokens, *, skip_special_tokens=False): + return bytes(tokens).decode() + + +def test_many_protocol_markers_do_not_multiply_long_prompt_storage(): + tokenizer = SimpleNamespace( + all_special_ids=[1, 2, 3], + encode=lambda text, **_: [ord(char) for char in text], + decode=lambda tokens, **_: "".join(chr(token) for token in tokens), + ) + prompt = [100] * 128_000 + output = [101, 1, 102, 2, 103, 3] * 100 + entries = output_prefix_observations(tokenizer, prompt, prompt, output) + assert len(entries) <= 2 + assert entries[-1][1] == prompt + output + + +@pytest.mark.parametrize("action", ["pass", '{"x": 1}']) +def test_normalized_reasoning_survives_an_edited_action(action): + tokenizer = Tokenizer() + encode = tokenizer.encode + prompt = encode("USER question ASSISTANT \n") + rendered = prompt + encode(f"reasoning\n\n\n{action}END\n") + sampled = encode(f"\nreasoning\n\n\n{action}END") + reasoning = prompt + encode("reasoning\n\n\nEND\n") + observations = chat_prefix_observations( + tokenizer, prompt, rendered, prompt, sampled, reasoning_prompt=reasoning + ) + assert len(observations) == 2 + cache = TokenPrefixCache() + for rendered_prefix, raw_prefix, edits in observations: + assert apply_prefix_edits(rendered_prefix, edits) == raw_prefix + cache.insert(rendered_prefix, raw_prefix, "rollout", edits) + + for next_action in (action, "corrected action"): + next_prompt = prompt + encode( + f"reasoning\n\n\n{next_action}END\nUSER next ASSISTANT" + ) + match = cache.lookup(next_prompt, "rollout") + assert match is not None + served = list(match.raw_prefix) + next_prompt[match.rendered_length :] + assert served == prompt + encode( + f"\nreasoning\n\n\n{next_action}END\nUSER next ASSISTANT" + ) + + +def test_truncation_does_not_replace_completed_turn_framing(): + tokenizer = Tokenizer() + encode = tokenizer.encode + assert not chat_prefix_observations( + tokenizer, + encode("prompt"), + encode("prompt partial END"), + encode("prompt"), + encode(" partial"), + complete=False, + ) + + +def test_changed_prompt_cannot_be_registered_as_a_continuation(): + assert not chat_prefix_observations(Tokenizer(), [1], [2, 3], [4], [5]) + + +@pytest.mark.parametrize( + "options", + [{"preserve_thinking": False}, {"clear_thinking": True}, {"drop_thinking": True}], +) +def test_explicit_rewriting_bypasses_observation(options): + class Request(BaseModel): + messages: list[dict] + chat_template_kwargs: dict + + async def render(_): + pytest.fail("Explicit rewriting must not create preservation mappings") + + assert ( + asyncio.run( + chat_response_prefixes( + Tokenizer(), + Request(messages=[], chat_template_kwargs=options), + [], + [], + render, + ) + ) + == [] + ) + + +def test_custom_stop_does_not_delete_the_template_terminator(): + encode = Tokenizer().encode + assert ( + chat_prefix_observations( + Tokenizer(), + encode("prompt:"), + encode("prompt:answerEND"), + encode("prompt:"), + encode("answer"), + complete=True, + ) + == [] + ) + + +def test_reasoning_boundary_does_not_absorb_whitespace_merged_into_action(): + from art_inference.append_only import _whitespace_prefix + + tokenizer = Tokenizer() + raw = tokenizer.encode("thought\n\naction") + boundary = _whitespace_prefix(tokenizer, tokenizer.encode("thought"), raw) + assert tokenizer.decode(raw[:boundary]) == "thought" + + +def test_missing_lineage_never_selects_an_ambiguous_tokenization(): + cache = TokenPrefixCache() + cache.insert([1, 2], [1, 3], "first") + match = cache.lookup([1, 2], None) + assert match is not None and match.raw_prefix == (1, 3) + cache.insert([1, 2], [1, 4], "second") + assert cache.lookup([1, 2], None) is None + first = cache.lookup([1, 2], "first") + second = cache.lookup([1, 2], "second") + assert first is not None and first.raw_prefix == (1, 3) + assert second is not None and second.raw_prefix == (1, 4) + + +def test_history_edits_leave_multimodal_markers_outside_replacements(): + from types import SimpleNamespace + + from art_inference.append_only import _rendering_edits + from art_inference.vllm import replace_prompt_tokens + + rendered = [1, 2, 9, 9, 3, 4] + raw = [1, 1, 2, 9, 9, 3, 4, 4] + edits = _rendering_edits(SimpleNamespace(all_special_ids=[9]), rendered, raw) + assert apply_prefix_edits(rendered, edits) == raw + original = { + "prompt_token_ids": rendered, + "mm_placeholders": {"image": [SimpleNamespace(offset=2, length=2)]}, + } + updated = replace_prompt_tokens(original, raw, edits) + assert updated["mm_placeholders"]["image"][0].offset == 3 diff --git a/tests/unit/test_inference_history.py b/tests/unit/test_inference_history.py new file mode 100644 index 000000000..a823a9bce --- /dev/null +++ b/tests/unit/test_inference_history.py @@ -0,0 +1,401 @@ +import asyncio +import json +from types import SimpleNamespace + +from pydantic import BaseModel, model_validator +import pytest + +from art_inference import vllm +from art_inference.token_prefix import TokenPrefixStore + + +class Request(BaseModel): + model: str = "model" + messages: list[dict] + stream: bool = False + stream_options: dict | None = None + max_tokens: int | None = None + max_completion_tokens: int | None = None + truncate_prompt_tokens: int | None = None + add_generation_prompt: bool = True + continue_final_message: bool = False + chat_template_kwargs: dict = {} + previous_response_id: str | None = None + tools: list = [] + tool_choice: str = "auto" + + @model_validator(mode="after") + def validate_stream(self): + if self.stream_options and not self.stream: + raise ValueError("Stream options can only be defined when stream=True") + return self + + +class Message(BaseModel): + role: str = "assistant" + reasoning_content: str + content: str + + +@pytest.fixture +def serving(monkeypatch): + monkeypatch.setattr(vllm, "_PREFIXES", TokenPrefixStore()) + + class Tokenizer: + all_special_ids = [ord("#"), ord("~")] + + def encode(self, text, **kwargs): + return list(text.encode()) + + def decode(self, tokens, **kwargs): + return bytes(tokens).decode() + + class Renderer: + def __init__(self, config, tokenizer): + self.tokenizer = tokenizer + + class Params: + def get_apply_chat_template_kwargs(self): + return {} + + class Engine: + def __init__(self): + self.prompts = [] + + async def generate(self, prompt, sampling_params=None): + self.prompts.append(prompt["prompt_token_ids"]) + chunks = ( + [b"\nthought\n#", b"action~"] + if sampling_params + else [b"\nthought\n#action~"] + ) + for index, tokens in enumerate(chunks): + yield SimpleNamespace( + outputs=[ + SimpleNamespace( + index=0, + token_ids=list(tokens), + finish_reason="stop" if index == len(chunks) - 1 else None, + ) + ] + ) + + class Serving: + def __init__(self): + self.renderer = Renderer(None, Tokenizer()) + self.engine = Engine() + self.online_renderer = SimpleNamespace(preprocess_chat=self.preprocess_chat) + self.chat_template = None + self.chat_template_content_format = "string" + self.parser = None + + async def preprocess_chat(self, request, messages, **kwargs): + return await self.render_chat_request(request) + + def _effective_chat_template_kwargs(self, request): + return request.chat_template_kwargs + + async def _make_request(self, request, previous): + return request.messages, (await self.render_chat_request(request))[1] + + async def render_chat_request(self, request): + text = "" + for message in request.messages: + if message["role"] == "assistant": + text += ( + "A" + + message.get("reasoning_content", "").strip() + + "#" + + (message.get("content") or "") + + "~\n" + ) + else: + text += "U" + message["content"] + ";" + if request.add_generation_prompt: + text += "A" + return [], [{"prompt_token_ids": list(text.encode())}] + + async def create_chat_completion(self, request, raw_request=None): + _, inputs = await self.render_chat_request(request) + message = Message(reasoning_content="\nthought\n", content="action") + + async def stream(): + params = SimpleNamespace(output_kind=SimpleNamespace(name="DELTA")) + async for _ in self.engine.generate(inputs[0], params): + pass + yield ( + "data: " + + json.dumps( + {"choices": [{"index": 0, "delta": message.model_dump()}]} + ) + + "\n\n" + ) + yield "data: [DONE]\n\n" + + if request.stream: + return stream() + async for _ in self.engine.generate(inputs[0]): + pass + return SimpleNamespace(choices=[SimpleNamespace(index=0, message=message)]) + + async def create_completion(self, request, raw_request=None): + return await self.create_chat_completion(request, raw_request) + + async def create_responses(self, request, raw_request=None): + _, inputs = await self.render_chat_request(request) + async for _ in self.engine.generate(inputs[0]): + pass + response = SimpleNamespace( + output=[ + Message( + reasoning_content="\nthought\n", content="action" + ).model_dump() + ] + ) + if not request.stream: + return response + + async def events(): + yield SimpleNamespace(type="response.completed", response=response) + + return events() + + def construct_input_messages(**kwargs): + return kwargs + + utils = SimpleNamespace( + construct_input_messages=construct_input_messages, + construct_chat_messages_with_tool_call=lambda items: items, + construct_tool_dicts=lambda *_: None, + ) + modules = { + "vllm.renderers.base": SimpleNamespace(BaseRenderer=Renderer), + "vllm.renderers.params": SimpleNamespace(ChatParams=Params), + "vllm.v1.engine.async_llm": SimpleNamespace(AsyncLLM=Engine), + "vllm.entrypoints.openai.chat_completion.serving": SimpleNamespace( + OpenAIServingChat=Serving + ), + "vllm.entrypoints.openai.completion.serving": SimpleNamespace( + OpenAIServingCompletion=Serving + ), + "vllm.entrypoints.openai.responses.serving": SimpleNamespace( + OpenAIServingResponses=Serving + ), + "vllm.entrypoints.openai.responses.utils": utils, + "vllm.entrypoints.openai.parser.harmony_utils": SimpleNamespace( + render_for_completion=lambda messages: messages + ), + "vllm.entrypoints.openai.chat_completion.protocol": SimpleNamespace( + ChatCompletionRequest=Request + ), + } + vllm.patch_history(modules.__getitem__) + return Serving(), modules + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("action", ["action", "edited action"]) +def test_served_history_survives_renderer_normalization(serving, stream, action): + server, _ = serving + + async def run(): + user = {"role": "user", "content": "question"} + first = Request( + messages=[user], + stream=stream, + stream_options={"include_usage": True} if stream else None, + ) + response = await server.create_chat_completion(first) + if stream: + assert (await anext(response)).startswith("data: ") + assert await anext(response) == "data: [DONE]\n\n" + await response.aclose() + prior = server.engine.prompts[0] + list(b"\nthought\n#") + assistant = Message( + reasoning_content="\nthought\n", content=action + ).model_dump() + await server.create_chat_completion( + Request(messages=[user, assistant, {"role": "user", "content": "next"}]) + ) + assert server.engine.prompts[-1][: len(prior)] == prior + assert ("#" + action + "~\nU").encode() in bytes(server.engine.prompts[-1]) + + asyncio.run(run()) + + +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("budget_field", ["max_tokens", "max_completion_tokens"]) +@pytest.mark.parametrize("truncate", [None, 11]) +def test_completed_observation_has_no_second_generation_budget_or_truncation( + serving, stream, budget_field, truncate +): + server, _ = serving + render = server.render_chat_request + completed = [] + + async def validated_render(request): + result = await render(request) + prompt = result[1][0] + if request.truncate_prompt_tokens is not None: + prompt["prompt_token_ids"] = prompt["prompt_token_ids"][ + -request.truncate_prompt_tokens : + ] + output_budget = request.max_completion_tokens or request.max_tokens or 0 + assert len(prompt["prompt_token_ids"]) + output_budget <= 48 + if not request.add_generation_prompt: + completed.append(bytes(prompt["prompt_token_ids"])) + return result + + server.render_chat_request = validated_render + request = Request( + messages=[{"role": "user", "content": "question"}], + stream=stream, + truncate_prompt_tokens=truncate, + max_tokens=32 if budget_field == "max_tokens" else None, + max_completion_tokens=32 if budget_field == "max_completion_tokens" else None, + ) + + async def run(): + response = await server.create_chat_completion(request) + if stream: + assert [event async for event in response][-1] == "data: [DONE]\n\n" + else: + assert response.choices[0].message.content == "action" + + asyncio.run(run()) + assert b"Uquestion;Athought#action~\n" in completed + assert getattr(request, budget_field) == 32 + assert request.truncate_prompt_tokens == truncate + + +def test_external_observations_use_the_host_store_without_local_publication(serving): + server, modules = serving + observations = [] + + async def observe(request, entries): + observations.extend(entries) + + vllm.set_external_history_observer(observe, modules.__getitem__) + + async def run(): + await server.create_chat_completion( + Request(messages=[{"role": "user", "content": "question"}]), + SimpleNamespace(headers={"x-caladan-prefix-scope": "scope"}), + ) + + asyncio.run(run()) + assert observations + assert not vllm._PREFIXES._scopes + + +def test_previous_responses_keep_reasoning_and_tool_calls(serving): + _, modules = serving + previous = [{"type": "reasoning"}, {"type": "function_call"}] + utils = modules["vllm.entrypoints.openai.responses.utils"] + assert utils.construct_input_messages( + prev_response_output=previous, request_input="next" + )["request_input"] == [*previous, {"role": "user", "content": "next"}] + + +def test_batched_generations_keep_separate_prompt_histories(serving): + server, _ = serving + turn = vllm._Turn("scope", server.renderer.tokenizer) + + async def run(): + token = vllm._CURRENT.set(turn) + try: + first = server.engine.generate({"prompt_token_ids": list(b"first")}) + second = server.engine.generate({"prompt_token_ids": list(b"second")}) + await anext(first) + await anext(second) + assert [item async for item in first] == [] + assert [item async for item in second] == [] + finally: + vllm._CURRENT.reset(token) + + asyncio.run(run()) + for prompt in (b"first", b"second"): + tokens = list(prompt + b"\nthought\n#action~") + match = vllm._PREFIXES.lookup("scope", tokens, None) + assert match is not None + assert list(match.raw_prefix) == tokens + + +def test_vllm_multimodal_offsets_follow_preserved_text(): + from art_inference.token_prefix import PrefixEdit + + span = SimpleNamespace(offset=4, length=2) + original = { + "prompt_token_ids": [1, 2, 3, 4, 9, 9], + "mm_placeholders": {"image": [span]}, + "token_type_ids": [0, 0, 0, 0, 1, 1], + } + edited = vllm.replace_prompt_tokens( + original, [1, 2, 2, 3, 4, 9, 9], [PrefixEdit(1, 2, (2, 2))] + ) + assert edited["mm_placeholders"]["image"][0].offset == 5 + assert span.offset == 4 + assert edited["token_type_ids"] == [0, 0, 0, 0, 0, 1, 1] + with pytest.raises(ValueError, match="multimodal placeholder"): + vllm.replace_prompt_tokens(original, [1], [PrefixEdit(4, 6, ())]) + + +def test_native_deepseek_v32_respects_preservation_options(): + encoding = SimpleNamespace( + encode_messages=lambda **kwargs: kwargs, + render_message=lambda index, messages, **kwargs: messages[index], + ) + + class Tokenizer: + def apply_chat_template(self, messages, **kwargs): + return encoding.encode_messages( + drop_thinking=messages[-1]["role"] == "user" + ) + + Tokenizer.apply_chat_template.__module__ = "vllm.tokenizers.deepseek_v32" + tokenizer = Tokenizer() + vllm._configure_native_tokenizer(tokenizer, lambda _: encoding) + messages = [{"role": "user"}] + assert tokenizer.apply_chat_template(messages)["drop_thinking"] is False + assert ( + tokenizer.apply_chat_template(messages, drop_thinking=True)["drop_thinking"] + is True + ) + assert ( + tokenizer.apply_chat_template(messages, preserve_thinking=False)[ + "drop_thinking" + ] + is True + ) + assert tokenizer.apply_chat_template(messages)["drop_thinking"] is False + + +@pytest.mark.parametrize("legacy", [False, True]) +@pytest.mark.parametrize("stream", [False, True]) +def test_responses_use_the_native_renderer_and_preserve_ids(serving, legacy, stream): + server, _ = serving + if legacy: + server.openai_serving_render = server.online_renderer + del server.online_renderer + + async def run(): + request = Request( + messages=[{"role": "user", "content": "question"}], stream=stream + ) + response = await server.create_responses(request) + if stream: + assert (await anext(response)).type == "response.completed" + await response.aclose() + previous = Request( + messages=[ + *request.messages, + Message(reasoning_content="\nthought\n", content="action").model_dump(), + {"role": "user", "content": "next"}, + ] + ) + await server.create_responses(previous) + assert bytes(server.engine.prompts[-1]).startswith( + b"Uquestion;A\nthought\n#action~" + ) + + asyncio.run(run()) diff --git a/tests/unit/test_local_sft.py b/tests/unit/test_local_sft.py index f9c423b2f..7b27eb38d 100644 --- a/tests/unit/test_local_sft.py +++ b/tests/unit/test_local_sft.py @@ -38,7 +38,7 @@ def test_qwen_rollout_server_uses_preserve_thinking_template( lambda *_args: None, ) monkeypatch.setattr( - "art.local.backend.AutoTokenizer.from_pretrained", + "art.local.backend.get_tokenizer", lambda _model: tokenizer, ) config: dict[str, Any] = {} @@ -49,24 +49,21 @@ def test_qwen_rollout_server_uses_preserve_thinking_template( assert "preserve_thinking is defined and preserve_thinking is true" in configured -def test_non_qwen_rollout_server_does_not_load_template( - monkeypatch: pytest.MonkeyPatch, -) -> None: +def test_all_model_families_use_the_default_tokenizer(monkeypatch): monkeypatch.setattr( - "art.local.backend._model_support_default_chat_template", - lambda *_args: None, + "art.local.backend._model_support_default_chat_template", lambda *_: None ) - monkeypatch.setattr( - "art.local.backend.AutoTokenizer.from_pretrained", - lambda _model: pytest.fail("non-Qwen templates should not be loaded eagerly"), - ) - config: dict[str, Any] = {} + calls = [] - _apply_configured_chat_template_server_args( - config, {}, base_model="meta-llama/Llama-3.1-8B-Instruct" - ) + def load(model): + calls.append(model) + return type("Tokenizer", (), {"chat_template": "{{ messages }}"})() - assert config == {} + monkeypatch.setattr("art.local.backend.get_tokenizer", load) + config = {} + _apply_configured_chat_template_server_args(config, {}, base_model="other/family") + assert calls == ["other/family"] + assert config["server_args"]["chat_template"] == "{{ messages }}" def test_explicit_server_template_avoids_tokenizer_loading( @@ -77,7 +74,7 @@ def test_explicit_server_template_avoids_tokenizer_loading( lambda *_args: pytest.fail("explicit template should win before model support"), ) monkeypatch.setattr( - "art.local.backend.AutoTokenizer.from_pretrained", + "art.local.backend.get_tokenizer", lambda _model: pytest.fail("explicit template should avoid hub access"), ) config: dict[str, Any] = {"server_args": {"chat_template": "explicit"}} @@ -97,7 +94,7 @@ def test_qwen_template_load_failure_is_a_warned_fallback( lambda *_args: None, ) monkeypatch.setattr( - "art.local.backend.AutoTokenizer.from_pretrained", + "art.local.backend.get_tokenizer", lambda _model: (_ for _ in ()).throw(ValueError("bad tokenizer")), ) config: dict[str, Any] = {} @@ -119,7 +116,7 @@ def _local_sft_patches( with ExitStack() as stack: for patcher in ( patch( - "art.local.backend.AutoTokenizer.from_pretrained", + "art.local.backend.get_tokenizer", return_value=object(), ), patch.object( @@ -257,3 +254,22 @@ async def train_sft( assert [call["learning_rate"] for call in calls] == [1e-4, 1e-4] assert [batch.learning_rate for batch in captured_batches] == [1e-4] assert results[0]["data/step_num_dropped_trajectories"] == 1.0 + + +def test_python_encoder_marker_is_not_passed_as_a_jinja_template(monkeypatch): + monkeypatch.setattr( + "art.local.backend._model_support_default_chat_template", lambda *_: None + ) + monkeypatch.setattr( + "art.local.backend.get_tokenizer", + lambda _: type( + "Tokenizer", + (), + {"chat_template": "deepseek_v4_python_encoder enable_thinking"}, + )(), + ) + config = {} + _apply_configured_chat_template_server_args( + config, {}, base_model="deepseek-ai/DeepSeek-V4-Flash" + ) + assert "chat_template" not in config.get("server_args", {}) diff --git a/tests/unit/test_pipeline_trainer_local_backend.py b/tests/unit/test_pipeline_trainer_local_backend.py index c59141eca..35b1753ec 100644 --- a/tests/unit/test_pipeline_trainer_local_backend.py +++ b/tests/unit/test_pipeline_trainer_local_backend.py @@ -746,7 +746,7 @@ def test_local_backend_get_packed_tensors_warns_and_drops_overlong_results( with ( patch( - "art.local.backend.AutoTokenizer.from_pretrained", + "art.local.backend.get_tokenizer", return_value=short_result._tokenizer, ), patch("transformers.AutoImageProcessor.from_pretrained", return_value=None), diff --git a/tests/unit/test_prefix_cache.py b/tests/unit/test_prefix_cache.py deleted file mode 100644 index 9c4b0ba65..000000000 --- a/tests/unit/test_prefix_cache.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Tests for the LRUTrieCache prefix rewrite helper.""" - -import pytest - -pytest.importorskip("datrie") - -from art.tinker.prefix_cache import LRUTrieCache - - -class TestLRUTrieCache: - def test_longest_prefix_match(self) -> None: - cache = LRUTrieCache(max_entries=10) - cache.insert([1, 2], [10, 11]) - cache.insert([1, 2, 3], [20, 21, 22]) - - entry = cache.lookup([1, 2, 3, 4]) - - assert entry is not None - assert entry.rendered_len == 3 - assert entry.raw_prefix == (20, 21, 22) - - def test_lru_eviction(self) -> None: - cache = LRUTrieCache(max_entries=2) - cache.insert([1], [10]) - cache.insert([2], [20]) - - assert cache.lookup([1, 99]) is not None - - cache.insert([3], [30]) - - assert cache.lookup([2, 0]) is None - assert cache.lookup([1, 0]) is not None - - def test_invalid_size(self) -> None: - with pytest.raises(ValueError): - LRUTrieCache(max_entries=0) diff --git a/tests/unit/test_preprocessing_tokenize.py b/tests/unit/test_preprocessing_tokenize.py index 40bc11349..1720bdffc 100644 --- a/tests/unit/test_preprocessing_tokenize.py +++ b/tests/unit/test_preprocessing_tokenize.py @@ -441,7 +441,8 @@ def test_native_or_unrecognized_thinking_templates_are_unchanged() -> None: ) unrelated = "{%- if loop.index0 > ns.last_query_index %}content{% endif %}" - assert chat_template_with_preserved_thinking(native) == native + configured = chat_template_with_preserved_thinking(native) + assert chat_template_with_preserved_thinking(configured) == configured assert chat_template_with_preserved_thinking(unrelated) == unrelated diff --git a/tests/unit/test_sglang_history.py b/tests/unit/test_sglang_history.py new file mode 100644 index 000000000..391a2d80d --- /dev/null +++ b/tests/unit/test_sglang_history.py @@ -0,0 +1,339 @@ +import asyncio +import json +from types import SimpleNamespace + +from pydantic import BaseModel +import pytest + +from art_inference import sglang +from art_inference.token_prefix import TokenPrefixStore + + +class Request(BaseModel): + model: str = "model" + messages: list[dict] = [] + input: str | list = "next" + chat_template_kwargs: dict = {} + add_generation_prompt: bool = True + skip_special_tokens: bool = True + previous_response_id: str | None = None + stream: bool = False + continue_final_message: bool = False + + +class Response(BaseModel): + id: str = "previous" + output: list = [] + + +@pytest.fixture +def serving(): + class Tokenizer: + def encode(self, text, **kwargs): + return list(text.encode()) + + def decode(self, tokens, **kwargs): + return bytes(tokens).decode() + + def apply_chat_template( + self, messages, *, add_generation_prompt=True, **kwargs + ): + text = "".join( + "A" + + m.get("reasoning_content", "").strip() + + "#" + + m.get("content", "") + + "~\n" + if m["role"] == "assistant" + else "U" + m["content"] + ";" + for m in messages + ) + return self.encode(text + ("A" if add_generation_prompt else "")) + + encoding = SimpleNamespace( + encode_messages=lambda **kwargs: kwargs, + render_message=lambda index, messages, thinking_mode: messages[index], + ) + + class Serving: + def __init__(self): + self.tokenizer_manager = SimpleNamespace( + tokenizer=Tokenizer(), model_config=SimpleNamespace(is_multimodal=False) + ) + self.msg_store = {} + self.response_store = {} + + def _process_messages(self, request, is_multimodal): + request.skip_special_tokens = False + self.options = request.chat_template_kwargs + self.dsv4_options = encoding.encode_messages() + messages, _ = self._handle_last_assistant_message( + [dict(message) for message in request.messages], request + ) + return SimpleNamespace( + prompt_ids=self.tokenizer_manager.tokenizer.apply_chat_template( + messages + ) + ) + + def _handle_last_assistant_message(self, messages, request): + # SGLang 0.5.15 request rendering rewrites a trailing assistant. + if messages and messages[-1]["role"] == "assistant": + if request.continue_final_message: + return messages[:-1], messages[-1]["content"] + messages[-1] = {"role": "user", "content": messages[-1]["content"]} + return messages, None + + async def _handle_non_streaming_request( + self, adapted_request, request, raw_request + ): + prompt = self._process_messages(request, False).prompt_ids + choice = { + "index": 0, + "message": { + "role": "assistant", + "reasoning_content": "\nthought\n", + "content": "action", + }, + "token_ids": list(b"\nthought\n#action~"), + "finish_reason": "stop", + } + return SimpleNamespace( + body=json.dumps( + {"prompt_token_ids": prompt, "choices": [choice]} + ).encode() + ) + + async def _generate_chat_stream(self, adapted_request, request, raw_request): + result = json.loads( + ( + await self._handle_non_streaming_request( + adapted_request, request, raw_request + ) + ).body + ) + choice = result["choices"][0] + choice["delta"] = choice.pop("message") + yield "data: " + json.dumps(result) + "\n\n" + yield "data: [DONE]\n\n" + + async def create_responses(self, request, raw_request=None): + prompt = self._process_messages( + Request(messages=request.input), False + ).prompt_ids + payload = { + "status": "completed", + "output": [ + {"type": "reasoning", "content": [{"text": "\nthought\n"}]}, + {"role": "assistant", "content": "action"}, + ], + "token_generations": [ + { + "prompt_token_ids": prompt, + "output_tokens": [ + {"token_id": token} for token in b"\nthought\n#action~" + ], + } + ], + } + if not request.stream: + return SimpleNamespace(body=json.dumps(payload).encode()) + + async def events(): + yield ( + "event: response.completed\ndata: " + + json.dumps({"response": payload}) + + "\n\n" + ) + + return events() + + @staticmethod + def _merge_consecutive_assistant_messages(messages): + combined = {} + for message in messages: + combined.update(message) + return [combined] + + def _response_tools_to_chat_tools(self, request): + return [] + + @classmethod + def _normalize_response_message_for_chat(cls, message): + return message + + def _construct_input_messages(self, request, prev_response=None): + return ( + [*prev_response.output, *request.input] + if prev_response + else request.input + ) + + def _construct_input_messages_with_harmony(self, request, prev_response): + previous = self.msg_store[prev_response.id] + previous[:] = [m for m in previous if m.channel != "analysis"] + return [*previous, request.input] + + entries = [] + + async def observe(request, observations): + entries.extend(observations) + + modules = { + "sglang.srt.entrypoints.openai.serving_chat": SimpleNamespace( + OpenAIServingChat=Serving + ), + "sglang.srt.entrypoints.openai.serving_responses": SimpleNamespace( + OpenAIServingResponses=Serving + ), + "sglang.srt.entrypoints.openai.encoding_dsv4": encoding, + "sglang.srt.entrypoints.openai.protocol": SimpleNamespace( + ChatCompletionRequest=Request + ), + "sglang.srt.entrypoints.harmony_utils": SimpleNamespace( + render_for_completion=lambda messages: messages + ), + "sglang.srt.entrypoints.openai.encoding_dsv32": SimpleNamespace( + encode_messages=lambda **kwargs: kwargs, + render_message=lambda index, messages, **kwargs: messages[index], + ), + } + sglang.patch_history(observe, modules.__getitem__) + return Serving(), entries + + +@pytest.mark.parametrize("stream", [False, True]) +def test_sglang_observes_normalized_history_and_edited_actions(serving, stream): + server, entries = serving + request = Request(messages=[{"role": "user", "content": "question"}]) + + async def run(): + if stream: + async for _ in server._generate_chat_stream(None, request, None): + pass + else: + await server._handle_non_streaming_request(None, request, None) + + asyncio.run(run()) + store = TokenPrefixStore() + for rendered, raw, edits in entries: + store.insert("scope", rendered, raw, "lineage", edits) + prompt = list(b"Uquestion;Athought#edited~\nUnext;A") + match = store.lookup("scope", prompt, "lineage") + assert match is not None + actual = list(match.raw_prefix) + prompt[match.rendered_length :] + assert bytes(actual) == b"Uquestion;A\nthought\n#edited~\nUnext;A" + assert server.options == {} + assert request.skip_special_tokens is False + assert server.dsv4_options["drop_thinking"] is False + + +@pytest.mark.parametrize("preserve", [False, True]) +def test_harmony_opt_out_does_not_mutate_stored_history(serving, preserve): + server, _ = serving + thought, final = ( + SimpleNamespace(channel="analysis"), + SimpleNamespace(channel="final"), + ) + server.msg_store["previous"] = [thought, final] + result = server._construct_input_messages_with_harmony( + Request(chat_template_kwargs={"preserve_thinking": preserve}), Response() + ) + assert result == ([thought, final, "next"] if preserve else [final, "next"]) + assert server.msg_store["previous"] == [thought, final] + + +def test_responses_use_full_reasoning_and_retain_output_items(serving): + server, _ = serving + assert server._normalize_response_message_for_chat( + { + "type": "reasoning", + "content": [{"text": "full "}, {"text": "reasoning"}], + "summary": [{"text": "short summary"}], + } + ) == {"role": "assistant", "reasoning_content": "full reasoning"} + items = [{"type": "reasoning"}, {"type": "function_call"}] + assert server._construct_input_messages(Request(), Response(output=items)) == [ + *items, + {"role": "user", "content": "next"}, + ] + + +def test_explicit_dsv4_opt_out_is_scoped_to_one_request(serving): + server, _ = serving + server._process_messages( + Request(chat_template_kwargs={"drop_thinking": True}), False + ) + assert server.dsv4_options["drop_thinking"] is True + server._process_messages(Request(), False) + assert server.dsv4_options["drop_thinking"] is False + + +@pytest.mark.parametrize("stream", [False, True]) +def test_responses_history_opt_out_reaches_nested_chat_rendering(serving, stream): + server, entries = serving + + async def run(): + result = await server.create_responses( + Request( + input=[{"role": "user", "content": "question"}], + chat_template_kwargs={"preserve_thinking": False}, + stream=stream, + ) + ) + if stream: + async for _ in result: + pass + + asyncio.run(run()) + assert server.dsv4_options["drop_thinking"] is True + assert not entries + server._process_messages(Request(), False) + assert server.dsv4_options["drop_thinking"] is False + + +def test_multimodal_offsets_and_rope_positions_follow_text_edits(): + import torch + + from art_inference.token_prefix import PrefixEdit + + mm = SimpleNamespace( + mm_items=[SimpleNamespace(offsets=[(4, 5)])], + input_ids=[1, 2, 3, 4, 9, 9], + padded_input_ids=[1, 2, 3, 4, -1, -1], + mrope_positions=torch.tensor( + [[0, 1, 2, 3, 4, 4], [0, 1, 2, 3, 4, 5], [0, 1, 2, 3, 4, 4]] + ), + mrope_position_delta=torch.tensor([0]), + token_type_ids=torch.tensor([0, 0, 0, 0, 1, 1]), + ) + tokenized = SimpleNamespace(mm_inputs=mm) + sglang.replace_prompt_tokens(tokenized, [PrefixEdit(1, 2, (2, 2))]) + actual = tokenized.mm_inputs + assert actual.mm_items[0].offsets == [(5, 6)] + assert mm.mm_items[0].offsets == [(4, 5)] + assert actual.padded_input_ids == [1, 2, 2, 3, 4, -1, -1] + assert actual.mrope_positions.tolist() == [ + [0, 1, 2, 3, 4, 5, 5], + [0, 1, 2, 3, 4, 5, 6], + [0, 1, 2, 3, 4, 5, 5], + ] + assert actual.token_type_ids.tolist() == [0, 0, 0, 0, 0, 1, 1] + assert actual.mrope_position_delta is mm.mrope_position_delta + + +@pytest.mark.parametrize("stream", [False, True]) +def test_responses_observe_the_complete_native_generation(serving, stream): + server, entries = serving + + async def run(): + result = await server.create_responses( + Request(input=[{"role": "user", "content": "question"}], stream=stream) + ) + if stream: + assert (await anext(result)).startswith("event: response.completed") + await result.aclose() + + asyncio.run(run()) + assert entries + assert bytes(entries[-1][1]) == b"Uquestion;A\nthought\n#action~" diff --git a/tests/unit/test_tinker_renderers.py b/tests/unit/test_tinker_renderers.py index b2f641a70..81e2bed8f 100644 --- a/tests/unit/test_tinker_renderers.py +++ b/tests/unit/test_tinker_renderers.py @@ -1,6 +1,8 @@ +import asyncio import json from pathlib import Path -from typing import cast +from types import SimpleNamespace +from typing import Any, cast from tinker import EncodedTextChunk, ModelInput from tinker_cookbook import renderers @@ -75,6 +77,32 @@ def test_get_renderer_name_autodetects_qwen3_5() -> None: assert get_renderer_name("Qwen/Qwen3.5-35B-A3B") == "qwen3_5_disable_thinking" +def test_tinker_normalizes_tool_arguments_for_mapping_templates(): + from art.tinker.server import OpenAICompatibleTinkerServerWorker + + class Tokenizer: + chat_template = ( + "{% for k, v in tool_call.function.arguments.items() %}{{ k }}{% endfor %}" + ) + + def apply_chat_template(self, messages, **kwargs): + assert messages[0]["tool_calls"][0]["function"]["arguments"] == {"x": 1} + return [1, 2] + + worker = OpenAICompatibleTinkerServerWorker( + _renderers={"model": cast(Any, SimpleNamespace(tokenizer=Tokenizer()))} + ) + message = { + "role": "assistant", + "tool_calls": [{"function": {"name": "lookup", "arguments": '{"x": 1}'}}], + } + assert asyncio.run(worker.prompt_tokens("model", cast(Any, [message]), None)) == [ + 1, + 2, + ] + assert message["tool_calls"][0]["function"]["arguments"] == '{"x": 1}' + + def test_qwen3_5_generation_prompt_matches_hf_suffixes() -> None: tokenizer = FakeTokenizer() diff --git a/tests/unit/test_token_prefix.py b/tests/unit/test_token_prefix.py new file mode 100644 index 000000000..c4457dbcf --- /dev/null +++ b/tests/unit/test_token_prefix.py @@ -0,0 +1,1074 @@ +import asyncio +from collections import defaultdict +from collections.abc import Awaitable, Callable +from dataclasses import replace +import json +import random + +import pytest + +from art.token_prefix import ( + COMPACT_PREFIX_VERSION, + CompactPrefixStore, + PrefixEdit, + PrefixMatch, + TokenPrefixCache, + TokenPrefixRuntime, + TokenPrefixStore, + apply_prefix_edits, + compact_candidate_from_payload, + compact_candidate_payload, + compact_candidates_from_header, + compact_candidates_header, + compact_prefix_candidate, + prefix_edits, + resolve_compact_prefix, + token_digest, + token_prefix_digests, +) + +_SharedEntry = tuple[str, list[int], list[int], str, tuple[PrefixEdit, ...]] + + +def _compact_entry( + rendered: list[int], + raw: list[int], + *, + scope: str = "a" * 64, + lineage: str = "rollout", +) -> dict[str, object]: + return { + "scope": scope, + "lineage": lineage, + **compact_candidate_payload(compact_prefix_candidate(rendered, raw)), + } + + +def _batch( + insert: Callable[[str, list[int], list[int], str], Awaitable[object]], +) -> Callable[[list[_SharedEntry]], Awaitable[None]]: + async def insert_many(entries: list[_SharedEntry]) -> None: + for scope, rendered, raw, lineage, _ in entries: + await insert(scope, rendered, raw, lineage) + + return insert_many + + +async def _discard(_entries: list[_SharedEntry]) -> None: + pass + + +async def _runtime_insert( + runtime: TokenPrefixRuntime, + scope: str, + rendered: list[int], + raw: list[int], + lineage: str = "rollout", +) -> None: + await runtime.insert_many( + [(scope, rendered, raw, lineage, prefix_edits(rendered, raw))] + ) + + +def test_store_retains_only_a_complete_bounded_lineage_snapshot() -> None: + store = CompactPrefixStore(max_serialized_bytes_per_lineage=512) + candidates = [compact_prefix_candidate([index], [index + 1]) for index in range(4)] + for candidate in candidates: + store.insert("a" * 64, "rollout", candidate) + + retained = store.candidates("a" * 64, "rollout", 4) + payload = { + "version": COMPACT_PREFIX_VERSION, + "candidates": [compact_candidate_payload(value) for value in retained], + } + assert len(json.dumps(payload, separators=(",", ":")).encode()) <= 512 + assert 0 < len(retained) < len(candidates) + assert retained[0] == candidates[-1] + + +def test_default_store_retains_busy_lineage_turns() -> None: + store = CompactPrefixStore() + candidates = [ + compact_prefix_candidate([index], [index + 1]) for index in range(240) + ] + for candidate in candidates: + assert store.insert("a" * 64, "shared-lineage", candidate) + + retained = store.candidates("a" * 64, "shared-lineage", 240) + assert len(retained) == len(candidates) + assert candidates[0] in retained + + +def test_complete_candidate_headers_round_trip_or_fall_back() -> None: + candidate = compact_prefix_candidate([1, 500], [1, 101, 102]) + + encoded = compact_candidates_header([candidate]) + empty = compact_candidates_header([]) + dense = compact_candidates_header( + [compact_prefix_candidate([index], [index]) for index in range(65)] + ) + oversized = compact_candidates_header( + [ + compact_prefix_candidate( + [index], + [index, *range(1_000 + index * 1_000, 2_000 + index * 1_000)], + ) + for index in range(32) + ] + ) + + assert encoded is not None + assert compact_candidates_from_header(encoded) == (candidate,) + assert empty is not None + assert compact_candidates_from_header(empty) == () + assert dense is not None + assert len(compact_candidates_from_header(dense)) == 65 + + assert oversized is None + + +def test_candidate_header_rejects_malformed_or_partial_snapshots() -> None: + for value in ( + "not-json", + json.dumps({"version": "old", "candidates": []}), + json.dumps({"version": COMPACT_PREFIX_VERSION}), + ): + try: + compact_candidates_from_header(value) + except ValueError: + pass + else: + raise AssertionError("invalid candidate header was accepted") + + +def test_compact_edits_cover_deletions_expansions_and_sparse_128k_changes() -> None: + cases = ( + ([9, 1, 2], [1, 2]), + ([1, 2], [9, 1, 2]), + ([1, 2, 3], [1, 8, 9, 3]), + ) + for rendered, raw in cases: + edits = prefix_edits(rendered, raw) + assert apply_prefix_edits(rendered, edits) == raw + candidate = compact_prefix_candidate(rendered, raw, edits) + assert resolve_compact_prefix([*rendered, 77], [candidate]) == PrefixMatch( + len(rendered), tuple(raw) + ) + + rendered = list(range(128_000)) + raw = rendered.copy() + raw[10] = 200_010 + raw[64_000] = 264_000 + raw[-2] = 327_998 + edits = prefix_edits(rendered, raw) + payload = compact_candidate_payload(compact_prefix_candidate(rendered, raw, edits)) + + assert len(edits) == 3 + assert len(json.dumps(payload, separators=(",", ":")).encode()) < 512 + assert resolve_compact_prefix( + [*rendered, 9], [compact_candidate_from_payload(payload)] + ) == PrefixMatch(len(rendered), tuple(raw)) + + +def test_compact_resolution_is_longest_newest_and_corruption_tolerant() -> None: + shorter = compact_prefix_candidate([1, 2], [7, 8]) + longer = compact_prefix_candidate([1, 2, 3], [7, 8, 9]) + corrupt = replace(longer, raw_digest="0" * 64) + + assert resolve_compact_prefix( + [1, 2, 3, 4], [shorter, corrupt, longer] + ) == PrefixMatch(3, (7, 8, 9)) + + first = compact_prefix_candidate([1, 2], [10, 11]) + newest = compact_prefix_candidate([1, 2], [20, 21]) + assert resolve_compact_prefix([1, 2, 3], [newest, first]) == PrefixMatch( + 2, (20, 21) + ) + + +def test_compact_identity_candidate_shadows_a_shorter_divergence() -> None: + store = CompactPrefixStore() + store.insert( + "scope", + "lineage", + compact_prefix_candidate([1, 2], [7, 8]), + ) + store.insert( + "scope", + "lineage", + compact_prefix_candidate([1, 2, 3], [1, 2, 3]), + ) + + candidates = store.candidates("scope", "lineage", 4) + assert candidates[0].edits == () + assert resolve_compact_prefix([1, 2, 3, 4], candidates) == PrefixMatch(3, (1, 2, 3)) + + +def test_compact_payload_rejects_invalid_token_and_edit_bounds() -> None: + payload = compact_candidate_payload(compact_prefix_candidate([1], [1])) + for edit in ( + [[0, 1, [True]]], + [[0, 1, [-1]]], + [[0, 1, [0x1_0000_0000]]], + [[1, 1, []], [0, 1, []]], + ): + invalid = payload | {"edits": edit} + try: + compact_candidate_from_payload(invalid) + except ValueError: + pass + else: + raise AssertionError(f"accepted invalid edit: {edit}") + + try: + compact_candidate_from_payload(payload | {"rendered_length": 262_145}) + except ValueError: + pass + else: + raise AssertionError("accepted an unreachable rendered length") + + +def test_compact_prefix_hashes_requested_lengths_consistently() -> None: + values = list(range(128_000)) + lengths = [128_000, 32, 64_000, 32] + digests = token_prefix_digests(values, lengths) + + assert list(digests) == [32, 64_000, 128_000] + for length, digest in digests.items(): + assert ( + digest + == compact_prefix_candidate( + values[:length], values[:length] + ).rendered_digest + ) + + +def test_compact_digest_encoding_has_a_stable_golden_vector() -> None: + assert token_digest([]) == ( + "ff127bd383d2204e4e5509b3c4001b75988e0256b4116ffb904e0864848b04c1" + ) + assert token_digest([0, 1, 0xFFFFFFFF]) == ( + "4cc20a06621a8b2cec34aef4b495b6439931b0faa260e9308a7298cfe9e862b5" + ) + + +def test_compact_candidates_match_a_randomized_reference() -> None: + randomizer = random.Random(1) + for _ in range(1_000): + rendered = [ + randomizer.randrange(32) for _ in range(randomizer.randrange(1, 65)) + ] + raw = rendered.copy() + start = randomizer.randrange(len(raw) + 1) + stop = randomizer.randrange(start, len(raw) + 1) + raw[start:stop] = [ + randomizer.randrange(32) for _ in range(randomizer.randrange(0, 12)) + ] + suffix = [randomizer.randrange(32) for _ in range(randomizer.randrange(0, 12))] + candidate = compact_prefix_candidate(rendered, raw) + + assert resolve_compact_prefix([*rendered, *suffix], [candidate]) == PrefixMatch( + len(rendered), tuple(raw) + ) + + +def test_longest_observed_prefix_rewrites_only_rendered_prefix() -> None: + cache = TokenPrefixCache() + cache.insert([1, 500], [1, 101, 102], "rollout") + cache.insert([1, 500, 3, 4], [1, 101, 102, 3, 4], "rollout") + + assert cache.lookup([1, 500, 3, 4, 5], "rollout") == PrefixMatch( + rendered_length=4, + raw_prefix=(1, 101, 102, 3, 4), + ) + + +def test_cache_matches_a_reference_model_across_random_prefixes() -> None: + randomizer = random.Random(0) + cache = TokenPrefixCache( + max_variants=10_000, + max_token_ids=1_000_000, + max_lineages_per_variant=100, + ) + reference: dict[tuple[int, ...], dict[tuple[int, ...], dict[str, int]]] = ( + defaultdict(lambda: defaultdict(dict)) + ) + + for observation in range(2_000): + rendered = tuple( + randomizer.randrange(8) for _ in range(randomizer.randrange(1, 12)) + ) + raw = tuple(randomizer.randrange(8) for _ in range(randomizer.randrange(1, 12))) + lineage = f"lineage-{randomizer.randrange(12)}" + cache.insert(rendered, raw, lineage) + reference[rendered][raw][lineage] = observation + + query = tuple( + randomizer.randrange(8) for _ in range(randomizer.randrange(1, 16)) + ) + query_lineage = f"lineage-{randomizer.randrange(12)}" + expected = None + candidates: list[tuple[int, ...]] = sorted( + ( + prefix + for prefix in reference + if len(prefix) <= len(query) and query[: len(prefix)] == prefix + ), + key=lambda prefix: len(prefix), + reverse=True, + ) + for prefix in candidates: + variants = reference[prefix] + matching = [ + (lineages[query_lineage], variant) + for variant, lineages in variants.items() + if query_lineage in lineages + ] + if not matching: + continue + _, selected = max(matching) + expected = PrefixMatch(len(prefix), selected) + break + + assert cache.lookup(query, query_lineage) == expected + + +def test_runtime_preserves_the_entire_prompt_on_a_miss() -> None: + async def exercise() -> None: + runtime = TokenPrefixRuntime( + lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep( + 0, result=None + ), + insert_shared_many=_discard, + ) + prompt = [1, 2, 3, 4] + + canonical, engine_input, _ = await runtime.rewrite_with_edits( + "model", prompt, "rollout", shared_candidate=True + ) + + assert canonical == prompt + assert engine_input == prompt + assert canonical is not prompt + assert engine_input is not canonical + + asyncio.run(exercise()) + + +def test_runtime_changes_exactly_the_matched_prefix() -> None: + async def exercise() -> None: + shared = TokenPrefixStore() + + async def lookup(scope, rendered, lineage): + return shared.lookup(scope, rendered, lineage) + + async def insert(scope, rendered, raw, lineage): + shared.insert(scope, rendered, raw, lineage) + + runtime = TokenPrefixRuntime( + lookup_shared=lookup, insert_shared_many=_batch(insert) + ) + cases = ( + ([500], [101, 102], []), + ([1, 500], [1, 101, 102], [3]), + ([1, 2, 3, 4], [7], [5, 6, 7, 8]), + (list(range(128)), [9, 10, 11], [200, 201]), + ) + for index, (rendered_prefix, raw_prefix, suffix) in enumerate(cases): + scope = f"model-{index}" + await _runtime_insert(runtime, scope, rendered_prefix, raw_prefix) + await runtime.flush() + prompt = [*rendered_prefix, *suffix] + + canonical, engine_input, _ = await runtime.rewrite_with_edits( + scope, prompt, "rollout", shared_candidate=True + ) + + assert canonical == prompt + assert engine_input == [*raw_prefix, *suffix] + assert engine_input[len(raw_prefix) :] == suffix + + asyncio.run(exercise()) + + +def test_replication_is_async_while_local_continuations_are_immediate() -> None: + async def exercise() -> None: + started = asyncio.Event() + release = asyncio.Event() + shared = TokenPrefixStore() + + async def insert(scope, rendered, raw, lineage): + started.set() + await release.wait() + shared.insert(scope, rendered, raw, lineage) + + runtime = TokenPrefixRuntime( + lookup_shared=lambda scope, rendered, lineage: asyncio.sleep( + 0, result=shared.lookup(scope, rendered, lineage) + ), + insert_shared_many=_batch(insert), + ) + await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102]) + + assert not started.is_set() + continuation = asyncio.create_task( + runtime.rewrite_with_edits( + "model", + [1, 500, 3], + "rollout", + shared_candidate=True, + ) + ) + _, rewritten, _ = await continuation + assert rewritten == [1, 101, 102, 3] + await started.wait() + release.set() + assert await runtime.close() + + asyncio.run(exercise()) + + +def test_completed_replication_does_not_delay_a_later_continuation() -> None: + async def exercise() -> None: + shared = TokenPrefixStore() + + async def insert(scope, rendered, raw, lineage): + shared.insert(scope, rendered, raw, lineage) + + runtime = TokenPrefixRuntime( + lookup_shared=lambda scope, rendered, lineage: asyncio.sleep( + 0, result=shared.lookup(scope, rendered, lineage) + ), + insert_shared_many=_batch(insert), + ) + await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102]) + assert await runtime.flush() + assert runtime._replicator is None + + _, rewritten, _ = await asyncio.wait_for( + runtime.rewrite_with_edits( + "model", + [1, 500, 3], + "rollout", + shared_candidate=True, + ), + 0.1, + ) + assert rewritten == [1, 101, 102, 3] + await runtime.close() + + asyncio.run(exercise()) + + +def test_failed_replication_can_reuse_an_older_matching_shared_prefix() -> None: + async def exercise() -> None: + shared = TokenPrefixStore() + shared.insert("model", [1, 500], [1, 41, 42], "rollout") + + async def fail(_scope, _rendered, _raw, _lineage): + raise RuntimeError("unavailable") + + first = TokenPrefixRuntime( + lookup_shared=lambda scope, rendered, lineage: asyncio.sleep( + 0, result=shared.lookup(scope, rendered, lineage) + ), + insert_shared_many=_batch(fail), + ) + await _runtime_insert(first, "model", [1, 500], [1, 101, 102]) + assert await first.flush() + second = TokenPrefixRuntime( + lookup_shared=lambda scope, rendered, lineage: asyncio.sleep( + 0, result=shared.lookup(scope, rendered, lineage) + ), + insert_shared_many=_discard, + ) + + canonical, rewritten, _ = await second.rewrite_with_edits( + "model", [1, 500, 3], "rollout", shared_candidate=True + ) + assert canonical == [1, 500, 3] + assert rewritten == [1, 41, 42, 3] + await first.close() + await second.close() + + asyncio.run(exercise()) + + +def test_replication_queue_is_bounded_and_drops_only_shared_work() -> None: + async def exercise() -> None: + started = asyncio.Event() + release = asyncio.Event() + replicated: list[str] = [] + + async def insert(scope, _rendered, _raw, _lineage): + replicated.append(scope) + if scope == "first": + started.set() + await release.wait() + + runtime = TokenPrefixRuntime( + lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep( + 0, result=None + ), + insert_shared_many=_batch(insert), + max_pending_batches=1, + ) + await _runtime_insert(runtime, "first", [1], [11]) + await started.wait() + await _runtime_insert(runtime, "second", [2], [22]) + await _runtime_insert(runtime, "dropped", [3], [33]) + + _, rewritten, _ = await runtime.rewrite_with_edits( + "dropped", [3, 4], "rollout", shared_candidate=False + ) + assert rewritten == [33, 4] + + release.set() + assert await runtime.close() + assert replicated == ["first", "second"] + + asyncio.run(exercise()) + + +def test_shutdown_waits_for_replication_then_stops_the_worker() -> None: + async def exercise() -> None: + started = asyncio.Event() + release = asyncio.Event() + + async def insert(_scope, _rendered, _raw, _lineage): + started.set() + await release.wait() + + runtime = TokenPrefixRuntime( + lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep( + 0, result=None + ), + insert_shared_many=_batch(insert), + ) + await _runtime_insert(runtime, "model", [1], [2]) + await started.wait() + closing = asyncio.create_task(runtime.close(timeout=1)) + await asyncio.sleep(0) + assert not closing.done() + + release.set() + assert await closing + + asyncio.run(exercise()) + + +def test_replication_coalesces_and_shutdown_timeout_is_bounded() -> None: + async def exercise() -> None: + calls = [] + + async def insert_many(entries): + calls.append(entries) + + runtime = TokenPrefixRuntime( + lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep( + 0, result=None + ), + insert_shared_many=insert_many, + ) + await _runtime_insert(runtime, "first", [1], [11]) + await _runtime_insert(runtime, "second", [2], [22]) + assert await runtime.flush() + assert len(calls) == 1 + assert [entry[0] for entry in calls[0]] == ["first", "second"] + + blocked = TokenPrefixRuntime( + lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep( + 0, result=None + ), + insert_shared_many=_batch(lambda *_args: asyncio.Event().wait()), + ) + await _runtime_insert(blocked, "model", [3], [33]) + assert not await blocked.close(timeout=0.001) + + asyncio.run(exercise()) + + +def test_target_concurrency_burst_is_retained_while_replication_is_blocked() -> None: + async def exercise() -> None: + replicated = [] + started = asyncio.Event() + release = asyncio.Event() + + async def insert_many(entries): + replicated.extend(entries) + if len(replicated) == 1: + started.set() + await release.wait() + + runtime = TokenPrefixRuntime( + lookup_shared=lambda _scope, _rendered, _lineage: asyncio.sleep( + 0, result=None + ), + insert_shared_many=insert_many, + ) + await runtime.insert_many([("model-0", [1], [1], "rollout", ())]) + await started.wait() + await asyncio.gather( + *( + runtime.insert_many( + [ + ( + f"model-{index}", + [index + 1], + [index + 1], + "rollout", + (), + ) + ] + ) + for index in range(1, 32) + ) + ) + release.set() + assert await runtime.flush() + assert len(replicated) == 32 + await runtime.close() + + asyncio.run(exercise()) + + +def test_identical_observations_are_idempotent() -> None: + cache = TokenPrefixCache() + cache.insert([1, 500], [1, 101, 102], "first") + cache.insert([1, 500], [1, 101, 102], "second") + + expected = PrefixMatch(2, (1, 101, 102)) + assert cache.lookup([1, 500, 3], "first") == expected + assert cache.lookup([1, 500, 3], "second") == expected + assert cache.lookup([1, 500, 3], "unknown") is None + + +def test_eviction_removes_all_variants_at_the_rendered_prefix() -> None: + cache = TokenPrefixCache(max_variants=2, max_token_ids=12) + cache.insert([1, 500], [1, 101, 102], "shared") + cache.insert([1, 500], [1, 500], "shared") + cache.insert([9], [9], "third") + + assert cache.lookup([1, 500, 3], "shared") is None + assert cache.lookup([9], "third") == PrefixMatch(1, (9,)) + + +def test_lineage_presence_tracks_only_retained_observations() -> None: + cache = TokenPrefixCache() + cache.insert([1], [1], "shared") + cache.insert([2], [2], "shared") + + assert cache.lookup([1], "shared") == PrefixMatch(1, (1,)) + assert cache.evict_oldest() is True + assert cache.lookup([1], "shared") == PrefixMatch(1, (1,)) + assert cache.evict_oldest() is True + assert cache.lookup([1], "shared") is None + + +def test_lineage_trimming_is_variant_local() -> None: + cache = TokenPrefixCache(max_lineages_per_variant=1) + cache.insert([1, 500], [1, 101, 102], "shared") + cache.insert([1, 500], [1, 500], "shared") + cache.insert([1, 500], [1, 101, 102], "new") + + assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (1, 500)) + assert cache.lookup([1, 500, 3], "new") == PrefixMatch(2, (1, 101, 102)) + + +def test_distinct_lineages_resolve_observed_tokenization_collision() -> None: + cache = TokenPrefixCache() + cache.insert([1], [1], "first") + cache.insert([1, 500], [1, 101, 102], "first") + cache.insert([1, 500], [1, 500], "second") + + assert cache.lookup([1, 500, 3], "first") == PrefixMatch(2, (1, 101, 102)) + assert cache.lookup([1, 500, 3], "second") == PrefixMatch(2, (1, 500)) + + +def test_collision_uses_most_recent_variant_at_longest_known_prefix() -> None: + cache = TokenPrefixCache() + cache.insert([1], [9], "shared") + cache.insert([1, 500], [9, 101, 102], "shared") + cache.insert([1, 500], [9, 500], "shared") + + assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (9, 500)) + cache.insert([1, 500], [9, 101, 102], "shared") + assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (9, 101, 102)) + + +def test_collision_without_shorter_prefix_uses_most_recent_variant() -> None: + cache = TokenPrefixCache() + cache.insert([1, 500], [1, 101, 102], "shared") + cache.insert([1, 500], [1, 500], "shared") + + assert cache.lookup([1, 500, 3], "shared") == PrefixMatch(2, (1, 500)) + + +def test_store_partitions_model_keys() -> None: + store = TokenPrefixStore() + store.insert("base", [1], [2], "lineage") + store.insert("checkpoint", [1], [3], "lineage") + + assert store.lookup("base", [1, 4], "lineage") == PrefixMatch(1, (2,)) + assert store.lookup("checkpoint", [1, 4], "lineage") == PrefixMatch(1, (3,)) + + +def test_store_enforces_per_scope_and_global_token_budgets() -> None: + store = TokenPrefixStore( + max_scopes=4, + max_variants_per_scope=4, + max_token_ids_per_scope=6, + max_token_ids=8, + ) + + assert store.insert("oversized", [1, 2, 3, 4], [5, 6, 7], "lineage") is False + assert store.lookup("oversized", [1, 2, 3, 4], "lineage") is None + + assert store.insert("first", [1, 2], [3, 4], "lineage") is True + assert store.insert("second", [5, 6], [7, 8], "lineage") is True + assert store.insert("third", [9], [10], "lineage") is True + + assert store.lookup("first", [1, 2], "lineage") is None + assert store.lookup("second", [5, 6], "lineage") == PrefixMatch(2, (7, 8)) + assert store.lookup("third", [9], "lineage") == PrefixMatch(1, (10,)) + + +def test_shared_store_reuses_identity_across_engine_replicas() -> None: + async def exercise() -> None: + shared = TokenPrefixStore() + + async def lookup(scope, rendered, lineage): + return shared.lookup(scope, rendered, lineage) + + async def insert(scope, rendered, raw, lineage): + shared.insert(scope, rendered, raw, lineage) + + first = TokenPrefixRuntime( + lookup_shared=lookup, insert_shared_many=_batch(insert) + ) + second = TokenPrefixRuntime( + lookup_shared=lookup, insert_shared_many=_batch(insert) + ) + await _runtime_insert(first, "model", [1, 500], [1, 101, 102]) + await first.flush() + + canonical, rewritten, _ = await second.rewrite_with_edits( + "model", + [1, 500, 3], + "rollout", + shared_candidate=True, + ) + + assert canonical == [1, 500, 3] + assert rewritten == [1, 101, 102, 3] + + asyncio.run(exercise()) + + +def test_runtime_migrates_an_implicit_lineage_without_a_client_session() -> None: + async def exercise() -> None: + shared = TokenPrefixStore() + shared.insert("model", [1, 500], [1, 101, 102], "scenario") + lookups = [] + + async def lookup(scope, rendered, lineage): + lookups.append(lineage) + return shared.lookup(scope, rendered, lineage) + + async def insert_many(_entries): + return None + + runtime = TokenPrefixRuntime( + lookup_shared=lookup, + insert_shared_many=insert_many, + ) + canonical, rewritten, _ = await runtime.rewrite_with_edits( + "model", + [1, 500, 3], + "tool-call", + fallback_lineage="scenario", + shared_candidate=True, + ) + + assert canonical == [1, 500, 3] + assert rewritten == [1, 101, 102, 3] + assert lookups == ["tool-call", "scenario"] + _, local, _ = await runtime.rewrite_with_edits( + "model", + [1, 500, 4], + "tool-call", + shared_candidate=False, + ) + assert local == [1, 101, 102, 4] + await runtime.close() + + asyncio.run(exercise()) + + +def test_shared_store_restores_text_equivalent_noncanonical_sampled_ids() -> None: + async def exercise() -> None: + shared = TokenPrefixStore() + + async def lookup(scope, rendered, lineage): + return shared.lookup(scope, rendered, lineage) + + async def insert(scope, rendered, raw, lineage): + shared.insert(scope, rendered, raw, lineage) + + first = TokenPrefixRuntime( + lookup_shared=lookup, insert_shared_many=_batch(insert) + ) + second = TokenPrefixRuntime( + lookup_shared=lookup, insert_shared_many=_batch(insert) + ) + # Qwen3.5 decodes both forms as " roaming", but only the two-token form + # carries the sampled logprobs that ART must train against. + await _runtime_insert(first, "model", [10, 65945], [10, 897, 6267]) + await first.flush() + + canonical, rewritten, edits = await second.rewrite_with_edits( + "model", [10, 65945, 20], "rollout", shared_candidate=True + ) + + assert canonical == [10, 65945, 20] + assert rewritten == [10, 897, 6267, 20] + assert edits == (PrefixEdit(1, 2, (897, 6267)),) + + asyncio.run(exercise()) + + +def test_replication_restarts_when_an_entry_arrives_during_shutdown() -> None: + async def exercise() -> None: + second = ("model", [3, 4], [3, 5], "rollout", (PrefixEdit(1, 2, (5,)),)) + + class RacingQueue(asyncio.Queue): + empty_calls = 0 + + def empty(self) -> bool: + self.empty_calls += 1 + if self.empty_calls == 2: + self.put_nowait([second]) + return True + return super().empty() + + replicated: list[ + tuple[str, list[int], list[int], str, tuple[PrefixEdit, ...]] + ] = [] + + async def insert_many(entries): + replicated.extend(entries) + + runtime = TokenPrefixRuntime( + lookup_shared=lambda *_args: asyncio.sleep(0, result=None), + insert_shared_many=insert_many, + ) + setattr(runtime, "_pending", RacingQueue(maxsize=16)) + + await _runtime_insert(runtime, "model", [1, 2], [1, 9]) + assert await runtime.flush(timeout=1) + + assert [entry[1] for entry in replicated] == [[1, 2], [3, 4]] + assert await runtime.close() + + asyncio.run(exercise()) + + +def test_pending_replication_can_reuse_an_older_matching_shared_prefix() -> None: + async def exercise() -> None: + started = asyncio.Event() + release = asyncio.Event() + shared = TokenPrefixStore() + + async def lookup(scope, rendered, lineage): + return shared.lookup(scope, rendered, lineage) + + async def insert(scope, rendered, raw, lineage): + started.set() + await release.wait() + shared.insert(scope, rendered, raw, lineage) + + shared.insert("model", [1, 500], [1, 41, 42], "rollout") + first = TokenPrefixRuntime( + lookup_shared=lookup, insert_shared_many=_batch(insert) + ) + await _runtime_insert(first, "model", [1, 500], [1, 101, 102]) + await started.wait() + second = TokenPrefixRuntime( + lookup_shared=lookup, insert_shared_many=_batch(insert) + ) + canonical, rewritten, _ = await second.rewrite_with_edits( + "model", [1, 500, 3], "rollout", shared_candidate=True + ) + assert canonical == [1, 500, 3] + assert rewritten == [1, 41, 42, 3] + + release.set() + assert await first.close() + await second.close() + + asyncio.run(exercise()) + + +def test_shared_candidate_compares_shared_with_local_identity() -> None: + async def exercise() -> None: + shared_lookups = 0 + + async def lookup(_scope, _rendered, _lineage): + nonlocal shared_lookups + shared_lookups += 1 + return None + + runtime = TokenPrefixRuntime( + lookup_shared=lookup, + insert_shared_many=_discard, + ) + await _runtime_insert(runtime, "model", [1, 2], [1, 2]) + + canonical, rewritten, _ = await runtime.rewrite_with_edits( + "model", [1, 2, 3], "rollout", shared_candidate=True + ) + + assert canonical == rewritten == [1, 2, 3] + assert shared_lookups == 1 + + asyncio.run(exercise()) + + +def test_local_nonidentity_and_latest_ambiguity_choice_are_immediate() -> None: + async def exercise() -> None: + shared_lookups = 0 + + async def lookup(_scope, _rendered, _lineage): + nonlocal shared_lookups + shared_lookups += 1 + return None + + runtime = TokenPrefixRuntime( + lookup_shared=lookup, + insert_shared_many=_discard, + ) + await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102]) + _, first, _ = await runtime.rewrite_with_edits( + "model", [1, 500, 3], "rollout", shared_candidate=True + ) + await _runtime_insert(runtime, "model", [1, 500], [1, 500]) + _, second, _ = await runtime.rewrite_with_edits( + "model", [1, 500, 3], "rollout", shared_candidate=True + ) + + assert first == [1, 101, 102, 3] + assert second == [1, 500, 3] + assert shared_lookups == 2 + await runtime.close() + + asyncio.run(exercise()) + + +def test_known_lineage_without_a_matching_prefix_consults_shared_lookup() -> None: + async def exercise() -> None: + shared_lookups = 0 + + async def lookup(_scope, _rendered, _lineage): + nonlocal shared_lookups + shared_lookups += 1 + return None + + runtime = TokenPrefixRuntime( + lookup_shared=lookup, + insert_shared_many=_discard, + ) + await _runtime_insert(runtime, "model", [1, 2], [1, 2]) + + canonical, rewritten, _ = await runtime.rewrite_with_edits( + "model", [8, 9], "rollout", shared_candidate=True + ) + + assert canonical == rewritten == [8, 9] + assert shared_lookups == 1 + + asyncio.run(exercise()) + + +def test_broad_lineage_without_a_local_match_consults_shared_lookup() -> None: + async def exercise() -> None: + shared_lookups = 0 + + async def lookup(_scope, _rendered, _lineage): + nonlocal shared_lookups + shared_lookups += 1 + return PrefixMatch(2, (8, 10)) + + runtime = TokenPrefixRuntime( + lookup_shared=lookup, + insert_shared_many=_discard, + ) + await _runtime_insert(runtime, "model", [1, 2], [1, 2]) + + canonical, rewritten, _ = await runtime.rewrite_with_edits( + "model", + [8, 9, 11], + "rollout", + shared_candidate=True, + ) + + assert canonical == [8, 9, 11] + assert rewritten == [8, 10, 11] + assert shared_lookups == 1 + + asyncio.run(exercise()) + + +def test_longer_shared_prefix_wins_after_routing_returns_to_stale_worker() -> None: + async def exercise() -> None: + shared = PrefixMatch( + 4, + (1, 101, 102, 3, 201, 202), + prefix_edits( + [1, 500, 3, 600], + [1, 101, 102, 3, 201, 202], + ), + ) + runtime = TokenPrefixRuntime( + lookup_shared=lambda *_args: asyncio.sleep(0, result=shared), + insert_shared_many=_discard, + ) + await _runtime_insert(runtime, "model", [1, 500], [1, 101, 102]) + + canonical, rewritten, _ = await runtime.rewrite_with_edits( + "model", + [1, 500, 3, 600, 4], + "rollout", + shared_candidate=True, + ) + + assert canonical == [1, 500, 3, 600, 4] + assert rewritten == [1, 101, 102, 3, 201, 202, 4] + await runtime.close() + + asyncio.run(exercise()) + + +def test_unknown_lineage_still_consults_shared_lookup() -> None: + async def exercise() -> None: + shared_lookups = 0 + + async def lookup(_scope, _rendered, _lineage): + nonlocal shared_lookups + shared_lookups += 1 + return None + + runtime = TokenPrefixRuntime( + lookup_shared=lookup, + insert_shared_many=_discard, + ) + await _runtime_insert(runtime, "model", [1, 2], [1, 2], "known") + await runtime.rewrite_with_edits( + "model", + [8, 9], + "unknown", + shared_candidate=True, + ) + + assert shared_lookups == 1 + + asyncio.run(exercise()) diff --git a/tests/unit/test_tokenizer.py b/tests/unit/test_tokenizer.py new file mode 100644 index 000000000..a7c8f94a8 --- /dev/null +++ b/tests/unit/test_tokenizer.py @@ -0,0 +1,211 @@ +from types import SimpleNamespace + +import jinja2 +import pytest + +from art import get_tokenizer + + +def test_default_tokenizer_preserves_reasoning_and_honors_explicit_opt_out(monkeypatch): + import transformers + + calls = [] + template = ( + "{% set enable_thinking = true %}" + "{% set ns = namespace(last_query_index=2) %}" + "{% for message in messages %}" + "{%- if loop.index0 > ns.last_query_index %}{{ message.reasoning }}{% endif %}" + "{{ message.content }}{% endfor %}" + ) + + def load(model, **kwargs): + calls.append((model, kwargs)) + return SimpleNamespace(chat_template=template) + + monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", load) + tokenizer = get_tokenizer( + "model:variant", revision="revision", local_files_only=True + ) + assert calls == [ + ( + "model", + { + "revision": "revision", + "local_files_only": True, + "trust_remote_code": False, + }, + ) + ] + messages = [{"reasoning": "thought", "content": "answer"}] + assert isinstance(tokenizer.chat_template, str) + render = jinja2.Environment().from_string(tokenizer.chat_template).render + assert render(messages=messages) == "thoughtanswer" + assert render(messages=messages, preserve_thinking=False) == "answer" + tokenizer.chat_template = "changed" + assert get_tokenizer("model").chat_template != "changed" + + +def test_pinned_llama_revision_does_not_switch_repositories(monkeypatch): + import transformers + + calls = [] + monkeypatch.setattr( + transformers.AutoTokenizer, + "from_pretrained", + lambda *a, **k: calls.append((a, k)), + ) + get_tokenizer("meta-llama/Llama-3.1-8B", revision="model-specific-commit") + assert calls[0][0] == ("meta-llama/Llama-3.1-8B",) + + +@pytest.mark.parametrize("revision", [None, "pinned-commit"]) +def test_ungated_llama_fallback_is_only_for_unpinned_loads(monkeypatch, revision): + import transformers + + calls = [] + + def load(model, **kwargs): + calls.append(model) + if model.startswith("meta-llama/"): + raise OSError("gated repository") + return SimpleNamespace(chat_template="native text template") + + monkeypatch.setattr(transformers.AutoTokenizer, "from_pretrained", load) + if revision: + with pytest.raises(OSError, match="gated"): + get_tokenizer("meta-llama/Llama-3.1-8B", revision=revision) + assert len(calls) == 1 + else: + get_tokenizer("meta-llama/Llama-3.1-8B") + assert calls[-1] == "thinkingmachineslabinc/meta-llama-3-instruct-tokenizer" + + +def test_llama_base_retains_native_vocabulary_with_default_chat_template(monkeypatch): + import transformers + + native = SimpleNamespace(chat_template=None) + monkeypatch.setattr( + transformers.AutoTokenizer, + "from_pretrained", + lambda model, **_: ( + native + if model.startswith("meta-llama/") + else SimpleNamespace(chat_template="public chat template") + ), + ) + actual = get_tokenizer("meta-llama/Llama-3.1-8B") + assert actual is native + assert actual.chat_template == "public chat template" + + +@pytest.mark.parametrize( + "template", + [ + "reasoning_content and loop.index0 > ns.last_user_index", + "thinking_text and loop.index0 > ns_turn.last_user_idx and message.get('tool_calls')", + "message.thinking and not future_final_message.found", + ], +) +def test_supported_template_patches_are_idempotent(template): + from art.utils.chat_template import chat_template_with_preserved_thinking + + patched = chat_template_with_preserved_thinking(template) + assert patched != template + assert chat_template_with_preserved_thinking(patched) == patched + + +def test_glm_preserves_content_whitespace_and_remains_valid_jinja(): + from art.utils.chat_template import chat_template_with_preserved_thinking + + template = "{% if clear_thinking %}legacy{% endif %}{%- if content.strip() -%}{{ content.strip() }}{%- endif -%}" + configured = chat_template_with_preserved_thinking(template) + assert isinstance(configured, str) + assert ( + jinja2.Environment().from_string(configured).render(content="\nanswer\n") + == "\nanswer\n" + ) + assert chat_template_with_preserved_thinking(configured) == configured + + +def test_kimi_preserves_old_reasoning_when_next_turn_does_not_think(): + from art.utils.chat_template import chat_template_with_preserved_thinking + + template = """{% set ns = namespace(last_non_tool_call_assistant_msg=0) %} + {%- set hist_msgs = messages[:ns.last_non_tool_call_assistant_msg+1] -%} + {%- set suffix_msgs = messages[ns.last_non_tool_call_assistant_msg+1:] -%} + {% for message in hist_msgs %}{{ message.content }}{% endfor %} + {% for message in suffix_msgs %}{%- if thinking is defined and thinking is false -%}{{ message.content }}{% else %}{{ message.reasoning_content }}{{ message.content }}{% endif %}{% endfor %}""" + configured = chat_template_with_preserved_thinking(template) + assert isinstance(configured, str) + render = jinja2.Environment().from_string(configured).render + messages = [{"reasoning_content": "thought", "content": "answer"}] + assert "thoughtanswer" in render(messages=messages, thinking=False) + assert "thought" not in render( + messages=messages, thinking=False, preserve_thinking=False + ) + assert chat_template_with_preserved_thinking(configured) == configured + + +@pytest.mark.parametrize("reasoning", [None, "", "\nthought\n"]) +@pytest.mark.parametrize("next_thinking", [False, True]) +@pytest.mark.parametrize( + "tools", + [ + None, + [ + { + "type": "function", + "function": { + "name": "lookup", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + ], +) +def test_deepseek_v4_retains_prior_turns_when_generation_mode_changes( + reasoning, next_thinking, tools +): + from tokenizers import Tokenizer + from tokenizers.models import WordLevel + from transformers import PreTrainedTokenizerFast + + from art.megatron.dsv4.tokenizer import get_dsv4_tokenizer + + tokenizer = get_dsv4_tokenizer( + PreTrainedTokenizerFast(tokenizer_object=Tokenizer(WordLevel({"[UNK]": 0}))) + ) + messages = [ + {"role": "user", "content": "question"}, + {"role": "assistant", "content": "answer", "reasoning_content": reasoning}, + ] + completed = tokenizer.apply_chat_template( + messages, tools=tools, tokenize=False, enable_thinking=reasoning is not None + ) + assert ( + tokenizer.apply_chat_template( + messages, + tools=tools, + tokenize=False, + enable_thinking=reasoning is not None, + chat_template=None, + ) + == completed + ) + continued = tokenizer.apply_chat_template( + [*messages, {"role": "user", "content": "next"}], + tools=tools, + tokenize=False, + enable_thinking=next_thinking, + ) + assert isinstance(completed, str) and isinstance(continued, str) + assert continued.startswith(completed) + if reasoning: + assert reasoning in continued + assert reasoning not in tokenizer.apply_chat_template( + [*messages, {"role": "user", "content": "next"}], + tools=tools, + tokenize=False, + enable_thinking=next_thinking, + preserve_thinking=False, + ) diff --git a/tests/unit/trajectories/test_tokenize.py b/tests/unit/trajectories/test_tokenize.py index 69f727247..f61777852 100644 --- a/tests/unit/trajectories/test_tokenize.py +++ b/tests/unit/trajectories/test_tokenize.py @@ -3421,12 +3421,15 @@ def test_loaded_tokenizers_are_cached_by_model_and_revision( class AutoTokenizer: @staticmethod - def from_pretrained(model: str, *, revision: str | None) -> object: + def from_pretrained( + model: str, *, revision: str | None = None, **kwargs + ) -> object: loaded.append((model, revision)) return object() transformers = ModuleType("transformers") setattr(transformers, "AutoTokenizer", AutoTokenizer) + setattr(transformers, "PreTrainedTokenizerFast", AutoTokenizer) monkeypatch.setitem(sys.modules, "transformers", transformers) _cached_tokenizer.cache_clear() try: @@ -3445,7 +3448,9 @@ def test_deepseek_v4_uses_arts_protocol_renderer( class AutoTokenizer: @staticmethod - def from_pretrained(model: str, *, revision: str | None) -> object: + def from_pretrained( + model: str, *, revision: str | None = None, **kwargs + ) -> object: assert model == "deepseek-ai/DeepSeek-V4-Flash" assert revision is None return raw @@ -3453,6 +3458,7 @@ def from_pretrained(model: str, *, revision: str | None) -> object: transformers = ModuleType("transformers") transformers.__path__ = [] # type: ignore[attr-defined] setattr(transformers, "AutoTokenizer", AutoTokenizer) + setattr(transformers, "PreTrainedTokenizerFast", AutoTokenizer) tokenizer_base = ModuleType("transformers.tokenization_utils_base") setattr(tokenizer_base, "PreTrainedTokenizerBase", object) monkeypatch.setitem(sys.modules, "transformers", transformers) diff --git a/vllm_runtime/hatch_build.py b/vllm_runtime/hatch_build.py new file mode 100644 index 000000000..23f31c9f6 --- /dev/null +++ b/vllm_runtime/hatch_build.py @@ -0,0 +1,15 @@ +"""Bundle ART's shared sources as real files, including in standalone sdists.""" + +from pathlib import Path + +from hatchling.builders.hooks.plugin.interface import BuildHookInterface + + +class CustomBuildHook(BuildHookInterface): + def initialize(self, version, build_data): + root = Path(self.root) + for source in (root / "src/art_vllm_runtime/_shared").glob("*.py"): + if source.is_symlink(): + build_data["force_include"][str(source.resolve())] = str( + source.relative_to(root) + ) diff --git a/vllm_runtime/pyproject.toml b/vllm_runtime/pyproject.toml index 76f8d8d1a..1ed818205 100644 --- a/vllm_runtime/pyproject.toml +++ b/vllm_runtime/pyproject.toml @@ -42,8 +42,8 @@ build-backend = "hatchling.build" [tool.hatch.build.targets.wheel] packages = ["src/art_vllm_runtime"] -[tool.hatch.build] -sources = ["src"] +[tool.hatch.build.hooks.custom] +path = "hatch_build.py" [tool.hatch.metadata] allow-direct-references = true diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/__init__.py b/vllm_runtime/src/art_vllm_runtime/_shared/__init__.py new file mode 100644 index 000000000..e302f59b7 --- /dev/null +++ b/vllm_runtime/src/art_vllm_runtime/_shared/__init__.py @@ -0,0 +1 @@ +"""Pure tokenizer/history helpers shared with ART without training imports.""" diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/append_only.py b/vllm_runtime/src/art_vllm_runtime/_shared/append_only.py new file mode 120000 index 000000000..9b4b33cde --- /dev/null +++ b/vllm_runtime/src/art_vllm_runtime/_shared/append_only.py @@ -0,0 +1 @@ +../../../../src/art_inference/append_only.py \ No newline at end of file diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/chat_template.py b/vllm_runtime/src/art_vllm_runtime/_shared/chat_template.py new file mode 120000 index 000000000..2c3caf7c2 --- /dev/null +++ b/vllm_runtime/src/art_vllm_runtime/_shared/chat_template.py @@ -0,0 +1 @@ +../../../../src/art_inference/chat_template.py \ No newline at end of file diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/token_prefix.py b/vllm_runtime/src/art_vllm_runtime/_shared/token_prefix.py new file mode 120000 index 000000000..aa3856092 --- /dev/null +++ b/vllm_runtime/src/art_vllm_runtime/_shared/token_prefix.py @@ -0,0 +1 @@ +../../../../src/art_inference/token_prefix.py \ No newline at end of file diff --git a/vllm_runtime/src/art_vllm_runtime/_shared/vllm.py b/vllm_runtime/src/art_vllm_runtime/_shared/vllm.py new file mode 120000 index 000000000..448a69e46 --- /dev/null +++ b/vllm_runtime/src/art_vllm_runtime/_shared/vllm.py @@ -0,0 +1 @@ +../../../../src/art_inference/vllm.py \ No newline at end of file diff --git a/vllm_runtime/src/art_vllm_runtime/history.py b/vllm_runtime/src/art_vllm_runtime/history.py new file mode 100644 index 000000000..cdcd3693e --- /dev/null +++ b/vllm_runtime/src/art_vllm_runtime/history.py @@ -0,0 +1,8 @@ +"""ART's shared vLLM history adapter, bundled without training dependencies.""" + +import sys + +from ._shared import vllm as _implementation +from ._shared.vllm import * # noqa: F403 + +sys.modules[__name__] = _implementation diff --git a/vllm_runtime/src/art_vllm_runtime/patches.py b/vllm_runtime/src/art_vllm_runtime/patches.py index 7c9d9c560..f6a16d97a 100644 --- a/vllm_runtime/src/art_vllm_runtime/patches.py +++ b/vllm_runtime/src/art_vllm_runtime/patches.py @@ -9,6 +9,7 @@ def apply_vllm_runtime_patches() -> None: patch_gemma4_moe_lora_support, ) from art_vllm_runtime.glm52_patches import apply_glm52_vllm_runtime_patches + from art_vllm_runtime.history import patch_history from art_vllm_runtime.moe_lora_patches import ( patch_local_3d_moe_dummy_lora, patch_small_batch_moe_lora_intermediate_dtype, @@ -19,6 +20,7 @@ def apply_vllm_runtime_patches() -> None: patch_policy_token_spans() patch_gemma4_moe_lora_support() subclass_chat_completion_request() + patch_history() patch_nonstreaming_chat_response_offload() patch_local_3d_moe_dummy_lora() patch_small_batch_moe_lora_intermediate_dtype()