This commit is contained in:
Evan
2026-03-24 11:54:04 +00:00
committed by rltakashige
parent 6355f3d8fb
commit a9c7b1c68a
70 changed files with 235 additions and 221 deletions
+5 -1
View File
@@ -5,7 +5,11 @@ description = "Add your description here"
readme = "README.md"
authors = [{ name = "Evan", email = "[email protected]" }]
requires-python = ">=3.13"
dependencies = ["pydantic>=2.13.0b2"]
dependencies = [
"mlx-lm", # TODO: depend on transformers or other
"openai-harmony", # inherit from workspace
"pydantic", # inherit from workspace
]
[build-system]
requires = ["uv_build>=0.9.24,<0.10.0"]
+10 -9
View File
@@ -1,14 +1,15 @@
from abc import ABC, abstractmethod
from collections.abc import Callable, Iterable
from typing import Self
from exo_core.types.tasks import TaskId
class Cancelled: pass
class Cancelled:
pass
class Finished: pass
class Finished:
pass
CANCEL_ALL_TASKS = TaskId("CANCEL_TALL_TASKS")
@@ -45,12 +46,12 @@ class Engine[TaskType, ResponseType](ABC):
class EngineBuilder[SetupType, TaskType, ResponseType](ABC):
@classmethod
@abstractmethod
def create(
cls,
bound_instance: SetupType,
) -> Self: ...
# @classmethod
# @abstractmethod
# def create(
# cls,
# bound_instance: SetupType,
# ) -> Self: ...
@abstractmethod
def connect(self) -> None: ...
@@ -0,0 +1,382 @@
from collections.abc import Generator
from functools import cache
from typing import Any
from loguru import logger
from mlx_engine.utils_mlx import (
detect_thinking_prompt_suffix,
)
from mlx_lm.tokenizer_utils import TokenizerWrapper
from openai_harmony import (
HarmonyEncodingName,
HarmonyError,
Role,
StreamableParser,
load_harmony_encoding,
)
from exo_core.tokenizers.tool_parsers import ToolParser
from exo_core.types.common import ModelId
from exo_core.types.runner_response import (
GenerationResponse,
ToolCallItem,
ToolCallResponse,
)
@cache
def get_gpt_oss_encoding():
encoding = load_harmony_encoding(HarmonyEncodingName.HARMONY_GPT_OSS)
return encoding
def apply_all_parsers(
receiver: Generator[GenerationResponse | None],
prompt: str,
tool_parser: ToolParser | None,
tokenizer: TokenizerWrapper,
model_id: ModelId,
tools: list[dict[str, Any]] | None,
) -> Generator[GenerationResponse | ToolCallResponse | None]:
gen = receiver
if tokenizer.has_thinking:
gen = parse_thinking_models(
gen,
tokenizer.think_start,
tokenizer.think_end,
starts_in_thinking=detect_thinking_prompt_suffix(prompt, tokenizer),
)
lower = model_id.normalize().lower()
if "gpt-oss" in lower or "gpt_oss" in lower:
gen = parse_gpt_oss(gen)
elif "deepseek" in lower:
gen = parse_deepseek_v32(gen)
elif tool_parser:
gen = parse_tool_calls(gen, tool_parser, tools)
return gen
_GPT_OSS_CHANNEL_TOKEN = 200005
_GPT_OSS_MESSAGE_TOKEN = 200008
def parse_gpt_oss(
responses: Generator[GenerationResponse | None],
) -> Generator[GenerationResponse | ToolCallResponse | None]:
encoding = get_gpt_oss_encoding()
stream = StreamableParser(encoding, role=Role.ASSISTANT)
thinking = False
current_tool_name: str | None = None
tool_arg_parts: list[str] = []
for response in responses:
if response is None:
yield None
continue
token_id = response.token
try:
stream.process(token_id)
except HarmonyError as e:
logger.error(
f"HarmonyError on token_id={response.token} text={response.text!r}: {e}"
)
return
delta = stream.last_content_delta
ch = stream.current_channel
recipient = stream.current_recipient
effective_recipient = (
recipient
if (recipient is not None and recipient.startswith("functions."))
else None
)
if effective_recipient != current_tool_name:
if current_tool_name is not None:
tool_name = current_tool_name.removeprefix("functions.")
logger.info(f"parse_gpt_oss yielding tool call: name={tool_name!r}")
yield ToolCallResponse(
tool_calls=[
ToolCallItem(
name=tool_name,
arguments="".join(tool_arg_parts).strip(),
)
],
usage=response.usage,
)
tool_arg_parts = []
current_tool_name = effective_recipient
if current_tool_name is not None:
if delta:
tool_arg_parts.append(delta)
if response.finish_reason is not None:
yield response.model_copy(update={"text": "".join(tool_arg_parts)})
tool_arg_parts = []
continue
is_suppressed = ch == "analysis" or (
recipient is not None and recipient.startswith("!")
)
if is_suppressed and not thinking:
thinking = True
if not is_suppressed and thinking:
thinking = False
if delta:
yield response.model_copy(update={"text": delta, "is_thinking": thinking})
if response.finish_reason is not None:
yield response.model_copy(update={"text": ""})
def parse_deepseek_v32(
responses: Generator[GenerationResponse | None],
) -> Generator[GenerationResponse | ToolCallResponse | None]:
"""Parse DeepSeek V3.2 DSML tool calls from the generation stream.
Uses accumulated-text matching (not per-token marker checks) because
DSML markers like <|DSML|function_calls> may span multiple tokens.
Also handles <think>...</think> blocks for thinking mode.
"""
from mlx_engine.dsml_encoding import (
THINKING_END,
THINKING_START,
TOOL_CALLS_END,
TOOL_CALLS_START,
parse_dsml_output,
)
accumulated = ""
in_tool_call = False
thinking = False
# Tokens buffered while we detect the start of a DSML block
pending_buffer: list[GenerationResponse] = []
# Text accumulated during a tool call block
tool_call_text = ""
for response in responses:
if response is None:
yield None
continue
# ── Handle thinking tags ──
if not thinking and THINKING_START in response.text:
thinking = True
# Yield any text before the <think> tag
before = response.text[: response.text.index(THINKING_START)]
if before:
yield response.model_copy(update={"text": before})
continue
if thinking and THINKING_END in response.text:
thinking = False
# Yield any text after the </think> tag
after = response.text[
response.text.index(THINKING_END) + len(THINKING_END) :
]
if after:
yield response.model_copy(update={"text": after, "is_thinking": False})
continue
if thinking:
yield response.model_copy(update={"is_thinking": True})
continue
# ── Handle tool call accumulation ──
if in_tool_call:
tool_call_text += response.text
if TOOL_CALLS_END in tool_call_text:
# Parse the accumulated DSML block
parsed = parse_dsml_output(tool_call_text)
if parsed is not None:
logger.info(f"parsed DSML tool calls: {parsed}")
yield ToolCallResponse(
tool_calls=parsed,
usage=response.usage,
stats=response.stats,
)
else:
logger.warning(
f"DSML tool call parsing failed for: {tool_call_text}"
)
yield response.model_copy(update={"text": tool_call_text})
in_tool_call = False
tool_call_text = ""
continue
# EOS reached before end marker — yield buffered text as-is
if response.finish_reason is not None:
logger.info("DSML tool call parsing interrupted by EOS")
yield response.model_copy(update={"text": tool_call_text})
in_tool_call = False
tool_call_text = ""
continue
# ── Detect start of tool call block ──
accumulated += response.text
if TOOL_CALLS_START in accumulated:
# The start marker might be split across pending_buffer + current token
start_idx = accumulated.index(TOOL_CALLS_START)
# Yield any pending tokens that are purely before the marker
pre_text = accumulated[:start_idx]
if pre_text:
# Flush pending buffer tokens that contributed text before the marker
for buf_resp in pending_buffer:
if pre_text:
chunk = buf_resp.text
if len(chunk) <= len(pre_text):
yield buf_resp
pre_text = pre_text[len(chunk) :]
else:
yield buf_resp.model_copy(update={"text": pre_text})
pre_text = ""
pending_buffer = []
tool_call_text = accumulated[start_idx:]
accumulated = ""
# Check if the end marker is already present (entire tool call in one token)
if TOOL_CALLS_END in tool_call_text:
parsed = parse_dsml_output(tool_call_text)
if parsed is not None:
logger.info(f"parsed DSML tool calls: {parsed}")
yield ToolCallResponse(
tool_calls=parsed,
usage=response.usage,
stats=response.stats,
)
else:
logger.warning(
f"DSML tool call parsing failed for: {tool_call_text}"
)
yield response.model_copy(update={"text": tool_call_text})
tool_call_text = ""
else:
in_tool_call = True
continue
# Check if accumulated text might be the start of a DSML marker
# Buffer tokens if we see a partial match at the end
if _could_be_dsml_prefix(accumulated):
pending_buffer.append(response)
continue
# No partial match — flush all pending tokens and the current one
for buf_resp in pending_buffer:
yield buf_resp
pending_buffer = []
accumulated = ""
yield response
# Flush any remaining pending buffer at generator end
for buf_resp in pending_buffer:
yield buf_resp
def _could_be_dsml_prefix(text: str) -> bool:
"""Check if the end of text could be the start of a DSML function_calls marker.
We look for suffixes of text that are prefixes of the TOOL_CALLS_START pattern.
This allows us to buffer tokens until we can determine if a tool call is starting.
"""
from mlx_engine.dsml_encoding import TOOL_CALLS_START
# Only check the last portion of text that could overlap with the marker
max_check = len(TOOL_CALLS_START)
tail = text[-max_check:] if len(text) > max_check else text
# Check if any suffix of tail is a prefix of TOOL_CALLS_START
for i in range(len(tail)):
suffix = tail[i:]
if TOOL_CALLS_START.startswith(suffix):
return True
return False
def parse_thinking_models(
responses: Generator[GenerationResponse | None],
think_start: str | None,
think_end: str | None,
starts_in_thinking: bool = True,
) -> Generator[GenerationResponse | None]:
"""Route thinking tokens via is_thinking flag.
Swallows think tag tokens, sets is_thinking on all others.
Always yields tokens with finish_reason to avoid hanging the chunk stream.
"""
is_thinking = starts_in_thinking
for response in responses:
if response is None:
yield None
continue
if response.finish_reason is not None:
yield response.model_copy(update={"is_thinking": False})
continue
if response.text == think_start:
is_thinking = True
continue
if response.text == think_end:
is_thinking = False
continue
yield response.model_copy(update={"is_thinking": is_thinking})
def parse_tool_calls(
responses: Generator[GenerationResponse | None],
tool_parser: ToolParser,
tools: list[dict[str, Any]] | None,
) -> Generator[GenerationResponse | ToolCallResponse | None]:
in_tool_call = False
tool_call_text_parts: list[str] = []
for response in responses:
if response is None:
yield None
continue
if not in_tool_call and response.text.startswith(tool_parser.start_parsing):
in_tool_call = True
if not in_tool_call:
yield response
continue
tool_call_text_parts.append(response.text)
if response.text.endswith(tool_parser.end_parsing):
# parse the actual tool calls from the tool call text
combined = "".join(tool_call_text_parts)
parsed = tool_parser.parse(combined.strip(), tools=tools)
logger.info(f"parsed {tool_call_text_parts=} into {parsed=}")
in_tool_call = False
tool_call_text_parts = []
if parsed is None:
logger.warning(f"tool call parsing failed for text {combined}")
yield response.model_copy(update={"text": combined})
continue
yield ToolCallResponse(
tool_calls=parsed, usage=response.usage, stats=response.stats
)
continue
if response.finish_reason is not None:
logger.info(
"tool call parsing interrupted, yield partial tool call as text"
)
response = response.model_copy(
update={
"text": "".join(tool_call_text_parts),
"token": 0,
}
)
yield response
+1 -1
View File
@@ -1,8 +1,8 @@
from collections.abc import Generator
from typing import Any, Literal
from exo_core.model_cards import ModelId
from exo_core.models import TaggedModel
from exo_core.types.common import ModelId
from exo_core.types.runner_response import (
FinishReason,
GenerationStats,
@@ -5,7 +5,8 @@ from pydantic import model_validator
from exo_core.model_cards import ModelTask
from exo_core.models import CamelCaseModel, TaggedModel
from exo_core.types.common import Host, Id, NodeId
from exo_core.types.runners import RunnerId, ShardAssignments, ShardMetadata
from exo_core.types.runners import RunnerId, ShardAssignments
from exo_core.types.shards import ShardMetadata
class InstanceId(Id):
@@ -2,9 +2,8 @@ from collections.abc import Mapping
from pydantic import model_validator
from exo_core.model_cards import ModelId
from exo_core.models import CamelCaseModel, TaggedModel
from exo_core.types.common import Id, NodeId
from exo_core.types.common import Id, ModelId, NodeId
from exo_core.types.shards import ShardMetadata