diff --git a/README.md b/README.md index 7e1b642..f0c24b4 100644 --- a/README.md +++ b/README.md @@ -19,11 +19,13 @@ uv add 'renderers[transformers]' uv add 'renderers[multimodal]' ``` -A BYO tokenizer must expose `encode`, `decode`, `convert_tokens_to_ids`, token -IDs such as `eos_token_id`, and `return_offsets_mapping=True` through its call -interface. `DefaultRenderer` additionally requires `apply_chat_template`. -This includes text-only Inkling training: `InklingRenderer` loads its -Transformers processor only when image or audio content is actually rendered. +A BYO tokenizer must expose `encode`, `decode`, `convert_tokens_to_ids`, and +token IDs such as `eos_token_id`. Character offsets are optional: tokenizers +supporting `return_offsets_mapping=True` also receive precise per-token +`is_content` attribution; without offsets, renderers return `is_content=[]`. +`DefaultRenderer` additionally requires `apply_chat_template`. This includes +text-only Inkling training: `InklingRenderer` loads its Transformers processor +only when image or audio content is actually rendered. ## At a glance diff --git a/pyproject.toml b/pyproject.toml index 845e70a..7aed44c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -39,8 +39,8 @@ dependencies = [ [project.optional-dependencies] # Tokenizer loading uses Hugging Face. Text-only renderers can instead be -# constructed with an offset-capable BYO tokenizer and do not import this -# dependency. +# constructed with a compatible BYO tokenizer and do not import this +# dependency. Character offsets are optional. transformers = [ # Keep this floor compatible with prime-rl's transformers pin. Inkling's # tokenizer and text-only renderer work on older releases; image/audio diff --git a/renderers/__init__.py b/renderers/__init__.py index 7635ddc..dc9f369 100644 --- a/renderers/__init__.py +++ b/renderers/__init__.py @@ -15,6 +15,7 @@ Message, MultiModalData, MultimodalRenderer, + OffsetTokenizer, ParsedResponse, ParsedToolCall, PlaceholderRange, @@ -181,6 +182,7 @@ def __dir__() -> list[str]: "Nemotron3RendererConfig", "Nemotron3UltraRenderer", "Nemotron3UltraRendererConfig", + "OffsetTokenizer", "OverlongPromptError", "ParsedResponse", "ParsedToolCall", diff --git a/renderers/base.py b/renderers/base.py index e82eb7d..e98637b 100644 --- a/renderers/base.py +++ b/renderers/base.py @@ -11,6 +11,7 @@ Literal, Protocol, TypedDict, + cast, runtime_checkable, ) @@ -243,7 +244,9 @@ class RenderedTokens: Empty ``sampled_mask`` (``[]``) means the renderer doesn't provide this signal — consumers should fall back to attribution-only masking. ``DefaultRenderer`` leaves it empty because the Jinja - template is opaque; hand-coded renderers populate it. + template is opaque. Hand-coded renderers normally populate it; a + renderer whose sampled/scaffold boundary depends on character + attribution may leave it empty for an offsetless tokenizer. ``is_content`` is a per-token signal generalizing the "scaffold vs body" distinction across all roles: ``True`` iff the token was @@ -269,7 +272,8 @@ class RenderedTokens: Empty ``is_content`` (``[]``) — like ``sampled_mask`` — means the renderer doesn't provide the signal. ``DefaultRenderer`` leaves it - empty for the same reason. + empty because its Jinja template is opaque; all renderers leave it + empty when the supplied tokenizer cannot return character offsets. ``message_tool_names`` is the per-message tool function name list, parallel to ``message_roles`` (same length). For tool-role @@ -476,7 +480,7 @@ def content_token_spans_by_role(self) -> dict[str, list[tuple[int, int]]]: Returns an empty dict when :attr:`is_content` or :attr:`message_roles` is empty (renderer didn't populate the - signal — e.g. ``DefaultRenderer``). + signal — e.g. ``DefaultRenderer`` or an offsetless tokenizer). Intended for selective loss masking: SFT on tool response bodies while RL acts only on assistant turns is the canonical @@ -667,8 +671,8 @@ class Tokenizer(Protocol): Hugging Face tokenizers satisfy this protocol, as can lightweight BYO adapters around ``tokenizers.Tokenizer`` or another tokenizer backend. Keeping the renderer-facing contract here makes ``transformers`` optional - for text rendering. Offset-capable ``__call__`` behavior is required by - :func:`attribute_text_segments` to preserve BPE boundary attribution. + for text rendering. Character offsets are a separate optional capability; + see :class:`OffsetTokenizer`. """ name_or_path: str @@ -681,6 +685,17 @@ def decode(self, token_ids: Any, *args: Any, **kwargs: Any) -> str: ... def convert_tokens_to_ids(self, tokens: Any) -> Any: ... + +@runtime_checkable +class OffsetTokenizer(Tokenizer, Protocol): + """Tokenizer that can return character offsets alongside token IDs. + + Hand-coded renderers use this optional capability to distinguish caller + content from adjacent template scaffold without changing the underlying + BPE pass. A basic :class:`Tokenizer` remains sufficient for rendering token + IDs; when offsets are unavailable, renderers leave ``is_content`` empty. + """ + def __call__(self, *args: Any, **kwargs: Any) -> Any: ... @@ -1064,7 +1079,7 @@ def is_multimodal(r: object) -> bool: "Install the optional dependency with " "`pip install 'renderers[transformers]'` (or " "`uv add 'renderers[transformers]'`). Text-only renderers work without " - "it when constructed with an offset-capable tokenizer object." + "it when constructed with a compatible tokenizer object." ) @@ -1759,36 +1774,122 @@ def trim_to_turn_close( return previous_ids -def _get_offset_tokenizer(tokenizer): - """Assert ``tokenizer`` supports ``return_offsets_mapping=True``. +class AttributedTextSegments(list[tuple[int, bool]]): + """Token/content pairs with an explicit attribution-availability flag.""" + + def __init__( + self, + values=(), + *, + has_content_attribution: bool, + ) -> None: + super().__init__(values) + self.has_content_attribution = has_content_attribution + + +def _get_offset_tokenizer(tokenizer: Tokenizer) -> OffsetTokenizer | None: + """Return ``tokenizer`` when it supports character offsets, else ``None``. Hand-coded renderers concatenate scaffold + body in one BPE pass to preserve cross-boundary merges, then attribute each resulting token back to its source segment via the fast tokenizer's - ``offset_mapping`` (see :func:`attribute_text_segments`). The - contract: every BYO tokenizer must be a fast tokenizer with offset - support. Tokenizers loaded via :func:`load_tokenizer` are - ``PreTrainedTokenizerFast`` instances that satisfy this trivially. + ``offset_mapping`` (see :func:`attribute_text_segments`). Tokenizers + loaded via :func:`load_tokenizer` are ``PreTrainedTokenizerFast`` + instances that satisfy this capability, but BYO tokenizers need not. """ + call = getattr(tokenizer, "__call__", None) + if not callable(call): + return None try: - tokenizer("a", add_special_tokens=False, return_offsets_mapping=True) - except (NotImplementedError, ValueError, TypeError) as exc: - raise RuntimeError( - "Hand-coded renderers require a fast tokenizer with " - "``return_offsets_mapping=True`` support for body/scaffold " - "attribution. Pass a tokenizer loaded via " - "``renderers.base.load_tokenizer``, or any " - "``transformers.PreTrainedTokenizerFast`` instance." - ) from exc - return tokenizer + encoding = call("a", add_special_tokens=False, return_offsets_mapping=True) + encoding["input_ids"] + encoding["offset_mapping"] + except (KeyError, NotImplementedError, TypeError, ValueError): + return None + return cast(OffsetTokenizer, tokenizer) + + +def _infer_offsets_from_decode( + tokenizer: Tokenizer, + token_ids: list[int], + text: str, +) -> list[tuple[int, int]] | None: + """Recover token character spans from an exact decoder round-trip. + + This is a narrow fallback for metadata that does not require exposing + content attribution. Some renderers join text from multiple messages in a + single BPE pass, so they still need to associate the resulting tokens with + the right message when a BYO tokenizer has no native offset mapping. + + Decoding individual tokens is linear and exact for the common BPE/SentencePiece + backends. Byte-fallback tokenizers can require multiple tokens before text + becomes valid, so a validated cumulative-prefix pass handles that case. + If either strategy cannot reconstruct ``text`` exactly, callers must use a + conservative renderer-specific message-index fallback. This helper never + upgrades the tokenizer's content-attribution capability: ``is_content`` + remains unavailable without native offsets. + """ + + def decode(ids: list[int]) -> str | None: + variants = ( + {"skip_special_tokens": False, "clean_up_tokenization_spaces": False}, + {"skip_special_tokens": False}, + {}, + ) + for kwargs in variants: + try: + decoded = tokenizer.decode(ids, **kwargs) + except TypeError: + continue + except (KeyError, NotImplementedError, UnicodeError, ValueError): + return None + return decoded if isinstance(decoded, str) else None + return None + + pieces: list[str] = [] + for token_id in token_ids: + piece = decode([token_id]) + if piece is None: + break + pieces.append(piece) + if len(pieces) == len(token_ids) and "".join(pieces) == text: + offsets: list[tuple[int, int]] = [] + position = 0 + for piece in pieces: + end = position + len(piece) + offsets.append((position, end)) + position = end + return offsets + + offsets = [] + previous_end = 0 + for end_index in range(1, len(token_ids) + 1): + prefix = decode(token_ids[:end_index]) + if prefix is None or len(prefix) < previous_end or not text.startswith(prefix): + return None + current_end = len(prefix) + offsets.append((previous_end, current_end)) + previous_end = current_end + if previous_end != len(text): + return None + return offsets + + +def _content_mask_or_empty( + tokenizer: Tokenizer, content_mask: list[bool] +) -> list[bool]: + """Return exact content attribution, or the empty-list unavailable sentinel.""" + if _get_offset_tokenizer(tokenizer) is None: + return [] + return content_mask def attribute_text_segments( - tokenizer, + tokenizer: Tokenizer, segments: "list[tuple[str, bool]]", *, overlap_is_content: bool = False, -) -> "list[tuple[int, bool]]": +) -> AttributedTextSegments: """Tokenize concatenated segments as a single BPE pass and return ``(token_id, is_content)`` pairs. @@ -1816,23 +1917,29 @@ def attribute_text_segments( every body byte inside the ``is_content=True`` run at the cost of a few adjacent wrap bytes. - Requires a HuggingFace fast tokenizer with offset tracking. Every - model in ``MODEL_RENDERER_MAP`` ships one, so the offset lookup - always succeeds for tokenizers obtained via :func:`load_tokenizer`. - BYO tokenizers must be a ``PreTrainedTokenizerFast`` (or anything - else exposing ``return_offsets_mapping=True``); slow tokenizers - aren't supported — BPE drift at the wrap/body boundary would - defeat the whole point. + When ``tokenizer`` implements :class:`OffsetTokenizer`, the result's + ``has_content_attribution`` flag is true and each bool is exact. For a + basic :class:`Tokenizer`, the joined text is still encoded in one pass so + token IDs remain identical, but the bools are placeholders and + ``has_content_attribution`` is false. Renderers propagate that state as an + empty ``RenderedTokens.is_content`` list rather than exposing a partial or + inaccurate mask. Empty input or empty joined text returns an empty list. """ if not segments: - return [] + return AttributedTextSegments([], has_content_attribution=True) full_text = "".join(text for text, _ in segments) if not full_text: - return [] + return AttributedTextSegments([], has_content_attribution=True) offset_tokenizer = _get_offset_tokenizer(tokenizer) + if offset_tokenizer is None: + token_ids = tokenizer.encode(full_text, add_special_tokens=False) + return AttributedTextSegments( + ((token_id, False) for token_id in token_ids), + has_content_attribution=False, + ) encoding = offset_tokenizer( full_text, add_special_tokens=False, @@ -1887,7 +1994,7 @@ def attribute_text_segments( # the last non-empty segment's bit. pass out.append((tok_id, is_content)) - return out + return AttributedTextSegments(out, has_content_attribution=True) def reject_assistant_in_extension(new_messages: list[Message]) -> bool: diff --git a/renderers/deepseek_v3.py b/renderers/deepseek_v3.py index 56b901f..5fbf7c5 100644 --- a/renderers/deepseek_v3.py +++ b/renderers/deepseek_v3.py @@ -20,6 +20,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -259,7 +260,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -408,7 +409,9 @@ def emit_text( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/gemma4.py b/renderers/gemma4.py index c519465..8de3878 100644 --- a/renderers/gemma4.py +++ b/renderers/gemma4.py @@ -34,6 +34,7 @@ ToolCallParseStatus, ToolSpec, Tokenizer, + _content_mask_or_empty, _require_transformers, attribute_text_segments, extract_message_tool_names, @@ -1025,7 +1026,7 @@ def render( token_ids=em.token_ids, message_indices=em.message_indices, sampled_mask=em.sampled, - is_content=em.is_content, + is_content=_content_mask_or_empty(self._tokenizer, em.is_content), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), multi_modal_data=multi_modal_data, @@ -1356,7 +1357,7 @@ def bridge_to_next_turn( token_ids=em.token_ids, message_indices=em.message_indices, sampled_mask=em.sampled, - is_content=em.is_content, + is_content=_content_mask_or_empty(self._tokenizer, em.is_content), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), multi_modal_data=self._merge_multi_modal_data( diff --git a/renderers/glm45.py b/renderers/glm45.py index d1b72c1..d65a342 100644 --- a/renderers/glm45.py +++ b/renderers/glm45.py @@ -19,6 +19,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -261,7 +262,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -448,7 +449,9 @@ def emit_text_segments( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/glm5.py b/renderers/glm5.py index 961c002..93d8670 100644 --- a/renderers/glm5.py +++ b/renderers/glm5.py @@ -20,6 +20,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -283,7 +284,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -465,7 +466,9 @@ def emit_text_segments( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/gpt_oss.py b/renderers/gpt_oss.py index 9cbf316..f8589e0 100644 --- a/renderers/gpt_oss.py +++ b/renderers/gpt_oss.py @@ -55,6 +55,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, extract_message_tool_names, reject_assistant_in_extension, resolve_thinking_retention, @@ -461,7 +462,7 @@ def emit_harmony_message( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -597,7 +598,9 @@ def bridge_to_next_turn( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/hy3.py b/renderers/hy3.py index 04a327f..1a55075 100644 --- a/renderers/hy3.py +++ b/renderers/hy3.py @@ -33,6 +33,9 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, + _get_offset_tokenizer, + _infer_offsets_from_decode, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -238,13 +241,34 @@ def _attribute_segments( if not segments: return [] full_text = "".join(text for text, _, _ in segments) - encoding = self._tokenizer( - full_text, - add_special_tokens=False, - return_offsets_mapping=True, - ) - token_ids = list(encoding["input_ids"]) - offsets = list(encoding["offset_mapping"]) + offset_tokenizer = _get_offset_tokenizer(self._tokenizer) + if offset_tokenizer is None: + token_ids = self._encode(full_text) + offsets = _infer_offsets_from_decode( + self._tokenizer, + token_ids, + full_text, + ) + if offsets is None: + # Token IDs remain exact even when a lossy decoder prevents + # reconstructing boundaries. Associate the opaque joined run + # with a contributing caller system message rather than + # silently classifying its body as global scaffold. + fallback_idx = next( + (msg_idx for text, _, msg_idx in segments if text and msg_idx >= 0), + -1, + ) + return [(token_id, False, fallback_idx) for token_id in token_ids] + has_content_attribution = False + else: + encoding = offset_tokenizer( + full_text, + add_special_tokens=False, + return_offsets_mapping=True, + ) + token_ids = list(encoding["input_ids"]) + offsets = list(encoding["offset_mapping"]) + has_content_attribution = True spans: list[tuple[int, int, bool, int]] = [] pos = 0 @@ -254,13 +278,19 @@ def _attribute_segments( total_len = pos out: list[tuple[int, bool, int]] = [] - last = (spans[-1][2], spans[-1][3]) + last = ( + spans[-1][2] if has_content_attribution else False, + spans[-1][3], + ) for tok_id, (start, _end) in zip(token_ids, offsets): attr = last if start < total_len: for seg_start, seg_end, seg_is_content, seg_idx in spans: if seg_start <= start < seg_end: - attr = (seg_is_content, seg_idx) + attr = ( + seg_is_content if has_content_attribution else False, + seg_idx, + ) break out.append((tok_id, attr[0], attr[1])) return out @@ -431,7 +461,7 @@ def emit_attributed(segments: list[tuple[str, bool, int]]) -> None: token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -699,7 +729,9 @@ def emit_text_segments(segments: list[tuple[str, bool]], msg_idx: int) -> None: token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/inkling.py b/renderers/inkling.py index 80f2f2b..6ef9074 100644 --- a/renderers/inkling.py +++ b/renderers/inkling.py @@ -48,6 +48,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, _require_transformers, extract_message_tool_names, reject_assistant_in_extension, @@ -492,7 +493,7 @@ def role_open(_role_id=role_id, _i=i): token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=tool_names, multi_modal_data=mm_data, @@ -989,7 +990,7 @@ def role_open(_role_id=role_id, _i=i): token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=tool_names, multi_modal_data=mm_data, diff --git a/renderers/kimi_k2.py b/renderers/kimi_k2.py index 64f4d00..e3175a4 100644 --- a/renderers/kimi_k2.py +++ b/renderers/kimi_k2.py @@ -22,6 +22,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, extract_message_tool_names, reject_assistant_in_extension, resolve_thinking_retention, @@ -309,7 +310,7 @@ def emit_text( token_ids=token_ids, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in caller_messages], message_tool_names=extract_message_tool_names(caller_messages), ) @@ -464,7 +465,9 @@ def emit_text( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/kimi_k25.py b/renderers/kimi_k25.py index ef7e4e3..e6dfe18 100644 --- a/renderers/kimi_k25.py +++ b/renderers/kimi_k25.py @@ -35,6 +35,7 @@ ToolCallParseStatus, ToolSpec, Tokenizer, + _content_mask_or_empty, _require_transformers, extract_message_tool_names, reject_assistant_in_extension, @@ -955,7 +956,7 @@ def emit_image( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), multi_modal_data=mm_data, @@ -1214,7 +1215,7 @@ def emit_image( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=bridge_roles, message_tool_names=bridge_tool_names, ) @@ -1228,7 +1229,7 @@ def emit_image( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=bridge_roles, message_tool_names=bridge_tool_names, multi_modal_data=mm_data, diff --git a/renderers/laguna_xs2.py b/renderers/laguna_xs2.py index a112f76..2ce25f8 100644 --- a/renderers/laguna_xs2.py +++ b/renderers/laguna_xs2.py @@ -57,6 +57,8 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, + _infer_offsets_from_decode, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -327,7 +329,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -484,7 +486,9 @@ def emit_text_segments( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) @@ -781,12 +785,52 @@ def emit_text_segments( tool_text += "" header_segs.append((tool_text, False)) header_segs.append(("\n", False)) - for tok_id, is_content in attribute_text_segments( + attributed = attribute_text_segments( self._tokenizer, header_segs, overlap_is_content=True - ): + ) + fallback_indices: list[int] | None = None + if not attributed.has_content_attribution: + full_header = "".join(text for text, _ in header_segs) + offsets = _infer_offsets_from_decode( + self._tokenizer, + [token_id for token_id, _ in attributed], + full_header, + ) + if offsets is None: + fallback_index = 0 if caller_has_system and has_sys_content else -1 + fallback_indices = [fallback_index] * len(attributed) + else: + content_spans: list[tuple[int, int]] = [] + position = 0 + for text, is_content in header_segs: + end = position + len(text) + if is_content: + content_spans.append((position, end)) + position = end + fallback_indices = [ + 0 + if ( + any( + span_start < end and start < span_end + for span_start, span_end in content_spans + ) + if end > start + else any( + span_start <= start < span_end + for span_start, span_end in content_spans + ) + ) + else -1 + for start, end in offsets + ] + for position, (tok_id, is_content) in enumerate(attributed): + if fallback_indices is not None: + message_index = fallback_indices[position] + else: + message_index = 0 if is_content else -1 emit_special( tok_id, - 0 if is_content else -1, + message_index, is_sampled=False, is_content=is_content, ) @@ -840,7 +884,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -958,7 +1002,9 @@ def emit_text_segments( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/llama_3.py b/renderers/llama_3.py index 950d1f2..89205e9 100644 --- a/renderers/llama_3.py +++ b/renderers/llama_3.py @@ -47,6 +47,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -388,7 +389,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -520,7 +521,9 @@ def emit_text_segments( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/minimax_m2.py b/renderers/minimax_m2.py index e9a2bda..312879e 100644 --- a/renderers/minimax_m2.py +++ b/renderers/minimax_m2.py @@ -20,6 +20,8 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, + _get_offset_tokenizer, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -164,9 +166,14 @@ def emit_token_overlap_body( BPE merges it with the wrap's trailing byte (``>The`` → single token). """ - from renderers.base import _get_offset_tokenizer - offset_tok = _get_offset_tokenizer(self._tokenizer) + if offset_tok is None: + ids = self._encode(full_text) + tokens.extend(ids) + indices.extend([msg_idx] * len(ids)) + sampled.extend([is_sampled] * len(ids)) + content_mask.extend([False] * len(ids)) + return encoding = offset_tok( full_text, add_special_tokens=False, return_offsets_mapping=True ) @@ -274,7 +281,7 @@ def emit_token_overlap_body( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -400,9 +407,14 @@ def emit_token_overlap_body( *, is_sampled: bool, ) -> None: - from renderers.base import _get_offset_tokenizer - offset_tok = _get_offset_tokenizer(self._tokenizer) + if offset_tok is None: + ids = self._encode(full_text) + ext.extend(ids) + ext_indices.extend([msg_idx] * len(ids)) + ext_sampled.extend([is_sampled] * len(ids)) + ext_content.extend([False] * len(ids)) + return encoding = offset_tok( full_text, add_special_tokens=False, return_offsets_mapping=True ) @@ -462,7 +474,9 @@ def emit_token_overlap_body( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/nemotron3.py b/renderers/nemotron3.py index 97a53c1..bdf26af 100644 --- a/renderers/nemotron3.py +++ b/renderers/nemotron3.py @@ -23,6 +23,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -483,7 +484,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in original_messages], message_tool_names=extract_message_tool_names(original_messages), ) @@ -667,7 +668,9 @@ def emit_text_segments( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/prime_qwen3.py b/renderers/prime_qwen3.py index 8ff88f5..9a4c537 100644 --- a/renderers/prime_qwen3.py +++ b/renderers/prime_qwen3.py @@ -12,6 +12,8 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, + _get_offset_tokenizer, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -315,8 +317,12 @@ def render( return RenderedTokens( token_ids=builder.token_ids, message_indices=builder.message_indices, - sampled_mask=builder.sampled_mask, - is_content=builder.is_content, + sampled_mask=( + builder.sampled_mask + if _get_offset_tokenizer(self._tokenizer) is not None + else [] + ), + is_content=_content_mask_or_empty(self._tokenizer, builder.is_content), message_roles=[message.get("role") or "" for message in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -683,8 +689,12 @@ def bridge_to_next_turn( return RenderedTokens( token_ids=builder.token_ids, message_indices=builder.message_indices, - sampled_mask=builder.sampled_mask, - is_content=builder.is_content, + sampled_mask=( + builder.sampled_mask + if _get_offset_tokenizer(self._tokenizer) is not None + else [] + ), + is_content=_content_mask_or_empty(self._tokenizer, builder.is_content), message_roles=[message.get("role") or "" for message in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/qwen3.py b/renderers/qwen3.py index 97c21f4..12d1658 100644 --- a/renderers/qwen3.py +++ b/renderers/qwen3.py @@ -26,6 +26,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, attribute_text_segments, extract_message_tool_names, reject_assistant_in_extension, @@ -268,7 +269,7 @@ def emit_text_segments( token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), ) @@ -433,7 +434,9 @@ def emit_text_segments( token_ids=previous_ids + ext, message_indices=[-1] * len(previous_ids) + ext_indices, sampled_mask=[False] * total_len, - is_content=[False] * len(previous_ids) + ext_content, + is_content=_content_mask_or_empty( + self._tokenizer, [False] * len(previous_ids) + ext_content + ), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), ) diff --git a/renderers/qwen35.py b/renderers/qwen35.py index 2e9de34..79e884d 100644 --- a/renderers/qwen35.py +++ b/renderers/qwen35.py @@ -33,6 +33,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, _require_transformers, attribute_text_segments, extract_message_tool_names, @@ -637,7 +638,7 @@ def flush_buf() -> None: token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), multi_modal_data=mm_data, @@ -931,7 +932,7 @@ def flush_buf() -> None: token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=bridge_roles, message_tool_names=bridge_tool_names, ) @@ -945,7 +946,7 @@ def flush_buf() -> None: token_ids=tokens, message_indices=indices, sampled_mask=sampled, - is_content=content_mask, + is_content=_content_mask_or_empty(self._tokenizer, content_mask), message_roles=bridge_roles, message_tool_names=bridge_tool_names, multi_modal_data=mm_data, diff --git a/renderers/qwen3_vl.py b/renderers/qwen3_vl.py index 121a77f..481b910 100644 --- a/renderers/qwen3_vl.py +++ b/renderers/qwen3_vl.py @@ -41,6 +41,7 @@ RenderedTokens, ToolSpec, Tokenizer, + _content_mask_or_empty, _require_transformers, attribute_text_segments, extract_message_tool_names, @@ -277,8 +278,8 @@ def _flush(self) -> None: self.is_content.extend([first_ic] * len(ids)) return # Mixed body/scaffold flush — encode once and attribute back to - # each segment via the fast tokenizer's offset_mapping. Requires - # a tokenizer (not just the encode fn) to look up offsets. + # each segment via offset_mapping when available. A basic tokenizer + # still preserves the joined token IDs but leaves attribution empty. assert self._tokenizer is not None, ( "_Emitter mixed-is_content flush requires a tokenizer; " "pass one to the constructor." @@ -619,7 +620,7 @@ def render_media_content(content: Any) -> None: token_ids=em.token_ids, message_indices=em.message_indices, sampled_mask=em.sampled, - is_content=em.is_content, + is_content=_content_mask_or_empty(self._tokenizer, em.is_content), message_roles=[m.get("role") or "" for m in messages], message_tool_names=extract_message_tool_names(messages), multi_modal_data=mm_data, @@ -864,7 +865,7 @@ def render_media_content(content: Any) -> None: token_ids=em.token_ids, message_indices=em.message_indices, sampled_mask=em.sampled, - is_content=em.is_content, + is_content=_content_mask_or_empty(self._tokenizer, em.is_content), message_roles=[m.get("role") or "" for m in new_messages], message_tool_names=extract_message_tool_names(new_messages), multi_modal_data=mm_data, diff --git a/tests/test_load_tokenizer.py b/tests/test_load_tokenizer.py index ea15d6a..aa9855e 100644 --- a/tests/test_load_tokenizer.py +++ b/tests/test_load_tokenizer.py @@ -12,8 +12,6 @@ from types import SimpleNamespace from unittest.mock import patch -import pytest - from renderers import base from renderers.base import TOKENIZER_SOURCE_OVERRIDES, TRUSTED_REVISIONS, load_tokenizer @@ -122,22 +120,29 @@ def test_tokenizer_source_overrides_are_exact_llama_mirrors(): } -def test_get_offset_tokenizer_rejects_offsetless_byo(): - """BYO tokenizers without ``return_offsets_mapping`` support raise a - clear error. Hand-coded renderers concatenate scaffold + body in one - BPE pass and attribute tokens via the fast tokenizer's offset map; - no transparent reload-from-name_or_path fallback exists. The - contract is: pass a fast tokenizer or get a loud error at construct - time, not silent BPE drift at the wrap/body boundary.""" +def test_offsetless_byo_preserves_ids_without_content_attribution(): + """Offsetless BYO tokenizers keep the joined BPE pass but omit attribution.""" class _NoOffsets: name_or_path = "anywhere/anything" + def encode(self, text, *, add_special_tokens=False): + assert add_special_tokens is False + return [len(text)] + def __call__(self, *args, **kwargs): raise NotImplementedError("BYO tokenizer has no offsets") - with pytest.raises(RuntimeError, match="fast tokenizer.*offsets"): - base._get_offset_tokenizer(_NoOffsets()) + tokenizer = _NoOffsets() + assert base._get_offset_tokenizer(tokenizer) is None + + attributed = base.attribute_text_segments( + tokenizer, + [("user\n", False), ("hello", True)], + ) + + assert attributed == [(10, False)] + assert attributed.has_content_attribution is False # --------------------------------------------------------------------------- diff --git a/tests/test_offsetless_tokenizers.py b/tests/test_offsetless_tokenizers.py new file mode 100644 index 0000000..255c2c6 --- /dev/null +++ b/tests/test_offsetless_tokenizers.py @@ -0,0 +1,223 @@ +"""Renderer-wide contract tests for BYO tokenizers without character offsets.""" + +from __future__ import annotations + +from typing import Any + +import pytest + +from renderers import create_renderer +from renderers.base import load_tokenizer + + +class OffsetlessTokenizer: + """Delegate the basic tokenizer protocol without exposing ``__call__``.""" + + def __init__(self, tokenizer: Any): + self._tokenizer = tokenizer + + def __getattr__(self, name: str) -> Any: + if name == "__call__": + raise AttributeError(name) + return getattr(self._tokenizer, name) + + def encode(self, *args: Any, **kwargs: Any) -> list[int]: + return self._tokenizer.encode(*args, **kwargs) + + def decode(self, *args: Any, **kwargs: Any) -> str: + return self._tokenizer.decode(*args, **kwargs) + + def convert_tokens_to_ids(self, *args: Any, **kwargs: Any) -> Any: + return self._tokenizer.convert_tokens_to_ids(*args, **kwargs) + + def apply_chat_template(self, *args: Any, **kwargs: Any) -> Any: + return self._tokenizer.apply_chat_template(*args, **kwargs) + + +TOOLS = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the current weather for a city", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], + }, + }, + } +] + + +CASES = [ + pytest.param( + [ + {"role": "system", "content": "You are concise."}, + {"role": "user", "content": "Hello!"}, + ], + None, + True, + id="system-generation-prompt", + ), + pytest.param( + [ + {"role": "system", "content": "You are a weather assistant."}, + {"role": "user", "content": "Weather?"}, + ], + TOOLS, + False, + id="system-and-tools", + ), + pytest.param( + [ + {"role": "user", "content": "What is 2+2?"}, + { + "role": "assistant", + "reasoning_content": "Simple arithmetic", + "content": "4", + }, + {"role": "user", "content": "And 3+3?"}, + ], + None, + True, + id="reasoning-history", + ), + pytest.param( + [ + {"role": "system", "content": "You are helpful."}, + {"role": "user", "content": "Weather in Paris?"}, + { + "role": "assistant", + "content": "Let me check.", + "tool_calls": [ + { + "function": { + "name": "get_weather", + "arguments": {"city": "Paris"}, + } + } + ], + }, + {"role": "tool", "content": '{"temp": 20}'}, + {"role": "assistant", "content": "It is 20 degrees."}, + ], + TOOLS, + False, + id="full-tool-cycle", + ), +] + + +def _offsetless_renderer(renderer: Any) -> Any: + return type(renderer)(OffsetlessTokenizer(renderer._tokenizer), renderer.config) + + +def _assert_offsetless_contract(expected: Any, actual: Any, renderer: Any) -> None: + assert actual.token_ids == expected.token_ids + assert actual.message_indices == expected.message_indices + if type(renderer).__name__ == "PrimeQwen3Renderer": + assert actual.sampled_mask == [] + else: + assert actual.sampled_mask == expected.sampled_mask + assert actual.is_content == [] + assert actual.message_roles == expected.message_roles + assert actual.message_tool_names == expected.message_tool_names + assert actual.multi_modal_data == expected.multi_modal_data + + +@pytest.mark.parametrize("messages,tools,add_generation_prompt", CASES) +def test_offsetless_contract_across_renderer_matrix( + model_name: str, + renderer: Any, + messages: list[dict[str, Any]], + tools: list[dict[str, Any]] | None, + add_generation_prompt: bool, +) -> None: + """Every configured renderer preserves tokens and non-content metadata.""" + + expected = renderer.render( + messages, + tools=tools, + add_generation_prompt=add_generation_prompt, + ) + offsetless = _offsetless_renderer(renderer) + actual = offsetless.render( + messages, + tools=tools, + add_generation_prompt=add_generation_prompt, + ) + + _assert_offsetless_contract(expected, actual, renderer) + + +def test_offsetless_contract_across_renderer_bridges( + model_name: str, + renderer: Any, +) -> None: + """Bridge paths preserve their tokens and metadata without offsets.""" + + prior_messages = [ + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi!"}, + ] + new_messages = [{"role": "user", "content": "Tell me more."}] + prior = renderer.render(prior_messages) + offsetless = _offsetless_renderer(renderer) + + expected = renderer.bridge_to_next_turn( + prior.token_ids, + [], + new_messages, + ) + actual = offsetless.bridge_to_next_turn( + prior.token_ids, + [], + new_messages, + ) + + assert (actual is None) == (expected is None) + if expected is not None and actual is not None: + _assert_offsetless_contract(expected, actual, renderer) + + +def test_hy3_offsetless_preserves_multiple_system_message_indices() -> None: + """Hy3's joined system header must not erase either caller message.""" + + tokenizer = load_tokenizer("tencent/Hy3") + renderer = create_renderer(tokenizer) + messages = [ + {"role": "system", "content": "First instruction."}, + {"role": "system", "content": "Second instruction."}, + {"role": "user", "content": "Hello"}, + ] + + expected = renderer.render(messages, tools=TOOLS) + actual = _offsetless_renderer(renderer).render(messages, tools=TOOLS) + + _assert_offsetless_contract(expected, actual, renderer) + assert {0, 1}.issubset(actual.message_indices) + + +@pytest.mark.parametrize( + "deepseek_model", + ["deepseek-ai/DeepSeek-V3", "deepseek-ai/DeepSeek-R1"], +) +def test_deepseek_offsetless_contract(deepseek_model: str) -> None: + """DeepSeek variants live outside the shared parity matrix.""" + + tokenizer = load_tokenizer(deepseek_model) + renderer = create_renderer(tokenizer) + messages = [ + {"role": "system", "content": "You are concise."}, + {"role": "user", "content": "What is 2+2?"}, + {"role": "assistant", "content": "4"}, + ] + + expected = renderer.render(messages, add_generation_prompt=True) + actual = _offsetless_renderer(renderer).render( + messages, + add_generation_prompt=True, + ) + + _assert_offsetless_contract(expected, actual, renderer)