Skip to content

vllm.parser.plamo3

Classes:

  • Plamo3Parser –

    PLaMo3 reasoning and tool-call parser backed by ParserEngine.

Plamo3Parser

Bases: ParserEngine

PLaMo3 reasoning and tool-call parser backed by ParserEngine.

Source code in vllm/parser/plamo3.py
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