class Plamo3Parser(ParserEngine):
"""PLaMo3 reasoning and tool-call parser backed by ParserEngine."""
def __init__(
self,
tokenizer: TokenizerLike,
tools: list[Tool] | None = None,
**kwargs,
) -> None:
# Configure PLaMo parsing and cache its multi-token reasoning markers.
chat_kwargs = kwargs.get("chat_template_kwargs", {}) or {}
self.thinking_enabled = chat_kwargs.get("enable_thinking", True)
kwargs.setdefault(
"parser_engine_config",
plamo3_config(thinking=self.thinking_enabled),
)
super().__init__(tokenizer, tools, **kwargs)
self._reasoning_start_token_id_sequence: list[int] = list(
tokenizer.encode(BEGIN_THINK, add_special_tokens=False)
)
self._reasoning_end_token_id_sequence: list[int] = list(
tokenizer.encode(END_THINK, add_special_tokens=False)
)
# Token-ID sequences of every terminal literal, for token-level
# unfinished-marker detection.
self._terminal_token_id_sequences: dict[str, tuple[int, ...]] = {
literal: tuple(tokenizer.encode(literal, add_special_tokens=False))
for literal in self.parser_engine_config.terminal_literals
}
def is_reasoning_end(self, input_ids: list[int]) -> bool:
# Detect PLaMo reasoning boundaries that span multiple token IDs.
if not self.thinking_enabled:
return True
start_ids = self._reasoning_start_token_id_sequence
end_ids = self._reasoning_end_token_id_sequence
for i in range(len(input_ids) - 1, -1, -1):
if input_ids[i : i + len(end_ids)] == end_ids:
return True
if input_ids[i : i + len(start_ids)] == start_ids:
return False
if input_ids[i] == self.vocab.get(EOT):
return False
return False
def extract_content_ids(self, input_ids: list[int]) -> list[int]:
# Extract content after PLaMo's multi-token reasoning end marker.
if not self.thinking_enabled:
return input_ids
end_ids = self._reasoning_end_token_id_sequence
for i in range(len(input_ids) - len(end_ids), -1, -1):
if input_ids[i : i + len(end_ids)] == end_ids:
return input_ids[i + len(end_ids) :]
return input_ids
def _strip_unfinished_marker(self, value: str | None) -> str | None:
if not value:
return value
marker_start = value.rfind("<|plamo:")
if marker_start < 0:
return value or None
suffix = value[marker_start:]
# Strip a trailing partial marker at the token-ID level.
suffix_ids = tuple(
self.model_tokenizer.encode(suffix, add_special_tokens=False)
)
if any(
len(seq) > len(suffix_ids) and seq[: len(suffix_ids)] == suffix_ids
for seq in self._terminal_token_id_sequences.values()
):
return value[:marker_start] or None
return value or None
def _coalesce_finished_tool_events(
self, events: list[SemanticEvent]
) -> list[SemanticEvent]:
combined = []
for (kind, _), group in groupby(
events, key=lambda event: (event.type, event.tool_index)
):
chunks = list(group)
if kind not in (EventType.TOOL_NAME, EventType.ARG_VALUE_CHUNK):
combined.extend(chunks)
continue
if kind == EventType.TOOL_NAME and chunks[0].tool_index < 0:
continue
value = "".join(chunk.value for chunk in chunks)
if kind == EventType.TOOL_NAME:
value = self._strip_unfinished_marker(value) or ""
if value:
combined.append(replace(chunks[0], value=value))
return combined
def _events_to_delta(
self,
events: list[SemanticEvent],
finished: bool = False,
) -> DeltaMessage | None:
# Coalesce flushed tool fragments before stripping truncated markers.
if not finished:
return super()._events_to_delta(events)
events = self._coalesce_finished_tool_events(events)
delta = super()._events_to_delta(events, finished=True)
if delta is None:
return None
delta.reasoning = self._strip_unfinished_marker(delta.reasoning)
if not self.skip_tool_parsing:
delta.content = self._strip_unfinished_marker(delta.content)
# The flush emits the remaining (never-streamed) tool-call
# arguments; an empty remainder must match non-streaming
# extraction (_extract_args_json), which normalizes it to '{}'.
for tc in delta.tool_calls or []:
if tc.function is not None and not (tc.function.arguments or "").strip():
tc.function.arguments = "{}"
if delta.reasoning is None and delta.content is None and not delta.tool_calls:
return None
return delta
def _handle_arg_chunk(
self,
event: SemanticEvent,
deltas: list[DeltaToolCall],
) -> None:
# Strip truncated marker suffixes and preserve the first argument delta.
idx = event.tool_index
slot = self._tool_slots[idx]
if (marker_pos := event.value.rfind("<|plamo:")) >= 0:
stripped = event.value[:marker_pos]
try:
json.loads(slot.args + stripped)
except ValueError:
pass
else:
event = replace(event, value=stripped)
name_sent_before = slot.name_sent
super()._handle_arg_chunk(event, deltas)
if event.value and not name_sent_before and slot.name_sent:
deltas.append(
DeltaToolCall(
index=idx,
function=DeltaFunctionCall(arguments=event.value),
)
)
def _extract_args_json(self, raw_args: str, func_name: str) -> str:
# Return the raw arguments as JSON
# since PLaMo3 generates function names and arguments separately.
return raw_args.strip() or "{}"
def extract_reasoning(
self,
model_output: str,
request: ChatCompletionRequest | ResponsesRequest,
) -> tuple[str | None, str | None]:
reasoning, content = super().extract_reasoning(model_output, request)
return self._strip_unfinished_marker(reasoning), content