mirror of
https://github.com/exo-explore/exo.git
synced 2026-09-09 12:02:25 -04:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
efa0da0912 |
No files matched your search
@@ -32,6 +32,7 @@ jobs:
|
||||
SPARKLE_ED25519_PRIVATE: ${{ secrets.SPARKLE_ED25519_PRIVATE }}
|
||||
SPARKLE_S3_BUCKET: ${{ secrets.SPARKLE_S3_BUCKET }}
|
||||
SPARKLE_S3_PREFIX: ${{ secrets.SPARKLE_S3_PREFIX }}
|
||||
EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT: ${{ secrets.EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT }}
|
||||
AWS_REGION: ${{ secrets.AWS_REGION }}
|
||||
EXO_BUILD_NUMBER: ${{ github.run_number }}
|
||||
EXO_LIBP2P_NAMESPACE: ${{ github.ref_name }}
|
||||
@@ -346,6 +347,7 @@ jobs:
|
||||
EXO_BUILD_COMMIT="$GITHUB_SHA" \
|
||||
SPARKLE_FEED_URL="$SPARKLE_FEED_URL" \
|
||||
SPARKLE_ED25519_PUBLIC="$SPARKLE_ED25519_PUBLIC" \
|
||||
EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT="$EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT" \
|
||||
CODE_SIGNING_IDENTITY="$SIGNING_IDENTITY" \
|
||||
CODE_SIGN_INJECT_BASE_ENTITLEMENTS=YES
|
||||
mkdir -p ../../output
|
||||
|
||||
@@ -1767,12 +1767,12 @@ def clip(
|
||||
array: The clipped array.
|
||||
"""
|
||||
|
||||
def compile[F: Callable[..., object]](
|
||||
fun: F,
|
||||
def compile(
|
||||
fun: Callable,
|
||||
inputs: object | None = ...,
|
||||
outputs: object | None = ...,
|
||||
shapeless: bool = ...,
|
||||
) -> F:
|
||||
) -> Callable:
|
||||
"""
|
||||
Returns a compiled function which produces the same output as ``fun``.
|
||||
|
||||
@@ -2915,8 +2915,8 @@ def gather_mm(
|
||||
a: array,
|
||||
b: array,
|
||||
/,
|
||||
lhs_indices: array | None = ...,
|
||||
rhs_indices: array | None = ...,
|
||||
lhs_indices: array,
|
||||
rhs_indices: array,
|
||||
*,
|
||||
sorted_indices: bool = ...,
|
||||
stream: Stream | Device | None = ...,
|
||||
@@ -4707,7 +4707,6 @@ def softmax(
|
||||
/,
|
||||
axis: int | Sequence[int] | None = ...,
|
||||
*,
|
||||
precise: bool = ...,
|
||||
stream: Stream | Device | None = ...,
|
||||
) -> array:
|
||||
"""
|
||||
|
||||
@@ -57,10 +57,6 @@ class Module(dict):
|
||||
def __init__(self) -> None:
|
||||
"""Should be called by the subclasses of ``Module``."""
|
||||
|
||||
def __getitem__(self, key: str) -> mx.array | Module: ...
|
||||
def get(
|
||||
self, key: str, default: mx.array | Module | None = ...
|
||||
) -> mx.array | Module | None: ...
|
||||
@property
|
||||
def training(self): # -> bool:
|
||||
"""Boolean indicating if the model is in training mode."""
|
||||
|
||||
@@ -3,7 +3,7 @@ This type stub file was generated by pyright.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
from typing import Optional
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
@@ -37,10 +37,10 @@ def quantized_scaled_dot_product_attention(
|
||||
bits: int = ...,
|
||||
) -> mx.array: ...
|
||||
def scaled_dot_product_attention(
|
||||
queries: mx.array,
|
||||
keys: mx.array,
|
||||
values: mx.array,
|
||||
cache: Optional[Any],
|
||||
queries,
|
||||
keys,
|
||||
values,
|
||||
cache,
|
||||
scale: float,
|
||||
mask: Optional[mx.array],
|
||||
sinks: Optional[mx.array] = ...,
|
||||
|
||||
@@ -1,103 +0,0 @@
|
||||
"""Type stubs for mlx_lm.models.gpt_oss"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
from .base import BaseModelArgs
|
||||
from .cache import KVCache
|
||||
from .switch_layers import SwitchGLU
|
||||
|
||||
@dataclass
|
||||
class ModelArgs(BaseModelArgs):
|
||||
model_type: str
|
||||
hidden_size: int
|
||||
intermediate_size: int
|
||||
num_hidden_layers: int
|
||||
num_attention_heads: int
|
||||
num_key_value_heads: int
|
||||
num_local_experts: int
|
||||
num_experts_per_tok: int
|
||||
vocab_size: int
|
||||
rms_norm_eps: float
|
||||
sliding_window: int
|
||||
layer_types: Optional[List[str]]
|
||||
|
||||
def mlx_topk(a: mx.array, k: int, axis: int = -1) -> tuple[mx.array, mx.array]: ...
|
||||
|
||||
class AttentionBlock(nn.Module):
|
||||
head_dim: int
|
||||
num_attention_heads: int
|
||||
num_key_value_heads: int
|
||||
num_key_value_groups: int
|
||||
sinks: mx.array
|
||||
q_proj: nn.Linear
|
||||
k_proj: nn.Linear
|
||||
v_proj: nn.Linear
|
||||
o_proj: nn.Linear
|
||||
sm_scale: float
|
||||
rope: nn.Module
|
||||
|
||||
def __init__(self, config: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
x: mx.array,
|
||||
mask: Optional[mx.array] = None,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
|
||||
class TransformerBlock(nn.Module):
|
||||
self_attn: AttentionBlock
|
||||
mlp: MLPBlock
|
||||
|
||||
def __init__(self, config: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
x: mx.array,
|
||||
mask: Optional[mx.array] = None,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
|
||||
class MLPBlock(nn.Module):
|
||||
hidden_size: int
|
||||
num_local_experts: int
|
||||
num_experts_per_tok: int
|
||||
experts: SwitchGLU
|
||||
router: nn.Linear
|
||||
sharding_group: Optional[mx.distributed.Group]
|
||||
|
||||
def __init__(self, config: ModelArgs) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
|
||||
class GptOssMoeModel(nn.Module):
|
||||
embed_tokens: nn.Embedding
|
||||
norm: nn.RMSNorm
|
||||
layer_types: List[str]
|
||||
layers: list[TransformerBlock]
|
||||
window_size: int
|
||||
swa_idx: int
|
||||
ga_idx: int
|
||||
|
||||
def __init__(self, args: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
inputs: mx.array,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
|
||||
class Model(nn.Module):
|
||||
model_type: str
|
||||
model: GptOssMoeModel
|
||||
lm_head: nn.Linear
|
||||
|
||||
def __init__(self, args: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
inputs: mx.array,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
@property
|
||||
def layers(self) -> list[nn.Module]: ...
|
||||
def make_cache(self) -> list[KVCache]: ...
|
||||
@@ -1,94 +0,0 @@
|
||||
"""Type stubs for mlx_lm.models.minimax"""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
|
||||
from .base import BaseModelArgs
|
||||
from .switch_layers import SwitchGLU
|
||||
|
||||
@dataclass
|
||||
class ModelArgs(BaseModelArgs):
|
||||
model_type: str
|
||||
hidden_size: int
|
||||
intermediate_size: int
|
||||
num_hidden_layers: int
|
||||
num_attention_heads: int
|
||||
num_key_value_heads: int
|
||||
num_local_experts: int
|
||||
num_experts_per_tok: int
|
||||
max_position_embeddings: int
|
||||
|
||||
class MiniMaxAttention(nn.Module):
|
||||
num_heads: int
|
||||
num_attention_heads: int
|
||||
num_key_value_heads: int
|
||||
head_dim: int
|
||||
scale: float
|
||||
q_proj: nn.Linear
|
||||
k_proj: nn.Linear
|
||||
v_proj: nn.Linear
|
||||
o_proj: nn.Linear
|
||||
q_norm: nn.Module
|
||||
k_norm: nn.Module
|
||||
rope: nn.Module
|
||||
|
||||
def __init__(self, args: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
x: mx.array,
|
||||
mask: Optional[mx.array] = None,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
|
||||
class MiniMaxSparseMoeBlock(nn.Module):
|
||||
num_experts_per_tok: int
|
||||
gate: nn.Linear
|
||||
switch_mlp: SwitchGLU
|
||||
e_score_correction_bias: mx.array
|
||||
sharding_group: Optional[mx.distributed.Group]
|
||||
|
||||
def __init__(self, args: ModelArgs) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
|
||||
class MiniMaxDecoderLayer(nn.Module):
|
||||
self_attn: MiniMaxAttention
|
||||
block_sparse_moe: MiniMaxSparseMoeBlock
|
||||
input_layernorm: nn.RMSNorm
|
||||
post_attention_layernorm: nn.RMSNorm
|
||||
|
||||
def __init__(self, args: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
x: mx.array,
|
||||
mask: Optional[mx.array] = None,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
|
||||
class MiniMaxModel(nn.Module):
|
||||
embed_tokens: nn.Embedding
|
||||
layers: list[MiniMaxDecoderLayer]
|
||||
norm: nn.RMSNorm
|
||||
|
||||
def __init__(self, args: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
inputs: mx.array,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
|
||||
class Model(nn.Module):
|
||||
model_type: str
|
||||
model: MiniMaxModel
|
||||
lm_head: nn.Linear
|
||||
|
||||
def __init__(self, args: ModelArgs) -> None: ...
|
||||
def __call__(
|
||||
self,
|
||||
inputs: mx.array,
|
||||
cache: Optional[Any] = None,
|
||||
) -> mx.array: ...
|
||||
@property
|
||||
def layers(self) -> list[MiniMaxDecoderLayer]: ...
|
||||
@@ -92,15 +92,6 @@ class NemotronHAttention(nn.Module):
|
||||
cache: Optional[KVCache] = None,
|
||||
) -> mx.array: ...
|
||||
|
||||
class MoEGate(nn.Module):
|
||||
config: ModelArgs
|
||||
top_k: int
|
||||
norm_topk_prob: bool
|
||||
weight: mx.array
|
||||
|
||||
def __init__(self, config: ModelArgs) -> None: ...
|
||||
def __call__(self, x: mx.array) -> tuple[mx.array, mx.array]: ...
|
||||
|
||||
class NemotronHMLP(nn.Module):
|
||||
up_proj: nn.Linear
|
||||
down_proj: nn.Linear
|
||||
@@ -111,14 +102,9 @@ class NemotronHMLP(nn.Module):
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
|
||||
class NemotronHMoE(nn.Module):
|
||||
config: ModelArgs
|
||||
num_experts_per_tok: int
|
||||
moe_latent_size: Optional[int]
|
||||
switch_mlp: SwitchMLP
|
||||
gate: MoEGate
|
||||
shared_experts: NemotronHMLP
|
||||
fc1_latent_proj: nn.Linear
|
||||
fc2_latent_proj: nn.Linear
|
||||
|
||||
def __init__(self, config: ModelArgs) -> None: ...
|
||||
def __call__(self, x: mx.array) -> mx.array: ...
|
||||
|
||||
@@ -71,7 +71,6 @@ class Qwen3NextAttention(nn.Module):
|
||||
class Qwen3NextSparseMoeBlock(nn.Module):
|
||||
norm_topk_prob: bool
|
||||
num_experts: int
|
||||
num_experts_per_tok: int
|
||||
top_k: int
|
||||
gate: nn.Linear
|
||||
switch_mlp: SwitchGLU
|
||||
|
||||
@@ -584,18 +584,9 @@ struct ContentView: View {
|
||||
|
||||
case .prompting:
|
||||
VStack(alignment: .leading, spacing: 6) {
|
||||
VStack(alignment: .leading, spacing: 2) {
|
||||
Text("Tell us what went wrong (optional)")
|
||||
.font(.caption2)
|
||||
.foregroundColor(.secondary)
|
||||
Text(
|
||||
"A quick description of what you were doing and what happened helps us track down the bug for you."
|
||||
)
|
||||
Text("What's the issue? (optional)")
|
||||
.font(.caption2)
|
||||
.foregroundColor(.secondary)
|
||||
.opacity(0.8)
|
||||
.fixedSize(horizontal: false, vertical: true)
|
||||
}
|
||||
TextEditor(text: $bugReportUserDescription)
|
||||
.font(.caption2)
|
||||
.frame(height: 60)
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
<key>EXOBuildCommit</key>
|
||||
<string>$(EXO_BUILD_COMMIT)</string>
|
||||
<key>EXOBugReportPresignedUrlEndpoint</key>
|
||||
<string>https://reports.exolabs.net/presigned-urls</string>
|
||||
<string>$(EXO_BUG_REPORT_PRESIGNED_URL_ENDPOINT)</string>
|
||||
<key>NSLocalNetworkUsageDescription</key>
|
||||
<string>EXO needs local network access to discover and connect to other devices in your cluster for distributed AI inference.</string>
|
||||
<key>NSBonjourServices</key>
|
||||
|
||||
+37
-82
@@ -3,13 +3,11 @@ from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import tomllib
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
@@ -211,7 +209,7 @@ def _openai_build_request(
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"tools": tools,
|
||||
"max_tokens": 4096,
|
||||
"max_tokens": 16384,
|
||||
"temperature": 0.0,
|
||||
}
|
||||
return "/v1/chat/completions", body
|
||||
@@ -278,7 +276,7 @@ def _openai_build_followup(
|
||||
"model": model,
|
||||
"messages": followup_messages,
|
||||
"tools": tools,
|
||||
"max_tokens": 4096,
|
||||
"max_tokens": 16384,
|
||||
"temperature": 0.0,
|
||||
}
|
||||
return "/v1/chat/completions", body
|
||||
@@ -381,7 +379,7 @@ def _claude_build_request(
|
||||
"model": model,
|
||||
"messages": claude_messages,
|
||||
"tools": claude_tools,
|
||||
"max_tokens": 4096,
|
||||
"max_tokens": 16384,
|
||||
"temperature": 0.0,
|
||||
}
|
||||
if system_content is not None:
|
||||
@@ -491,7 +489,7 @@ def _claude_build_followup(
|
||||
"model": model,
|
||||
"messages": claude_messages,
|
||||
"tools": claude_tools,
|
||||
"max_tokens": 4096,
|
||||
"max_tokens": 16384,
|
||||
"temperature": 0.0,
|
||||
}
|
||||
if system_content is not None:
|
||||
@@ -915,12 +913,6 @@ Examples:
|
||||
default=1,
|
||||
help="Repeat each scenario N times (default: 1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--concurrency",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Run up to N scenarios in parallel against the same instance (default: 1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--scenarios",
|
||||
nargs="*",
|
||||
@@ -943,13 +935,6 @@ Examples:
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.concurrency < 1:
|
||||
print(
|
||||
f"--concurrency must be >= 1 (got {args.concurrency})",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(2)
|
||||
|
||||
all_scenarios = load_scenarios(SCENARIOS_PATH)
|
||||
if args.scenarios:
|
||||
scenarios = [s for s in all_scenarios if s.name in args.scenarios]
|
||||
@@ -1025,72 +1010,42 @@ Examples:
|
||||
cluster_snapshot = capture_cluster_snapshot(exo)
|
||||
all_results: list[ScenarioResult] = []
|
||||
|
||||
tasks: list[tuple[int, Scenario, ApiName]] = [
|
||||
(run_idx, scenario, api_name)
|
||||
for run_idx in range(args.repeat)
|
||||
for scenario in scenarios
|
||||
for api_name in api_names
|
||||
]
|
||||
|
||||
def _run_one(
|
||||
http_client: httpx.Client,
|
||||
task: tuple[int, Scenario, ApiName],
|
||||
) -> tuple[tuple[int, Scenario, ApiName], list[ScenarioResult], str]:
|
||||
run_idx, scenario, api_name = task
|
||||
buf = io.StringIO()
|
||||
run_tag = f"[run {run_idx + 1}/{args.repeat}]" if args.repeat > 1 else ""
|
||||
print(
|
||||
f"\n {run_tag}[{api_name:>9}] {scenario.name}: {scenario.description}",
|
||||
file=buf,
|
||||
)
|
||||
scenario_results = run_scenario(
|
||||
http_client,
|
||||
args.host,
|
||||
args.port,
|
||||
full_model_id,
|
||||
scenario,
|
||||
api_name,
|
||||
args.timeout,
|
||||
args.verbose,
|
||||
)
|
||||
for r in scenario_results:
|
||||
status = "PASS" if r.passed else "FAIL"
|
||||
print(
|
||||
f" [{r.phase:>10}] {status} ({r.latency_ms:.0f}ms)",
|
||||
file=buf,
|
||||
)
|
||||
for check_name, check_ok in r.checks.items():
|
||||
mark = "+" if check_ok else "-"
|
||||
print(f" {mark} {check_name}", file=buf)
|
||||
if r.error:
|
||||
print(f" ! {r.error}", file=buf)
|
||||
return task, scenario_results, buf.getvalue()
|
||||
|
||||
try:
|
||||
with httpx.Client() as http_client:
|
||||
if args.concurrency == 1:
|
||||
current_run = -1
|
||||
for task in tasks:
|
||||
run_idx = task[0]
|
||||
if args.repeat > 1 and run_idx != current_run:
|
||||
print(f"\n--- Run {run_idx + 1}/{args.repeat} ---", file=log)
|
||||
current_run = run_idx
|
||||
_, scenario_results, buffered = _run_one(http_client, task)
|
||||
all_results.extend(scenario_results)
|
||||
log.write(buffered)
|
||||
log.flush()
|
||||
else:
|
||||
print(
|
||||
f"Running {len(tasks)} tasks with concurrency={args.concurrency}",
|
||||
file=log,
|
||||
)
|
||||
with ThreadPoolExecutor(max_workers=args.concurrency) as pool:
|
||||
futures = [pool.submit(_run_one, http_client, t) for t in tasks]
|
||||
for fut in as_completed(futures):
|
||||
_, scenario_results, buffered = fut.result()
|
||||
for run_idx in range(args.repeat):
|
||||
if args.repeat > 1:
|
||||
print(f"\n--- Run {run_idx + 1}/{args.repeat} ---", file=log)
|
||||
|
||||
for scenario in scenarios:
|
||||
for api_name in api_names:
|
||||
print(
|
||||
f"\n [{api_name:>9}] {scenario.name}: {scenario.description}",
|
||||
file=log,
|
||||
)
|
||||
|
||||
scenario_results = run_scenario(
|
||||
http_client,
|
||||
args.host,
|
||||
args.port,
|
||||
full_model_id,
|
||||
scenario,
|
||||
api_name,
|
||||
args.timeout,
|
||||
args.verbose,
|
||||
)
|
||||
all_results.extend(scenario_results)
|
||||
log.write(buffered)
|
||||
log.flush()
|
||||
|
||||
for r in scenario_results:
|
||||
status = "PASS" if r.passed else "FAIL"
|
||||
print(
|
||||
f" [{r.phase:>10}] {status} ({r.latency_ms:.0f}ms)",
|
||||
file=log,
|
||||
)
|
||||
for check_name, check_ok in r.checks.items():
|
||||
mark = "+" if check_ok else "-"
|
||||
print(f" {mark} {check_name}", file=log)
|
||||
if r.error:
|
||||
print(f" ! {r.error}", file=log)
|
||||
finally:
|
||||
try:
|
||||
exo.request_json("DELETE", f"/instance/{instance_id}")
|
||||
|
||||
+1
-1
@@ -564,7 +564,7 @@ def add_common_instance_args(ap: argparse.ArgumentParser) -> None:
|
||||
ap.add_argument(
|
||||
"--settle-timeout",
|
||||
type=float,
|
||||
default=60.0,
|
||||
default=0,
|
||||
help="Max seconds to wait for the cluster to produce valid placements (0 = try once).",
|
||||
)
|
||||
ap.add_argument(
|
||||
|
||||
@@ -88,12 +88,10 @@
|
||||
let codexModel = $state("");
|
||||
let codexMcpPath = $state("/Users/username");
|
||||
let openClawModel = $state("");
|
||||
let piModel = $state("");
|
||||
$effect(() => {
|
||||
const def = modelsBySize.length > 0 ? modelsBySize[0] : "your-model-id";
|
||||
codexModel = def;
|
||||
openClawModel = def;
|
||||
piModel = def;
|
||||
});
|
||||
|
||||
const claudeShellCommand = $derived(
|
||||
@@ -220,55 +218,6 @@
|
||||
),
|
||||
);
|
||||
|
||||
const piModelsJson = $derived.by(() => {
|
||||
const models: Record<string, unknown>[] = [];
|
||||
for (const modelId of runningModels) {
|
||||
const caps = modelCapabilities[modelId] || [];
|
||||
const ctxLen = modelContextLengths[modelId] || 0;
|
||||
const entry: Record<string, unknown> = { id: modelId };
|
||||
if (caps.includes("vision")) {
|
||||
entry.input = ["text", "image"];
|
||||
}
|
||||
// Mark thinking-capable models so pi surfaces its thinking-level selector
|
||||
// for them. exo capability strings: "thinking" (model emits reasoning
|
||||
// content) and "thinking_toggle" (user can turn it on/off).
|
||||
if (caps.includes("thinking") || caps.includes("thinking_toggle")) {
|
||||
entry.reasoning = true;
|
||||
}
|
||||
if (ctxLen > 0) {
|
||||
entry.contextWindow = ctxLen;
|
||||
}
|
||||
models.push(entry);
|
||||
}
|
||||
if (models.length === 0) {
|
||||
models.push({ id: "your-model-id" });
|
||||
}
|
||||
return JSON.stringify(
|
||||
{
|
||||
providers: {
|
||||
exo: {
|
||||
baseUrl: `${apiUrl}/v1`,
|
||||
api: "openai-completions",
|
||||
apiKey: "exo",
|
||||
compat: {
|
||||
supportsDeveloperRole: false,
|
||||
// exo's OpenAI surface takes a boolean `enable_thinking` toggle,
|
||||
// not graded effort levels, so disable pi's `reasoning_effort`
|
||||
// parameter and use the matching top-level-boolean format.
|
||||
supportsReasoningEffort: false,
|
||||
thinkingFormat: "qwen",
|
||||
},
|
||||
models,
|
||||
},
|
||||
},
|
||||
},
|
||||
null,
|
||||
2,
|
||||
);
|
||||
});
|
||||
|
||||
const piShellCommand = $derived(`pi --provider exo --model ${piModel}`);
|
||||
|
||||
const ollamaCommand = $derived(
|
||||
`OLLAMA_HOST=${apiUrl}/ollama ollama run ${modelsBySize.length > 0 ? modelsBySize[0] : "your-model-id"}`,
|
||||
);
|
||||
@@ -328,7 +277,6 @@
|
||||
"OpenCode",
|
||||
"Codex",
|
||||
"OpenClaw",
|
||||
"Pi",
|
||||
"Open WebUI",
|
||||
"n8n",
|
||||
"Firefox",
|
||||
@@ -567,33 +515,6 @@
|
||||
config={`openclaw doctor --fix${(modelCapabilities[openClawModel] || []).includes("vision") ? `\nopenclaw models set-image exo/${openClawModel}` : ""}\nopenclaw gateway &\nopenclaw dashboard`}
|
||||
language="bash"
|
||||
/>
|
||||
{:else if activeTab === "Pi"}
|
||||
{#if runningModels.length > 1}
|
||||
<div class="text-xs">
|
||||
<span
|
||||
class="text-exo-light-gray/50 text-[10px] uppercase tracking-wider block mb-1"
|
||||
>Model</span
|
||||
>
|
||||
<select bind:value={piModel} class={selectClass}>
|
||||
{#each runningModels as model}
|
||||
<option value={model}>{model.split("/").pop()}</option>
|
||||
{/each}
|
||||
</select>
|
||||
</div>
|
||||
{/if}
|
||||
<IntegrationCard
|
||||
title="Models Config"
|
||||
subtitle="~/.pi/agent/models.json"
|
||||
description="Register exo as a custom provider in pi. Create or edit this file, then run pi and pick an exo model via /model. Install pi with: npm install -g @mariozechner/pi-coding-agent"
|
||||
config={piModelsJson}
|
||||
/>
|
||||
<IntegrationCard
|
||||
title="Shell Command"
|
||||
subtitle="Run in terminal"
|
||||
description="Launch pi directly with the exo provider and model selected."
|
||||
config={piShellCommand}
|
||||
language="bash"
|
||||
/>
|
||||
{:else if activeTab === "Open WebUI"}
|
||||
<IntegrationCard
|
||||
title="1. Start Open WebUI"
|
||||
|
||||
Generated
-6
@@ -1,6 +0,0 @@
|
||||
{
|
||||
"name": "exo",
|
||||
"lockfileVersion": 3,
|
||||
"requires": true,
|
||||
"packages": {}
|
||||
}
|
||||
+5
-14
@@ -3,7 +3,7 @@ name = "exo"
|
||||
version = "0.3.70"
|
||||
description = "Exo"
|
||||
readme = "README.md"
|
||||
requires-python = "==3.13.*"
|
||||
requires-python = ">=3.13"
|
||||
dependencies = [
|
||||
"aiofiles>=24.1.0",
|
||||
"aiohttp>=3.12.14",
|
||||
@@ -17,8 +17,8 @@ dependencies = [
|
||||
"loguru>=0.7.3",
|
||||
"exo-pyo3-bindings", # rust bindings
|
||||
"anyio==4.11.0",
|
||||
"mlx==0.31.2; sys_platform == 'darwin'",
|
||||
"mlx-lm; sys_platform=='darwin'",
|
||||
"mlx==0.31.1; sys_platform == 'darwin'",
|
||||
"mlx-lm",
|
||||
"tiktoken>=0.12.0", # required for kimi k2 tokenizer
|
||||
"hypercorn>=0.18.0",
|
||||
"openai-harmony>=0.0.8",
|
||||
@@ -30,6 +30,7 @@ dependencies = [
|
||||
"zstandard>=0.23.0",
|
||||
"mlx-vlm>=0.3.11",
|
||||
"transformers>=5.0.0,<5.4.0",
|
||||
"pydantic-settings>=2.13.1",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
@@ -49,21 +50,15 @@ dev = [
|
||||
[project.optional-dependencies]
|
||||
build = ["nanobind"]
|
||||
cpu = [
|
||||
"mlx==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-cpu==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-lm; sys_platform == 'linux'",
|
||||
"torch>=2.10.0; sys_platform == 'linux'",
|
||||
]
|
||||
cuda12 = [
|
||||
"mlx==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-cuda-12==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-lm; sys_platform == 'linux'",
|
||||
"torch>=2.10.0; sys_platform == 'linux'",
|
||||
]
|
||||
cuda13 = [
|
||||
"mlx==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-cuda-13==0.31.1; sys_platform == 'linux'",
|
||||
"mlx-lm; sys_platform == 'linux'",
|
||||
"torch>=2.10.0; sys_platform == 'linux'",
|
||||
]
|
||||
|
||||
@@ -81,7 +76,7 @@ mlx-lm = { git = "https://github.com/rltakashige/mlx-lm", branch = "leo/fix-arra
|
||||
torch = [
|
||||
{ index = "pytorch-cu130", marker = "sys_platform == 'linux' and extra == 'cuda13' and extra != 'cpu' and extra != 'cuda12'" },
|
||||
{ index = "pytorch-cu120", marker = "sys_platform == 'linux' and extra == 'cuda12' and extra != 'cpu' and extra != 'cuda13'" },
|
||||
{ index = "pytorch-cpu", marker = "(extra != 'cuda12' and extra != 'cuda13' and sys_platform == 'linux') or sys_platform == 'darwin'" },
|
||||
{ index = "pytorch-cpu", marker = "sys_platform == 'linux' and extra == 'cpu' and extra != 'cuda12' and extra != 'cuda13'" },
|
||||
]
|
||||
vllm = { git = "https://github.com/hmellor/vllm.git", branch = "transformers-v5" }
|
||||
|
||||
@@ -155,10 +150,6 @@ prerelease = "allow"
|
||||
environments = ["sys_platform == 'darwin'", "sys_platform == 'linux'"]
|
||||
conflicts = [[{ extra = "cuda12" }, { extra = "cuda13" }, { extra = "cpu" }]]
|
||||
constraint-dependencies = ["transformers>=5.0.0,<5.4.0"]
|
||||
override-dependencies = [
|
||||
"mlx==0.31.1; sys_platform=='linux'",
|
||||
"mlx; sys_platform=='darwin'",
|
||||
]
|
||||
|
||||
[tool.uv.extra-build-dependencies]
|
||||
miniaudio = ["setuptools", "cffi", "pycparser"]
|
||||
|
||||
+68
-62
@@ -5,13 +5,14 @@ let
|
||||
workspaceRoot = ../.;
|
||||
};
|
||||
|
||||
mkPythonSet = { pkgs, lib, self', members }:
|
||||
mkPythonSet = { pkgs, lib, self' }:
|
||||
let
|
||||
inherit (pkgs.stdenv.hostPlatform) isLinux isDarwin isx86_64;
|
||||
inherit (pkgs.config) cudaSupport;
|
||||
inherit (pkgs) cudaPackages;
|
||||
cuda13Support = cudaSupport && cudaPackages.cudaMajorVersion == "13";
|
||||
libmlx_source = if cuda13Support then "mlx-cuda-13" else if cudaSupport then "mlx-cuda-12" else "mlx-cpu";
|
||||
uv_extra = if cuda13Support then "cuda13" else if cudaSupport then "cuda12" else "cpu";
|
||||
python = pkgs.python313;
|
||||
cudaLibs = with cudaPackages; [
|
||||
cuda_cudart
|
||||
@@ -50,7 +51,7 @@ let
|
||||
'';
|
||||
};
|
||||
};
|
||||
buildSystemsOverlay = final: prev:
|
||||
buildSystemsOverlay = final: prev: { } //
|
||||
lib.optionalAttrs isDarwin
|
||||
{
|
||||
mlx = prev.mlx.overrideAttrs (old:
|
||||
@@ -80,7 +81,7 @@ let
|
||||
nativeBuildInputs = (old.nativeBuildInputs or [ ]) ++ [ pkgs.cmake self'.packages.metal-toolchain ];
|
||||
# TODO: non-sdk_26 support
|
||||
buildInputs = (old.buildInputs or [ ])
|
||||
++ [ gguf-tools pkgs.fmt pkgs.nlohmann_json pkgs.apple-sdk_26 ];
|
||||
++ [ gguf-tools pkgs.fmt pkgs.nlohmann_json pkgs.apple-sdk_26 ];
|
||||
patches = [
|
||||
(pkgs.replaceVars ../nix/darwin-build-fixes.patch {
|
||||
sdkVersion = pkgs.apple-sdk_26.version;
|
||||
@@ -112,42 +113,42 @@ let
|
||||
MACOSX_DEPLOYMENT_TARGET = pkgs.apple-sdk_26.version;
|
||||
});
|
||||
} // lib.optionalAttrs isLinux {
|
||||
mlx = prev.mlx.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ lib.optionals cudaSupport cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = lib.optionals cudaSupport [ "libcuda.so.1" ];
|
||||
postInstall = ''
|
||||
cp -r "${final.${libmlx_source}}/${final.python.sitePackages}/mlx" "$out/${final.python.sitePackages}/mlx/"
|
||||
'';
|
||||
});
|
||||
} // lib.optionalAttrs cudaSupport {
|
||||
"${libmlx_source}" = prev."${libmlx_source}".overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cufile = prev.nvidia-cufile.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ [ pkgs.rdma-core ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cusolver = prev.nvidia-cusolver.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-nvshmem-cu13 = prev.nvidia-nvshmem-cu13.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ [ pkgs.rdma-core pkgs.pmix pkgs.libfabric pkgs.ucx pkgs.openmpi ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cusparse = prev.nvidia-cusparse.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ [ cudaLibs ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
torch = prev.torch.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
};
|
||||
mlx = prev.mlx.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ lib.optionals cudaSupport cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = lib.optionals cudaSupport [ "libcuda.so.1" ];
|
||||
postInstall = (old.postInstall or "") + ''
|
||||
cp -r "${final.${libmlx_source}}/${final.python.sitePackages}/mlx" "$out/${final.python.sitePackages}/mlx/"
|
||||
'';
|
||||
});
|
||||
} // lib.optionalAttrs cudaSupport {
|
||||
"${libmlx_source}" = prev."${libmlx_source}".overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cufile = prev.nvidia-cufile.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ [ pkgs.rdma-core ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cusolver = prev.nvidia-cusolver.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-nvshmem-cu13 = prev.nvidia-nvshmem-cu13.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ [ pkgs.rdma-core pkgs.pmix pkgs.libfabric pkgs.ucx pkgs.openmpi ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
nvidia-cusparse = prev.nvidia-cusparse.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ [ cudaLibs ];
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
torch = prev.torch.overrideAttrs (old: {
|
||||
buildInputs = old.buildInputs ++ cudaLibs;
|
||||
autoPatchelfIgnoreMissingDeps = [ "libcuda.so.1" ];
|
||||
});
|
||||
};
|
||||
pyprojectOverlay = workspace.mkPyprojectOverlay {
|
||||
sourcePreference = "wheel";
|
||||
dependencies = members;
|
||||
dependencies = { exo = [ uv_extra ]; exo-bench = [ ]; };
|
||||
};
|
||||
editableOverlay = workspace.mkEditablePyprojectOverlay {
|
||||
# Use environment variable pointing to editable root directory
|
||||
@@ -164,8 +165,8 @@ let
|
||||
buildSystemsOverlay
|
||||
]
|
||||
);
|
||||
venv = name: (pythonSet.mkVirtualEnv "${name}-env" members).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" ]; });
|
||||
mkApp = cmd: name: pkgs.writeShellApplication {
|
||||
|
||||
mkApp = cmd: name: members: pkgs.writeShellApplication {
|
||||
inherit name;
|
||||
runtimeEnv = {
|
||||
EXO_DASHBOARD_DIR = self'.packages.dashboard;
|
||||
@@ -173,17 +174,17 @@ let
|
||||
};
|
||||
runtimeInputs = [
|
||||
# mlx and mlx-cuda ship clashing cmake files - we dont need them at runtime anyway
|
||||
(venv name)
|
||||
((pythonSet.mkVirtualEnv "${name}-env" members).overrideAttrs (_: { venvSkip = [ "lib/python${python.pythonVersion}/site-packages/mlx/share/cmake/*" ]; }))
|
||||
]
|
||||
++ lib.optionals isDarwin [ pkgs.macmon ];
|
||||
text = "exec " + lib.optionalString cudaSupport "${lib.getExe pkgs.nix-gl-host} " + cmd;
|
||||
};
|
||||
in
|
||||
{
|
||||
inherit venv;
|
||||
inherit pythonSet;
|
||||
editablePythonSet = pythonSet.overrideScope editableOverlay;
|
||||
mkPythonScript = path: mkApp ''python ${path} "$@"'';
|
||||
mkExo = mkApp ''exo "$@"'';
|
||||
mkPythonScript = members: name: path: mkApp ''python ${path} "$@"'' name members;
|
||||
mkExo = name: members: mkApp ''exo "$@"'' name members;
|
||||
};
|
||||
in
|
||||
{
|
||||
@@ -191,21 +192,16 @@ in
|
||||
{ self', pkgs, unfreePkgs, lib, ... }:
|
||||
let
|
||||
inherit (pkgs.stdenv.hostPlatform) isLinux;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; members = { exo = [ "cpu" ]; }; }) editablePythonSet mkExo;
|
||||
inherit (mkPythonSet { inherit self' pkgs lib; }) pythonSet editablePythonSet mkPythonScript mkExo;
|
||||
|
||||
exoVenv = pythonSet.mkVirtualEnv "exo-env" { exo = lib.optionals isLinux [ "cpu" ]; };
|
||||
|
||||
# Virtual environment with dev dependencies for testing
|
||||
testVenv = (mkPythonSet {
|
||||
inherit self' pkgs lib; members = {
|
||||
exo = [ "dev" "cpu" ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
testVenv = pythonSet.mkVirtualEnv "exo-test-env" {
|
||||
exo = [ "dev" ] ++ lib.optionals isLinux [ "cpu" ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
};
|
||||
}).venv "exo-test";
|
||||
|
||||
mkBenchScript = (mkPythonSet {
|
||||
inherit self' pkgs lib; members = {
|
||||
exo = [ "cpu" ];
|
||||
exo-bench = [ ]; # Include pytest, pytest-asyncio, pytest-env
|
||||
};
|
||||
}).mkPythonScript;
|
||||
mkBenchScript = mkPythonScript { exo-bench = [ ]; };
|
||||
|
||||
mkSimplePythonScript = name: path: pkgs.writeShellApplication {
|
||||
inherit name;
|
||||
@@ -216,7 +212,9 @@ in
|
||||
in
|
||||
{
|
||||
packages = {
|
||||
exo = mkExo "exo";
|
||||
exo = mkExo "exo" { exo = lib.optionals isLinux [ "cpu" ]; };
|
||||
# for devShell
|
||||
exo-venv = exoVenv;
|
||||
editableVenv = editablePythonSet.mkVirtualEnv "exo-dev-env" { exo = [ "dev" ]; };
|
||||
# for running tests in ci
|
||||
exo-test-env = testVenv;
|
||||
@@ -226,8 +224,8 @@ in
|
||||
# used by ./tests/run_exo_on.sh
|
||||
exo-get-all-models-on-cluster = mkSimplePythonScript "exo-get-all-models-on-cluster" (inputs.self + /tests/get_all_models_on_cluster.py);
|
||||
} // lib.optionalAttrs isLinux {
|
||||
exo-cuda-12 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; members = { exo = [ "cuda12" ]; }; }).mkExo "exo-cuda-12";
|
||||
exo-cuda-13 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; members = { exo = [ "cuda13" ]; }; }).mkExo "exo-cuda-13";
|
||||
exo-cuda-12 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_12) pkgs; }).mkExo "exo-cuda-12" { exo = [ "cuda12" ]; };
|
||||
exo-cuda-13 = (mkPythonSet { inherit self' lib; inherit (unfreePkgs.pkgsCuda.cudaPackages_13) pkgs; }).mkExo "exo-cuda-13" { exo = [ "cuda13" ]; };
|
||||
};
|
||||
|
||||
checks = {
|
||||
@@ -237,11 +235,19 @@ in
|
||||
touch $out
|
||||
'';
|
||||
|
||||
typecheck = pkgs.runCommand "typecheck" { nativeBuildInputs = [ testVenv ]; } ''
|
||||
cd ${inputs.self}
|
||||
basedpyright
|
||||
touch $out
|
||||
'';
|
||||
typecheck = pkgs.runCommand "typecheck"
|
||||
{
|
||||
nativeBuildInputs = [
|
||||
testVenv
|
||||
pkgs.basedpyright
|
||||
];
|
||||
}
|
||||
''
|
||||
cd ${inputs.self}
|
||||
export HOME=$TMPDIR
|
||||
basedpyright --pythonpath ${testVenv}/bin/python --project ${inputs.self}/pyproject.toml
|
||||
touch $out
|
||||
'';
|
||||
};
|
||||
};
|
||||
}
|
||||
@@ -1,21 +0,0 @@
|
||||
model_id = "mlx-community/GLM-5.1-DQ4plus-q8"
|
||||
n_layers = 78
|
||||
hidden_size = 6144
|
||||
num_key_value_heads = 64
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "glm"
|
||||
quantization = "8bit"
|
||||
base_model = "GLM-5.1"
|
||||
capabilities = ["text", "thinking"]
|
||||
|
||||
context_length = 202752
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 465173655552
|
||||
|
||||
# Source: https://huggingface.co/zai-org/GLM-5.1
|
||||
# Source: https://docs.z.ai/api-reference/llm/chat-completion
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
@@ -1,21 +0,0 @@
|
||||
model_id = "mlx-community/GLM-5.1-MXFP4-Q8"
|
||||
n_layers = 78
|
||||
hidden_size = 6144
|
||||
num_key_value_heads = 64
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "glm"
|
||||
quantization = "MXFP4-Q8"
|
||||
base_model = "GLM-5.1"
|
||||
capabilities = ["text", "thinking"]
|
||||
|
||||
context_length = 202752
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 405480321024
|
||||
|
||||
# Source: https://huggingface.co/zai-org/GLM-5.1
|
||||
# Source: https://docs.z.ai/api-reference/llm/chat-completion
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
@@ -1,21 +0,0 @@
|
||||
model_id = "mlx-community/GLM-5.1"
|
||||
n_layers = 78
|
||||
hidden_size = 6144
|
||||
num_key_value_heads = 64
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "glm"
|
||||
quantization = "bf16"
|
||||
base_model = "GLM-5.1"
|
||||
capabilities = ["text", "thinking"]
|
||||
|
||||
context_length = 202752
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 1487822475264
|
||||
|
||||
# Source: https://huggingface.co/zai-org/GLM-5.1
|
||||
# Source: https://docs.z.ai/api-reference/llm/chat-completion
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
@@ -1,33 +0,0 @@
|
||||
model_id = "mlx-community/Kimi-K2.6-mlx-DQ3_K_M-q8"
|
||||
n_layers = 61
|
||||
hidden_size = 7168
|
||||
num_key_value_heads = 64
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "kimi"
|
||||
quantization = "3bit"
|
||||
base_model = "Kimi K2.6"
|
||||
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
|
||||
|
||||
context_length = 262144
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 470628683776
|
||||
|
||||
[vision]
|
||||
image_token_id = 163605
|
||||
model_type = "kimi_vl"
|
||||
weights_repo = "exolabs/Kimi-K2.6-vision"
|
||||
processor_repo = "moonshotai/Kimi-K2.6"
|
||||
|
||||
# Source: https://huggingface.co/moonshotai/Kimi-K2.6
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
min_p = 0.01
|
||||
|
||||
# Source: https://huggingface.co/moonshotai/Kimi-K2.6
|
||||
[sampling_defaults.non_thinking]
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
min_p = 0.01
|
||||
@@ -1,35 +0,0 @@
|
||||
model_id = "mlx-community/Qwen3.6-27B-4bit"
|
||||
n_layers = 64
|
||||
hidden_size = 5120
|
||||
num_key_value_heads = 4
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "qwen"
|
||||
quantization = "4bit"
|
||||
base_model = "Qwen3.6 27B"
|
||||
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
|
||||
|
||||
context_length = 262144
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 16054262240
|
||||
|
||||
# Source: https://huggingface.co/Qwen/Qwen3.6-27B#best-practices
|
||||
# Source: https://unsloth.ai/docs/models/qwen3.5
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
|
||||
# Source: https://huggingface.co/Qwen/Qwen3.6-27B#best-practices
|
||||
# Source: https://unsloth.ai/docs/models/qwen3.5
|
||||
[sampling_defaults.non_thinking]
|
||||
temperature = 0.7
|
||||
top_p = 0.8
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
@@ -1,35 +0,0 @@
|
||||
model_id = "mlx-community/Qwen3.6-27B-8bit"
|
||||
n_layers = 64
|
||||
hidden_size = 5120
|
||||
num_key_value_heads = 4
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "qwen"
|
||||
quantization = "8bit"
|
||||
base_model = "Qwen3.6 27B"
|
||||
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
|
||||
|
||||
context_length = 262144
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 29500938720
|
||||
|
||||
# Source: https://huggingface.co/Qwen/Qwen3.6-27B#best-practices
|
||||
# Source: https://unsloth.ai/docs/models/qwen3.5
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
|
||||
# Source: https://huggingface.co/Qwen/Qwen3.6-27B#best-practices
|
||||
# Source: https://unsloth.ai/docs/models/qwen3.5
|
||||
[sampling_defaults.non_thinking]
|
||||
temperature = 0.7
|
||||
top_p = 0.8
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
@@ -1,35 +0,0 @@
|
||||
model_id = "mlx-community/Qwen3.6-27B-bf16"
|
||||
n_layers = 64
|
||||
hidden_size = 5120
|
||||
num_key_value_heads = 4
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "qwen"
|
||||
quantization = "bf16"
|
||||
base_model = "Qwen3.6 27B"
|
||||
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
|
||||
|
||||
context_length = 262144
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 54713457120
|
||||
|
||||
# Source: https://huggingface.co/Qwen/Qwen3.6-27B#best-practices
|
||||
# Source: https://unsloth.ai/docs/models/qwen3.5
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
|
||||
# Source: https://huggingface.co/Qwen/Qwen3.6-27B#best-practices
|
||||
# Source: https://unsloth.ai/docs/models/qwen3.5
|
||||
[sampling_defaults.non_thinking]
|
||||
temperature = 0.7
|
||||
top_p = 0.8
|
||||
top_k = 20
|
||||
min_p = 0.0
|
||||
repetition_penalty = 1.0
|
||||
presence_penalty = 1.5
|
||||
@@ -1,33 +0,0 @@
|
||||
model_id = "moonshotai/Kimi-K2.6"
|
||||
n_layers = 61
|
||||
hidden_size = 7168
|
||||
num_key_value_heads = 64
|
||||
supports_tensor = true
|
||||
tasks = ["TextGeneration"]
|
||||
family = "kimi"
|
||||
quantization = ""
|
||||
base_model = "Kimi K2.6"
|
||||
capabilities = ["text", "thinking", "thinking_toggle", "vision"]
|
||||
|
||||
context_length = 262144
|
||||
|
||||
[storage_size]
|
||||
in_bytes = 595148192736
|
||||
|
||||
[vision]
|
||||
image_token_id = 163605
|
||||
model_type = "kimi_vl"
|
||||
weights_repo = "exolabs/Kimi-K2.6-vision"
|
||||
processor_repo = "moonshotai/Kimi-K2.6"
|
||||
|
||||
# Source: https://huggingface.co/moonshotai/Kimi-K2.6
|
||||
[sampling_defaults]
|
||||
temperature = 1.0
|
||||
top_p = 0.95
|
||||
min_p = 0.01
|
||||
|
||||
# Source: https://huggingface.co/moonshotai/Kimi-K2.6
|
||||
[sampling_defaults.non_thinking]
|
||||
temperature = 0.6
|
||||
top_p = 0.95
|
||||
min_p = 0.01
|
||||
@@ -113,23 +113,6 @@ def _extract_content(content: str | list[ResponseContentPart]) -> str:
|
||||
)
|
||||
|
||||
|
||||
def _append_tool_call(
|
||||
chat_template_messages: list[dict[str, Any]], tool_call: dict[str, Any]
|
||||
) -> None:
|
||||
if chat_template_messages:
|
||||
prev = chat_template_messages[-1]
|
||||
if prev.get("role") == "assistant" and isinstance(prev.get("content"), str):
|
||||
existing: list[dict[str, Any]] | None = prev.get("tool_calls")
|
||||
if existing is None:
|
||||
prev["tool_calls"] = [tool_call]
|
||||
else:
|
||||
existing.append(tool_call)
|
||||
return
|
||||
chat_template_messages.append(
|
||||
{"role": "assistant", "content": "", "tool_calls": [tool_call]}
|
||||
)
|
||||
|
||||
|
||||
async def responses_request_to_text_generation(
|
||||
request: ResponsesRequest,
|
||||
) -> TextGenerationTaskParams:
|
||||
@@ -199,44 +182,59 @@ async def responses_request_to_text_generation(
|
||||
| McpCallInputItem()
|
||||
| CustomToolCallInputItem()
|
||||
):
|
||||
_append_tool_call(
|
||||
chat_template_messages,
|
||||
chat_template_messages.append(
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.name,
|
||||
"arguments": item.arguments,
|
||||
},
|
||||
},
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.name,
|
||||
"arguments": item.arguments,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
case (
|
||||
LocalShellCallInputItem()
|
||||
| ShellCallInputItem()
|
||||
| ComputerCallInputItem()
|
||||
):
|
||||
_append_tool_call(
|
||||
chat_template_messages,
|
||||
chat_template_messages.append(
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.type,
|
||||
"arguments": json.dumps(item.action),
|
||||
},
|
||||
},
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.type,
|
||||
"arguments": json.dumps(item.action),
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
case ApplyPatchCallInputItem():
|
||||
_append_tool_call(
|
||||
chat_template_messages,
|
||||
chat_template_messages.append(
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "apply_patch",
|
||||
"arguments": json.dumps({"patch": item.patch}),
|
||||
},
|
||||
},
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "apply_patch",
|
||||
"arguments": json.dumps({"patch": item.patch}),
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
case (
|
||||
WebSearchCallInputItem()
|
||||
@@ -256,16 +254,21 @@ async def responses_request_to_text_generation(
|
||||
args = {"prompt": item.prompt}
|
||||
else:
|
||||
args = {"query": item.query}
|
||||
_append_tool_call(
|
||||
chat_template_messages,
|
||||
chat_template_messages.append(
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.type,
|
||||
"arguments": json.dumps(args),
|
||||
},
|
||||
},
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.type,
|
||||
"arguments": json.dumps(args),
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
case (
|
||||
FunctionCallOutputInputItem()
|
||||
@@ -317,16 +320,21 @@ async def responses_request_to_text_generation(
|
||||
}
|
||||
)
|
||||
case McpApprovalRequestInputItem():
|
||||
_append_tool_call(
|
||||
chat_template_messages,
|
||||
chat_template_messages.append(
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.name,
|
||||
"arguments": item.arguments,
|
||||
},
|
||||
},
|
||||
"role": "assistant",
|
||||
"content": "",
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": item.call_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": item.name,
|
||||
"arguments": item.arguments,
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
case McpApprovalResponseInputItem():
|
||||
chat_template_messages.append(
|
||||
|
||||
+23
-12
@@ -185,10 +185,7 @@ from exo.shared.types.tasks import (
|
||||
from exo.shared.types.tasks import (
|
||||
TextGeneration as TextGenerationTask,
|
||||
)
|
||||
from exo.shared.types.text_generation import (
|
||||
Base64ImageHash,
|
||||
TextGenerationTaskParams,
|
||||
)
|
||||
from exo.shared.types.text_generation import Base64Image, TextGenerationTaskParams
|
||||
from exo.shared.types.worker.downloads import DownloadCompleted
|
||||
from exo.shared.types.worker.instances import Instance, InstanceId, InstanceMeta
|
||||
from exo.shared.types.worker.shards import Sharding
|
||||
@@ -237,7 +234,6 @@ class API:
|
||||
self.node_id: NodeId = node_id
|
||||
self.last_completed_election: int = 0
|
||||
self.port = port
|
||||
self._sent_image_hashes: set[str] = set()
|
||||
|
||||
self.paused: bool = False
|
||||
self.paused_ev: anyio.Event = anyio.Event()
|
||||
@@ -287,7 +283,6 @@ class API:
|
||||
self.event_receiver.close()
|
||||
self.event_receiver = event_receiver
|
||||
self._tg.start_soon(self._apply_state)
|
||||
self._sent_image_hashes = set()
|
||||
|
||||
def unpause(self, result_clock: int):
|
||||
logger.info("Unpausing API")
|
||||
@@ -742,6 +737,8 @@ class API:
|
||||
"TODO: we should send a notification to the user to download the model"
|
||||
)
|
||||
|
||||
_sent_image_hashes: set[str] = set()
|
||||
|
||||
async def _send_text_generation_with_images(
|
||||
self, task_params: TextGenerationTaskParams
|
||||
) -> TextGeneration:
|
||||
@@ -753,19 +750,23 @@ class API:
|
||||
return command
|
||||
|
||||
hashes = [hashlib.sha256(img.encode("ascii")).hexdigest() for img in images]
|
||||
all_hashes = {idx: Base64ImageHash(h) for idx, h in enumerate(hashes)}
|
||||
task_params = task_params.model_copy(
|
||||
update={"images": [], "image_hashes": all_hashes}
|
||||
)
|
||||
command = TextGeneration(task_params=task_params)
|
||||
|
||||
cached_hashes: dict[int, str] = {}
|
||||
new_images: list[tuple[int, str]] = []
|
||||
for idx, (img, h) in enumerate(zip(images, hashes, strict=True)):
|
||||
if h not in self._sent_image_hashes:
|
||||
if h in self._sent_image_hashes:
|
||||
cached_hashes[idx] = h
|
||||
else:
|
||||
self._sent_image_hashes.add(h)
|
||||
new_images.append((idx, img))
|
||||
|
||||
wrapped_hashes = {idx: Base64Image(h) for idx, h in cached_hashes.items()}
|
||||
|
||||
if not new_images:
|
||||
task_params = task_params.model_copy(
|
||||
update={"images": [], "image_hashes": wrapped_hashes}
|
||||
)
|
||||
command = TextGeneration(task_params=task_params)
|
||||
await self._send(command)
|
||||
return command
|
||||
|
||||
@@ -774,6 +775,16 @@ class API:
|
||||
for i in range(0, len(img_data), EXO_MAX_CHUNK_SIZE):
|
||||
all_chunks.append((img_idx, img_data[i : i + EXO_MAX_CHUNK_SIZE]))
|
||||
|
||||
task_params = task_params.model_copy(
|
||||
update={
|
||||
"images": [],
|
||||
"image_hashes": wrapped_hashes,
|
||||
"total_input_chunks": len(all_chunks),
|
||||
"image_count": len(new_images),
|
||||
}
|
||||
)
|
||||
command = TextGeneration(task_params=task_params)
|
||||
|
||||
for global_idx, (img_idx, chunk_data) in enumerate(all_chunks):
|
||||
await self._send(
|
||||
SendInputChunk(
|
||||
|
||||
@@ -88,9 +88,7 @@ class DownloadCoordinator:
|
||||
|
||||
try:
|
||||
if progress.status == "complete":
|
||||
found = await to_thread.run_sync(
|
||||
resolve_existing_model, model_id, callback_shard.model_card
|
||||
)
|
||||
found = await to_thread.run_sync(resolve_existing_model, model_id)
|
||||
if found is not None:
|
||||
completed = self._completed_from_path(
|
||||
callback_shard, found, progress.total
|
||||
@@ -195,9 +193,7 @@ class DownloadCoordinator:
|
||||
return
|
||||
|
||||
# Check all model directories for pre-existing complete models
|
||||
found_path = await to_thread.run_sync(
|
||||
resolve_existing_model, model_id, shard.model_card
|
||||
)
|
||||
found_path = await to_thread.run_sync(resolve_existing_model, model_id)
|
||||
if found_path is not None:
|
||||
logger.info(f"DownloadCoordinator: Model {model_id} found at {found_path}")
|
||||
completed = self._completed_from_path(
|
||||
@@ -224,9 +220,7 @@ class DownloadCoordinator:
|
||||
)
|
||||
|
||||
if initial_progress.status == "complete":
|
||||
found = await to_thread.run_sync(
|
||||
resolve_existing_model, model_id, shard.model_card
|
||||
)
|
||||
found = await to_thread.run_sync(resolve_existing_model, model_id)
|
||||
if found is not None:
|
||||
completed = self._completed_from_path(
|
||||
shard, found, initial_progress.total
|
||||
@@ -357,9 +351,7 @@ class DownloadCoordinator:
|
||||
|
||||
if progress.status == "complete":
|
||||
found = await to_thread.run_sync(
|
||||
resolve_existing_model,
|
||||
model_id,
|
||||
progress.shard.model_card,
|
||||
resolve_existing_model, model_id
|
||||
)
|
||||
if found is not None:
|
||||
status: DownloadProgress = self._completed_from_path(
|
||||
@@ -388,9 +380,7 @@ class DownloadCoordinator:
|
||||
# (is_model_directory_complete) which validates that all
|
||||
# safetensors weight files are present.
|
||||
found = await to_thread.run_sync(
|
||||
resolve_existing_model,
|
||||
model_id,
|
||||
progress.shard.model_card,
|
||||
resolve_existing_model, model_id
|
||||
)
|
||||
if found is not None:
|
||||
status = self._completed_from_path(
|
||||
@@ -431,9 +421,7 @@ class DownloadCoordinator:
|
||||
(DownloadCompleted, DownloadOngoing, DownloadFailed),
|
||||
):
|
||||
continue
|
||||
found = await to_thread.run_sync(
|
||||
resolve_existing_model, mid, card
|
||||
)
|
||||
found = await to_thread.run_sync(resolve_existing_model, mid)
|
||||
if found is not None and is_read_only_model_dir(found):
|
||||
path_shard = PipelineShardMetadata(
|
||||
model_card=card,
|
||||
|
||||
@@ -35,7 +35,7 @@ from exo.shared.constants import (
|
||||
EXO_MODELS_DIRS,
|
||||
EXO_MODELS_READ_ONLY_DIRS,
|
||||
)
|
||||
from exo.shared.models.model_cards import ModelCard, ModelTask
|
||||
from exo.shared.models.model_cards import ModelTask
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.worker.downloads import (
|
||||
@@ -118,9 +118,7 @@ class InsufficientDiskSpaceError(Exception):
|
||||
"""Raised when no writable model directory has enough free space."""
|
||||
|
||||
|
||||
def resolve_existing_model(
|
||||
model_id: ModelId, card: ModelCard | None = None
|
||||
) -> Path | None:
|
||||
def resolve_existing_model(model_id: ModelId) -> Path | None:
|
||||
"""Search all model directories for a complete, pre-existing model.
|
||||
|
||||
Checks read-only directories first, then writable directories.
|
||||
@@ -130,7 +128,7 @@ def resolve_existing_model(
|
||||
normalized = model_id.normalize()
|
||||
for search_dir in (*EXO_MODELS_READ_ONLY_DIRS, *EXO_MODELS_DIRS):
|
||||
candidate = search_dir / normalized
|
||||
if candidate.is_dir() and is_model_directory_complete(candidate, card):
|
||||
if candidate.is_dir() and is_model_directory_complete(candidate):
|
||||
return candidate
|
||||
return None
|
||||
|
||||
@@ -167,29 +165,6 @@ def select_download_dir(required_bytes: int) -> Path:
|
||||
)
|
||||
|
||||
|
||||
async def select_download_dir_for_shard(
|
||||
model_id: ModelId,
|
||||
filtered_file_list: list[FileListEntry],
|
||||
total_size: int,
|
||||
) -> Path:
|
||||
for candidate_dir in EXO_MODELS_DIRS:
|
||||
if not candidate_dir.exists():
|
||||
continue
|
||||
sub = candidate_dir / model_id.normalize()
|
||||
if not await aios.path.isdir(sub):
|
||||
continue
|
||||
existing_bytes = 0
|
||||
for file_entry in filtered_file_list:
|
||||
existing_bytes += await get_downloaded_size(sub / file_entry.path)
|
||||
remaining = max(total_size - existing_bytes, 0)
|
||||
try:
|
||||
if shutil.disk_usage(candidate_dir).free >= remaining:
|
||||
return candidate_dir
|
||||
except OSError:
|
||||
continue
|
||||
return select_download_dir(total_size)
|
||||
|
||||
|
||||
async def resolve_model_dir(model_id: ModelId) -> Path:
|
||||
"""Return the directory for a model's files, creating it if needed.
|
||||
|
||||
@@ -304,26 +279,10 @@ def _scan_model_directory(
|
||||
return list(entries_by_path.values())
|
||||
|
||||
|
||||
def is_model_directory_complete(model_dir: Path, card: ModelCard | None = None) -> bool:
|
||||
"""Check if a model directory contains all required weight files.
|
||||
Also checks for sibling weights repo.
|
||||
"""
|
||||
def is_model_directory_complete(model_dir: Path) -> bool:
|
||||
"""Check if a model directory contains all required weight files."""
|
||||
file_list = _scan_model_directory(model_dir, recursive=True)
|
||||
if file_list is None or not all(f.size is not None for f in file_list):
|
||||
return False
|
||||
if (
|
||||
card is not None
|
||||
and card.vision is not None
|
||||
and card.vision.weights_repo != str(card.model_id)
|
||||
):
|
||||
vision_id = ModelId(card.vision.weights_repo)
|
||||
normalized = vision_id.normalize()
|
||||
for search_dir in (*EXO_MODELS_READ_ONLY_DIRS, *EXO_MODELS_DIRS):
|
||||
candidate = search_dir / normalized
|
||||
if candidate.is_dir() and is_model_directory_complete(candidate):
|
||||
return True
|
||||
return False
|
||||
return True
|
||||
return file_list is not None and all(f.size is not None for f in file_list)
|
||||
|
||||
|
||||
async def _build_file_list_from_local_directory(
|
||||
@@ -875,9 +834,7 @@ async def download_shard(
|
||||
else EXO_DEFAULT_MODELS_DIR / model_id.normalize()
|
||||
)
|
||||
else:
|
||||
models_dir = await select_download_dir_for_shard(
|
||||
model_id, filtered_file_list, total_size
|
||||
)
|
||||
models_dir = select_download_dir(total_size)
|
||||
target_dir = models_dir / model_id.normalize()
|
||||
await aios.makedirs(target_dir, exist_ok=True)
|
||||
file_progress: dict[str, RepoFileDownloadProgress] = {}
|
||||
|
||||
@@ -117,39 +117,40 @@ class ResumableShardDownloader(ShardDownloader):
|
||||
) -> Path:
|
||||
allow_patterns = ["config.json"] if config_only else None
|
||||
|
||||
has_vision_sibling = (
|
||||
not config_only
|
||||
and not self.offline
|
||||
and shard.model_card.vision is not None
|
||||
and shard.model_card.vision.weights_repo != str(shard.model_card.model_id)
|
||||
)
|
||||
|
||||
async def main_progress(
|
||||
cb_shard: ShardMetadata, progress: RepoDownloadProgress
|
||||
) -> None:
|
||||
if has_vision_sibling and progress.status == "complete":
|
||||
return
|
||||
await self.on_progress_wrapper(cb_shard, progress)
|
||||
|
||||
target_dir, _ = await download_shard(
|
||||
shard,
|
||||
main_progress,
|
||||
self.on_progress_wrapper,
|
||||
max_parallel_downloads=self.max_parallel_downloads,
|
||||
allow_patterns=allow_patterns,
|
||||
skip_internet=self.offline,
|
||||
)
|
||||
|
||||
if has_vision_sibling:
|
||||
vision_shard = self._build_vision_shard(shard)
|
||||
|
||||
async def vision_progress(
|
||||
_cb_shard: ShardMetadata, progress: RepoDownloadProgress
|
||||
) -> None:
|
||||
await self.on_progress_wrapper(shard, progress)
|
||||
|
||||
if (
|
||||
not config_only
|
||||
and not self.offline
|
||||
and shard.model_card.vision
|
||||
and shard.model_card.vision.weights_repo != str(shard.model_card.model_id)
|
||||
):
|
||||
vision_repo = shard.model_card.vision.weights_repo
|
||||
vision_card = ModelCard(
|
||||
model_id=ModelId(vision_repo),
|
||||
storage_size=Memory.from_bytes(0),
|
||||
n_layers=1,
|
||||
hidden_size=1,
|
||||
supports_tensor=False,
|
||||
tasks=[ModelTask.TextGeneration],
|
||||
)
|
||||
vision_shard = PipelineShardMetadata(
|
||||
model_card=vision_card,
|
||||
device_rank=0,
|
||||
world_size=1,
|
||||
start_layer=0,
|
||||
end_layer=1,
|
||||
n_layers=1,
|
||||
)
|
||||
await download_shard(
|
||||
vision_shard,
|
||||
vision_progress,
|
||||
self.on_progress_wrapper,
|
||||
max_parallel_downloads=self.max_parallel_downloads,
|
||||
allow_patterns=["*.safetensors", "config.json"],
|
||||
skip_internet=self.offline,
|
||||
@@ -157,87 +158,6 @@ class ResumableShardDownloader(ShardDownloader):
|
||||
|
||||
return target_dir
|
||||
|
||||
async def _status_for_shard(
|
||||
self, shard: ShardMetadata
|
||||
) -> tuple[Path, RepoDownloadProgress]:
|
||||
async def _noop(
|
||||
_cb_shard: ShardMetadata, _progress: RepoDownloadProgress
|
||||
) -> None:
|
||||
return
|
||||
|
||||
path, main_progress = await download_shard(
|
||||
shard,
|
||||
_noop,
|
||||
skip_download=True,
|
||||
skip_internet=self.offline,
|
||||
)
|
||||
|
||||
has_vision_sibling = (
|
||||
shard.model_card.vision is not None
|
||||
and shard.model_card.vision.weights_repo != str(shard.model_card.model_id)
|
||||
)
|
||||
if not has_vision_sibling:
|
||||
return path, main_progress
|
||||
|
||||
vision_shard = self._build_vision_shard(shard)
|
||||
_, vision_progress = await download_shard(
|
||||
vision_shard,
|
||||
_noop,
|
||||
skip_download=True,
|
||||
skip_internet=self.offline,
|
||||
)
|
||||
combined = self._combine_progress(shard, main_progress, vision_progress)
|
||||
return path, combined
|
||||
|
||||
@staticmethod
|
||||
def _build_vision_shard(shard: ShardMetadata) -> PipelineShardMetadata:
|
||||
assert shard.model_card.vision is not None
|
||||
vision_card = ModelCard(
|
||||
model_id=ModelId(shard.model_card.vision.weights_repo),
|
||||
storage_size=Memory.from_bytes(0),
|
||||
n_layers=1,
|
||||
hidden_size=1,
|
||||
supports_tensor=False,
|
||||
tasks=[ModelTask.TextGeneration],
|
||||
)
|
||||
return PipelineShardMetadata(
|
||||
model_card=vision_card,
|
||||
device_rank=0,
|
||||
world_size=1,
|
||||
start_layer=0,
|
||||
end_layer=1,
|
||||
n_layers=1,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _combine_progress(
|
||||
shard: ShardMetadata,
|
||||
main: RepoDownloadProgress,
|
||||
vision: RepoDownloadProgress,
|
||||
) -> RepoDownloadProgress:
|
||||
status_rank = {"not_started": 0, "in_progress": 1, "complete": 2}
|
||||
combined_status = min(
|
||||
(main.status, vision.status), key=lambda s: status_rank[s]
|
||||
)
|
||||
file_progress = dict(main.file_progress)
|
||||
for file_path, fp in vision.file_progress.items():
|
||||
file_progress[f"{vision.repo_id}/{file_path}"] = fp
|
||||
return RepoDownloadProgress(
|
||||
repo_id=main.repo_id,
|
||||
repo_revision=main.repo_revision,
|
||||
shard=shard,
|
||||
completed_files=main.completed_files + vision.completed_files,
|
||||
total_files=main.total_files + vision.total_files,
|
||||
downloaded=main.downloaded + vision.downloaded,
|
||||
downloaded_this_session=main.downloaded_this_session
|
||||
+ vision.downloaded_this_session,
|
||||
total=main.total + vision.total,
|
||||
overall_speed=main.overall_speed + vision.overall_speed,
|
||||
overall_eta=max(main.overall_eta, vision.overall_eta),
|
||||
status=combined_status,
|
||||
file_progress=file_progress,
|
||||
)
|
||||
|
||||
async def get_shard_download_status(
|
||||
self,
|
||||
) -> AsyncIterator[tuple[Path, RepoDownloadProgress]]:
|
||||
@@ -246,7 +166,12 @@ class ResumableShardDownloader(ShardDownloader):
|
||||
) -> tuple[Path, RepoDownloadProgress]:
|
||||
"""Helper coroutine that builds the shard for a model and gets its download status."""
|
||||
shard = await build_full_shard(model_id)
|
||||
return await self._status_for_shard(shard)
|
||||
return await download_shard(
|
||||
shard,
|
||||
self.on_progress_wrapper,
|
||||
skip_download=True,
|
||||
skip_internet=self.offline,
|
||||
)
|
||||
|
||||
semaphore = asyncio.Semaphore(self.max_parallel_downloads)
|
||||
|
||||
@@ -270,5 +195,10 @@ class ResumableShardDownloader(ShardDownloader):
|
||||
async def get_shard_download_status_for_shard(
|
||||
self, shard: ShardMetadata
|
||||
) -> RepoDownloadProgress:
|
||||
_, progress = await self._status_for_shard(shard)
|
||||
_, progress = await download_shard(
|
||||
shard,
|
||||
self.on_progress_wrapper,
|
||||
skip_download=True,
|
||||
skip_internet=self.offline,
|
||||
)
|
||||
return progress
|
||||
@@ -410,6 +410,8 @@ class Master:
|
||||
continue
|
||||
|
||||
logger.debug(f"Master indexing event: {str(event)[:100]}")
|
||||
indexed = IndexedEvent(event=event, idx=len(self._event_log))
|
||||
self.state = apply(self.state, indexed)
|
||||
|
||||
event = event.model_copy(
|
||||
update={"_master_time_stamp": datetime.now(tz=timezone.utc)}
|
||||
@@ -419,9 +421,6 @@ class Master:
|
||||
update={"when": str(datetime.now(tz=timezone.utc))}
|
||||
)
|
||||
|
||||
indexed = IndexedEvent(event=event, idx=len(self._event_log))
|
||||
self.state = apply(self.state, indexed)
|
||||
|
||||
self._event_log.append(event)
|
||||
await self._send_event(indexed)
|
||||
|
||||
|
||||
@@ -68,7 +68,6 @@ DASHBOARD_DIR = (
|
||||
# Log files (data/logs or cache)
|
||||
EXO_LOG_DIR = EXO_CACHE_HOME / "exo_log"
|
||||
EXO_LOG = EXO_LOG_DIR / "exo.log"
|
||||
EXO_TEST_LOG = EXO_CACHE_HOME / "exo_test.log"
|
||||
|
||||
# Identity (config)
|
||||
EXO_NODE_ID_KEYPAIR = EXO_CONFIG_HOME / "node_id.keypair"
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
from pathlib import Path
|
||||
from collections.abc import Sequence
|
||||
import tomlkit
|
||||
from exo.utils.pydantic_ext import FrozenModel
|
||||
from typing import Self, Any
|
||||
from pydantic import Field, BaseModel, model_validator
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict, PydanticBaseSettingsSource, TomlConfigSettingsSource
|
||||
from exo.shared.types.common import NodeId, ModelId
|
||||
from exo.shared.types.worker.instances import InstanceId
|
||||
from exo.shared.constants import EXO_CONFIG_HOME, EXO_DATA_HOME, EXO_CACHE_HOME
|
||||
from exo.utils.dashboard_path import find_dashboard, find_resources
|
||||
|
||||
|
||||
def default_merge[T: BaseModel](left: T, right: T) -> T:
|
||||
if left == right:
|
||||
return left
|
||||
merged_dict = {}
|
||||
for key in type(left).model_fields:
|
||||
try:
|
||||
merged_dict[key] = getattr(left, key).merge( # pyright: ignore[reportAny]
|
||||
getattr(right, key, None)
|
||||
)
|
||||
except AttributeError:
|
||||
raise NotImplementedError("Cluster Option using default implementation incorrectly")
|
||||
|
||||
return type(left).model_validate(merged_dict)
|
||||
|
||||
|
||||
def _parse_colon_separated_dirs(obj: Any) -> set[Path]: # pyright: ignore[reportAny]
|
||||
if isinstance(obj, (list, set)):
|
||||
return set(Path(d).expanduser() for d in obj) # pyright: ignore[reportUnknownArgumentType, reportUnknownVariableType]
|
||||
else:
|
||||
return set(Path(d).expanduser() for d in str(obj).split(":")) # pyright: ignore[reportAny]
|
||||
|
||||
class ModelDirsSettings(BaseModel, frozen=True):
|
||||
# env: EXO_MODEL_DIRS_DEFAULT prepends to WRITEABLE, defaults to EXO_DATA_HOME/models
|
||||
# env: EXO_MODEL_DIRS_WRITEABLE, defaults to []
|
||||
writeable: list[Path] = []
|
||||
# env: EXO_MODEL_DIRS_READONLY, defaults to []
|
||||
readonly: list[Path] = []
|
||||
|
||||
@model_validator(mode="before")
|
||||
@classmethod
|
||||
def build_defaults(cls, data: Any) -> Any: # pyright: ignore[reportAny]
|
||||
if not isinstance(data, dict):
|
||||
return data # pyright: ignore[reportAny]
|
||||
default = Path(data.get("default", EXO_DATA_HOME / "models")).expanduser() # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType]
|
||||
readonly = _parse_colon_separated_dirs(data.get("readonly", [])) # pyright: ignore[reportUnknownMemberType]
|
||||
writeable = _parse_colon_separated_dirs(data.get("writeable", [])).difference(readonly) # pyright: ignore[reportUnknownMemberType]
|
||||
if default not in readonly:
|
||||
writeable = [default, *writeable]
|
||||
return {**data, "writeable": writeable, "readonly": readonly} # pyright: ignore[reportUnknownVariableType]
|
||||
|
||||
class RuntimeDirsSettings(BaseModel, frozen=True):
|
||||
dashboard: Path = Field(default_factory=find_dashboard)
|
||||
resources: Path = Field(default_factory=find_resources)
|
||||
logs: Path = EXO_CACHE_HOME / "log"
|
||||
log_file: str = "latest.log"
|
||||
|
||||
def log_file_path(self):
|
||||
return self.logs / self.log_file
|
||||
|
||||
# doesnt require merge
|
||||
class LocalSettings(FrozenModel):
|
||||
runtime_dirs: RuntimeDirsSettings
|
||||
model_dirs: ModelDirsSettings
|
||||
|
||||
class InstanceSettings(FrozenModel):
|
||||
# env: EXO_INSTANCE_DEFAULTS_BATCH_CONCURRENCY
|
||||
batch_concurrency: int
|
||||
|
||||
def merge(self, other: Self) -> Self:
|
||||
return type(self)(batch_concurrency=min(self.batch_concurrency, other.batch_concurrency))
|
||||
|
||||
class ClusterSettings(FrozenModel):
|
||||
instance_defaults: InstanceSettings = InstanceSettings(batch_concurrency=8)
|
||||
model_settings_overrides: dict[ModelId, InstanceSettings] = {}
|
||||
|
||||
def merge(self, other: Self) -> Self:
|
||||
return default_merge(self, other)
|
||||
|
||||
class SettingsFile(BaseSettings):
|
||||
model_config = SettingsConfigDict(
|
||||
extra='ignore',
|
||||
frozen=True,
|
||||
toml_file=EXO_CONFIG_HOME / "config.toml",
|
||||
env_prefix="EXO_",
|
||||
env_nested_delimiter="_",
|
||||
env_ignore_empty=True,
|
||||
)
|
||||
|
||||
model_dirs: ModelDirsSettings
|
||||
runtime_dirs: RuntimeDirsSettings
|
||||
model_settings_overrides: dict[ModelId, InstanceSettings] = {}
|
||||
instance_defaults: InstanceSettings
|
||||
|
||||
def get_local(self) -> LocalSettings:
|
||||
...
|
||||
def get_cluster(self) -> ClusterSettings:
|
||||
...
|
||||
|
||||
@classmethod
|
||||
def settings_customise_sources(
|
||||
cls,
|
||||
settings_cls: type[BaseSettings],
|
||||
init_settings: PydanticBaseSettingsSource,
|
||||
env_settings: PydanticBaseSettingsSource,
|
||||
dotenv_settings: PydanticBaseSettingsSource,
|
||||
file_secret_settings: PydanticBaseSettingsSource,
|
||||
) -> tuple[PydanticBaseSettingsSource, ...]:
|
||||
return (init_settings, env_settings, TomlConfigSettingsSource(settings_cls),)
|
||||
|
||||
def sync(self):
|
||||
"""nb: only call this once per save"""
|
||||
cfg_path = type(self).model_config.get("toml_file", None)
|
||||
if isinstance(cfg_path, Sequence):
|
||||
cfg_path=cfg_path[0]
|
||||
if cfg_path:
|
||||
with open(cfg_path, "w") as fp:
|
||||
tomlkit.dump(self.model_dump(exclude_defaults=True), fp) # pyright: ignore[reportUnknownMemberType]
|
||||
|
||||
|
||||
|
||||
class StateSettings(FrozenModel):
|
||||
per_node: dict[NodeId, LocalSettings]
|
||||
per_instance: dict[InstanceId, InstanceSettings]
|
||||
cluster: ClusterSettings
|
||||
|
||||
def model_merge_local(self, node_id: NodeId, settings: LocalSettings) -> Self:
|
||||
return self.model_copy(update={
|
||||
"per_node": {
|
||||
**self.per_node,
|
||||
node_id: settings
|
||||
}
|
||||
})
|
||||
|
||||
def model_merge_cluster(self, settings: ClusterSettings) -> Self:
|
||||
return self.model_copy(update={
|
||||
"cluster": self.cluster.merge(settings)
|
||||
})
|
||||
|
||||
def settings_for(self, node_id: NodeId) -> StoredSettings:
|
||||
merged = {}
|
||||
for key, val in self.cluster.model_dump(exclude_defaults=True).items(): # pyright: ignore[reportAny]
|
||||
merged[key] = val
|
||||
|
||||
if (local := self.per_node.get(node_id, None)) is not None:
|
||||
for key, val in local.model_dump(exclude_defaults=True).items(): # pyright: ignore[reportAny]
|
||||
merged[key] = val
|
||||
|
||||
return StoredSettings.model_validate(merged)
|
||||
|
||||
def sync(self, node_id: NodeId):
|
||||
"""nb: only call this once per save"""
|
||||
toml_file=EXO_CONFIG_HOME / "config.toml"
|
||||
with open(toml_file, "w") as fp:
|
||||
tomlkit.dump(self.settings_for(node_id).model_dump(exclude_defaults=True), fp) # pyright: ignore[reportUnknownMemberType]
|
||||
|
||||
@@ -114,6 +114,8 @@ class TextGenerationTaskParams(BaseModel, frozen=True):
|
||||
frequency_penalty: float | None = None
|
||||
images: list[Base64Image] = Field(default_factory=list)
|
||||
image_hashes: dict[int, Base64ImageHash] = Field(default_factory=dict)
|
||||
total_input_chunks: int = 0
|
||||
image_count: int = 0
|
||||
|
||||
def with_card_sampling_defaults(self) -> "TextGenerationTaskParams":
|
||||
from exo.shared.models.model_cards import get_card
|
||||
|
||||
@@ -70,11 +70,6 @@ class FinishedResponse(BaseRunnerResponse):
|
||||
pass
|
||||
|
||||
|
||||
class ModelLoadingResponse(BaseRunnerResponse):
|
||||
layers_loaded: int
|
||||
total: int
|
||||
|
||||
|
||||
class PrefillProgressResponse(BaseRunnerResponse):
|
||||
processed_tokens: int
|
||||
total_tokens: int
|
||||
@@ -1,8 +1,10 @@
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import cast
|
||||
from functools import cache
|
||||
|
||||
|
||||
@cache
|
||||
def find_resources() -> Path:
|
||||
resources = _find_resources_in_repo() or _find_resources_in_bundle()
|
||||
if resources is None:
|
||||
@@ -31,6 +33,7 @@ def _find_resources_in_bundle() -> Path | None:
|
||||
return None
|
||||
|
||||
|
||||
@cache
|
||||
def find_dashboard() -> Path:
|
||||
dashboard = _find_dashboard_in_repo() or _find_dashboard_in_bundle()
|
||||
if not dashboard:
|
||||
|
||||
@@ -8,7 +8,6 @@ from PIL import Image
|
||||
|
||||
from exo.api.types import AdvancedImageParams
|
||||
from exo.download.download_utils import build_model_path
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.worker.instances import BoundInstance
|
||||
from exo.shared.types.worker.shards import CfgShardMetadata, PipelineShardMetadata
|
||||
from exo.worker.engines.image.config import ImageModelConfig
|
||||
@@ -23,14 +22,13 @@ from exo.worker.runner.bootstrap import logger
|
||||
|
||||
|
||||
class DistributedImageModel:
|
||||
model_id: ModelId
|
||||
_config: ImageModelConfig
|
||||
_adapter: ModelAdapter[Any, Any]
|
||||
_runner: DiffusionRunner
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_id: ModelId,
|
||||
model_id: str,
|
||||
local_path: Path,
|
||||
shard_metadata: PipelineShardMetadata | CfgShardMetadata,
|
||||
group: Optional[mx.distributed.Group] = None,
|
||||
@@ -70,7 +68,6 @@ class DistributedImageModel:
|
||||
else:
|
||||
logger.info("Single-node initialization")
|
||||
|
||||
self.model_id = model_id
|
||||
self._config = config
|
||||
self._adapter = adapter
|
||||
self._runner = runner
|
||||
|
||||
@@ -3,7 +3,7 @@ import io
|
||||
import random
|
||||
import tempfile
|
||||
import time
|
||||
from collections.abc import Callable, Iterator
|
||||
from collections.abc import Callable
|
||||
from pathlib import Path
|
||||
from typing import Generator, Literal
|
||||
|
||||
@@ -17,10 +17,11 @@ from exo.api.types import (
|
||||
ImageGenerationTaskParams,
|
||||
ImageSize,
|
||||
)
|
||||
from exo.shared.constants import EXO_MAX_CHUNK_SIZE
|
||||
from exo.shared.types.chunks import ImageChunk
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.memory import Memory
|
||||
from exo.shared.types.worker.runner_response import (
|
||||
ImageGenerationResponse,
|
||||
PartialImageResponse,
|
||||
)
|
||||
from exo.worker.engines.image.distributed_model import DistributedImageModel
|
||||
|
||||
|
||||
@@ -70,8 +71,16 @@ def generate_image(
|
||||
model: DistributedImageModel,
|
||||
task: ImageGenerationTaskParams | ImageEditsTaskParams,
|
||||
cancel_checker: Callable[[], bool] | None = None,
|
||||
) -> Generator[ImageChunk, None, None]:
|
||||
"""Generate image(s), optionally yielding partial results."""
|
||||
) -> Generator[ImageGenerationResponse | PartialImageResponse, None, None]:
|
||||
"""Generate image(s), optionally yielding partial results.
|
||||
|
||||
When partial_images > 0 or stream=True, yields PartialImageResponse for
|
||||
intermediate images, then ImageGenerationResponse for the final image.
|
||||
|
||||
Yields:
|
||||
PartialImageResponse for intermediate images (if partial_images > 0, first image only)
|
||||
ImageGenerationResponse for final complete images
|
||||
"""
|
||||
width, height = parse_size(task.size)
|
||||
quality: Literal["low", "medium", "high"] = task.quality or "medium"
|
||||
|
||||
@@ -133,14 +142,12 @@ def generate_image(
|
||||
image = image.convert("RGB")
|
||||
image.save(buffer, format=image_format)
|
||||
|
||||
yield from _process_image_response(
|
||||
yield PartialImageResponse(
|
||||
image_data=buffer.getvalue(),
|
||||
image_format=task.output_format,
|
||||
format=task.output_format,
|
||||
partial_index=partial_idx,
|
||||
total_partials=total_partials,
|
||||
image_index=image_num,
|
||||
model_id=model.model_id,
|
||||
stats=None,
|
||||
)
|
||||
else:
|
||||
image = result
|
||||
@@ -182,54 +189,9 @@ def generate_image(
|
||||
image = image.convert("RGB")
|
||||
image.save(buffer, format=image_format)
|
||||
|
||||
yield from _process_image_response(
|
||||
yield ImageGenerationResponse(
|
||||
image_data=buffer.getvalue(),
|
||||
image_format=task.output_format,
|
||||
format=task.output_format,
|
||||
stats=stats,
|
||||
image_index=image_num,
|
||||
model_id=model.model_id,
|
||||
partial_index=None,
|
||||
total_partials=None,
|
||||
)
|
||||
|
||||
|
||||
def _process_image_response(
|
||||
image_data: bytes,
|
||||
image_index: int,
|
||||
image_format: Literal["png", "jpeg", "webp"],
|
||||
partial_index: int | None,
|
||||
total_partials: int | None,
|
||||
stats: ImageGenerationStats | None,
|
||||
model_id: ModelId,
|
||||
) -> Iterator[ImageChunk]:
|
||||
"""Process a single image response and send chunks."""
|
||||
is_partial = partial_index is not None
|
||||
encoded_data = base64.b64encode(image_data).decode("utf-8")
|
||||
# Extract stats from final ImageGenerationResponse if available
|
||||
data_chunks = [
|
||||
encoded_data[i : i + EXO_MAX_CHUNK_SIZE]
|
||||
for i in range(0, len(encoded_data), EXO_MAX_CHUNK_SIZE)
|
||||
]
|
||||
total_chunks = len(data_chunks)
|
||||
|
||||
def _data_to_chunk(item: tuple[int, str]) -> ImageChunk:
|
||||
chunk_index, chunk_data = item
|
||||
# Only include stats on the last chunk of the final image
|
||||
chunk_stats = (
|
||||
stats if chunk_index == total_chunks - 1 and not is_partial else None
|
||||
)
|
||||
|
||||
return ImageChunk(
|
||||
model=model_id,
|
||||
data=chunk_data,
|
||||
chunk_index=chunk_index,
|
||||
total_chunks=total_chunks,
|
||||
image_index=image_index,
|
||||
is_partial=is_partial,
|
||||
partial_index=partial_index,
|
||||
total_partials=total_partials,
|
||||
stats=chunk_stats,
|
||||
format=image_format,
|
||||
)
|
||||
|
||||
return map(_data_to_chunk, enumerate(data_chunks))
|
||||
@@ -1,5 +1,5 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Callable, Generator
|
||||
from collections.abc import Callable
|
||||
from functools import partial
|
||||
from inspect import signature
|
||||
from typing import TYPE_CHECKING, Literal, Protocol, cast
|
||||
@@ -12,7 +12,7 @@ from mlx.nn.layers.distributed import (
|
||||
sum_gradients,
|
||||
)
|
||||
from mlx_lm.models.base import (
|
||||
scaled_dot_product_attention,
|
||||
scaled_dot_product_attention, # pyright: ignore[reportUnknownVariableType]
|
||||
)
|
||||
from mlx_lm.models.cache import ArraysCache, KVCache
|
||||
from mlx_lm.models.deepseek_v3 import DeepseekV3MLP
|
||||
@@ -59,13 +59,14 @@ from mlx_lm.models.step3p5 import Model as Step35Model
|
||||
from mlx_lm.models.step3p5 import Step3p5MLP as Step35MLP
|
||||
from mlx_lm.models.step3p5 import Step3p5Model as Step35InnerModel
|
||||
|
||||
from exo.shared.types.worker.runner_response import ModelLoadingResponse
|
||||
from exo.shared.types.worker.shards import PipelineShardMetadata
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from mlx_lm.models.cache import Cache
|
||||
|
||||
LayerLoadedCallback = Callable[[int, int], None] # (layers_loaded, total_layers)
|
||||
|
||||
|
||||
_pending_prefill_sends: list[tuple[mx.array, int, mx.distributed.Group]] = []
|
||||
|
||||
@@ -275,7 +276,8 @@ def pipeline_auto_parallel(
|
||||
model: nn.Module,
|
||||
group: mx.distributed.Group,
|
||||
model_shard_meta: PipelineShardMetadata,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
"""
|
||||
Automatically parallelize a model across multiple devices.
|
||||
Args:
|
||||
@@ -295,7 +297,8 @@ def pipeline_auto_parallel(
|
||||
total = len(layers)
|
||||
for i, layer in enumerate(layers):
|
||||
mx.eval(layer) # type: ignore
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
|
||||
layers[0] = PipelineFirstLayer(layers[0], device_rank, group=group)
|
||||
layers[-1] = PipelineLastLayer(
|
||||
@@ -306,20 +309,24 @@ def pipeline_auto_parallel(
|
||||
)
|
||||
|
||||
if isinstance(inner_model_instance, GptOssMoeModel):
|
||||
inner_model_instance.layer_types = inner_model_instance.layer_types[
|
||||
inner_model_instance.layer_types = inner_model_instance.layer_types[ # type: ignore
|
||||
start_layer:end_layer
|
||||
]
|
||||
# We can assume the model has at least one layer thanks to placement.
|
||||
# If a layer type doesn't exist, we can set it to 0.
|
||||
inner_model_instance.swa_idx = (
|
||||
0
|
||||
if "sliding_attention" not in inner_model_instance.layer_types
|
||||
else inner_model_instance.layer_types.index("sliding_attention")
|
||||
if "sliding_attention" not in inner_model_instance.layer_types # type: ignore
|
||||
else inner_model_instance.layer_types.index( # type: ignore
|
||||
"sliding_attention"
|
||||
)
|
||||
)
|
||||
inner_model_instance.ga_idx = (
|
||||
0
|
||||
if "full_attention" not in inner_model_instance.layer_types
|
||||
else inner_model_instance.layer_types.index("full_attention")
|
||||
if "full_attention" not in inner_model_instance.layer_types # type: ignore
|
||||
else inner_model_instance.layer_types.index( # type: ignore
|
||||
"full_attention"
|
||||
)
|
||||
)
|
||||
|
||||
if isinstance(inner_model_instance, Step35InnerModel):
|
||||
@@ -453,7 +460,8 @@ def patch_tensor_model[T](model: T) -> T:
|
||||
def tensor_auto_parallel(
|
||||
model: nn.Module,
|
||||
group: mx.distributed.Group,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
all_to_sharded_linear = partial(
|
||||
shard_linear,
|
||||
sharding="all-to-sharded",
|
||||
@@ -587,7 +595,7 @@ def tensor_auto_parallel(
|
||||
else:
|
||||
raise ValueError(f"Unsupported model type: {type(model)}")
|
||||
|
||||
model = yield from tensor_parallel_sharding_strategy.shard_model(model)
|
||||
model = tensor_parallel_sharding_strategy.shard_model(model, on_layer_loaded)
|
||||
return patch_tensor_model(model)
|
||||
|
||||
|
||||
@@ -611,14 +619,16 @@ class TensorParallelShardingStrategy(ABC):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]: ...
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module: ...
|
||||
|
||||
|
||||
class LlamaShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(LlamaModel, model)
|
||||
total = len(model.layers)
|
||||
for i, layer in enumerate(model.layers):
|
||||
@@ -636,8 +646,8 @@ class LlamaShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.mlp.down_proj = self.sharded_to_all_linear(layer.mlp.down_proj)
|
||||
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
|
||||
|
||||
@@ -671,7 +681,8 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(DeepseekV3Model, model)
|
||||
total = len(model.layers)
|
||||
|
||||
@@ -727,8 +738,8 @@ class DeepSeekShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.mlp.sharding_group = self.group
|
||||
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
|
||||
return model
|
||||
|
||||
@@ -753,7 +764,8 @@ class GLM4MoeLiteShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(GLM4MoeLiteModel, model)
|
||||
total = len(model.layers) # type: ignore
|
||||
for i, layer in enumerate(model.layers): # type: ignore
|
||||
@@ -804,8 +816,8 @@ class GLM4MoeLiteShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.mlp = ShardedMoE(layer.mlp) # type: ignore
|
||||
layer.mlp.sharding_group = self.group # type: ignore
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
|
||||
return model
|
||||
|
||||
@@ -879,7 +891,7 @@ class WrappedMiniMaxAttention(CustomMlxLayer):
|
||||
keys,
|
||||
values,
|
||||
cache=cache,
|
||||
scale=self._original_layer.scale,
|
||||
scale=self._original_layer.scale, # type: ignore
|
||||
mask=mask,
|
||||
)
|
||||
|
||||
@@ -892,7 +904,8 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(MiniMaxModel, model)
|
||||
total = len(model.layers)
|
||||
for i, layer in enumerate(model.layers):
|
||||
@@ -918,11 +931,11 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy):
|
||||
self.all_to_sharded_linear_in_place(
|
||||
layer.block_sparse_moe.switch_mlp.up_proj
|
||||
)
|
||||
layer.block_sparse_moe = ShardedMoE(layer.block_sparse_moe) # type: ignore
|
||||
layer.block_sparse_moe.sharding_group = self.group
|
||||
layer.block_sparse_moe = ShardedMoE(layer.block_sparse_moe) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
|
||||
layer.block_sparse_moe.sharding_group = self.group # pyright: ignore[reportAttributeAccessIssue]
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
|
||||
|
||||
@@ -930,7 +943,8 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(
|
||||
Qwen3Model
|
||||
| Qwen3MoeModel
|
||||
@@ -1085,8 +1099,8 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
|
||||
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
|
||||
|
||||
@@ -1094,7 +1108,8 @@ class Glm4MoeShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(Glm4MoeModel, model)
|
||||
total = len(model.layers)
|
||||
for i, layer in enumerate(model.layers):
|
||||
@@ -1130,8 +1145,8 @@ class Glm4MoeShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.mlp.up_proj = self.all_to_sharded_linear(layer.mlp.up_proj)
|
||||
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
|
||||
|
||||
@@ -1139,7 +1154,8 @@ class GptOssShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(GptOssMoeModel, model)
|
||||
total = len(model.layers)
|
||||
|
||||
@@ -1168,10 +1184,10 @@ class GptOssShardingStrategy(TensorParallelShardingStrategy):
|
||||
self.all_to_sharded_linear_in_place(layer.mlp.experts.up_proj)
|
||||
|
||||
layer.mlp = ShardedMoE(layer.mlp) # type: ignore
|
||||
layer.mlp.sharding_group = self.group
|
||||
layer.mlp.sharding_group = self.group # pyright: ignore[reportAttributeAccessIssue]
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
|
||||
|
||||
@@ -1179,7 +1195,8 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(Step35Model, model)
|
||||
total = len(model.layers)
|
||||
|
||||
@@ -1212,8 +1229,8 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
|
||||
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
|
||||
|
||||
mx.eval(layer)
|
||||
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
|
||||
|
||||
@@ -1221,7 +1238,8 @@ class NemotronHShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(NemotronHModel, model)
|
||||
rank = self.group.rank()
|
||||
total = len(model.layers)
|
||||
@@ -1254,7 +1272,8 @@ class NemotronHShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.mixer = mixer # pyright: ignore[reportAttributeAccessIssue]
|
||||
|
||||
mx.eval(layer)
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
|
||||
def _shard_mamba2_mixer(self, mixer: NemotronHMamba2Mixer, rank: int) -> None:
|
||||
@@ -1361,7 +1380,8 @@ class Gemma4ShardingStrategy(TensorParallelShardingStrategy):
|
||||
def shard_model(
|
||||
self,
|
||||
model: nn.Module,
|
||||
) -> Generator[ModelLoadingResponse, None, nn.Module]:
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> nn.Module:
|
||||
model = cast(Gemma4Model, model)
|
||||
layers = model.language_model.model.layers
|
||||
total = len(layers)
|
||||
@@ -1370,11 +1390,9 @@ class Gemma4ShardingStrategy(TensorParallelShardingStrategy):
|
||||
|
||||
attn = layer.self_attn
|
||||
attn.q_proj = self.all_to_sharded_linear(attn.q_proj)
|
||||
has_kv: bool = cast(bool, attn.has_kv)
|
||||
if has_kv:
|
||||
attn.k_proj = self.all_to_sharded_linear(attn.k_proj)
|
||||
if not attn.use_k_eq_v:
|
||||
attn.v_proj = self.all_to_sharded_linear(attn.v_proj)
|
||||
attn.k_proj = self.all_to_sharded_linear(attn.k_proj)
|
||||
if not attn.use_k_eq_v:
|
||||
attn.v_proj = self.all_to_sharded_linear(attn.v_proj)
|
||||
attn.o_proj = self.sharded_to_all_linear(attn.o_proj)
|
||||
attn.n_heads //= self.N
|
||||
attn.n_kv_heads //= self.N
|
||||
@@ -1391,5 +1409,6 @@ class Gemma4ShardingStrategy(TensorParallelShardingStrategy):
|
||||
layer.experts.sharding_group = self.group
|
||||
|
||||
mx.eval(layer)
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
return model
|
||||
@@ -249,12 +249,7 @@ class KVPrefixCache:
|
||||
# For partial match: trim to best_length, remaining has suffix to prefill
|
||||
# This ensures stream_generate always has at least one token to start with
|
||||
has_ssm = has_non_kv_caches(self.caches[best_index])
|
||||
cached_length = cache_length(self.caches[best_index])
|
||||
if has_ssm:
|
||||
target = best_length
|
||||
else:
|
||||
desired = (max_length - 1) if is_exact else best_length
|
||||
target = min(cached_length, desired)
|
||||
target = (max_length - 1) if is_exact and not has_ssm else best_length
|
||||
restore_pos, restore_snap = self._get_snapshot(best_index, target)
|
||||
|
||||
# No usable snapshot — need fresh cache
|
||||
@@ -262,6 +257,7 @@ class KVPrefixCache:
|
||||
return make_kv_cache(model), prompt_tokens, None, False
|
||||
|
||||
prompt_cache = deepcopy(self.caches[best_index])
|
||||
cached_length = cache_length(self.caches[best_index])
|
||||
tokens_to_trim = cached_length - restore_pos
|
||||
if tokens_to_trim > 0:
|
||||
trim_cache(prompt_cache, tokens_to_trim, restore_snap)
|
||||
|
||||
@@ -17,13 +17,6 @@ TOOL_CALLS_START = f"<{DSML_TOKEN}function_calls>"
|
||||
TOOL_CALLS_END = f"</{DSML_TOKEN}function_calls>"
|
||||
_ORPHAN_THINK_END = ASSISTANT_TOKEN + THINKING_END
|
||||
_FIXED_THINK_BLOCK = ASSISTANT_TOKEN + THINKING_START + "\n" + THINKING_END
|
||||
_FUNCTION_RESULTS_CLOSE = "</function_results>"
|
||||
_ORPHAN_TOOL_RESULT_SUFFIX = _FUNCTION_RESULTS_CLOSE + "\n\n" + THINKING_END
|
||||
_EMPTY_THINK_BLOCKS = (
|
||||
THINKING_START + "\n\n" + THINKING_END,
|
||||
THINKING_START + "\n" + THINKING_END,
|
||||
THINKING_START + THINKING_END,
|
||||
)
|
||||
|
||||
|
||||
def encode_messages(
|
||||
@@ -42,11 +35,7 @@ def encode_messages(
|
||||
add_default_bos_token=add_default_bos_token,
|
||||
tools=tools,
|
||||
)
|
||||
prompt = prompt.replace(_ORPHAN_TOOL_RESULT_SUFFIX, _FUNCTION_RESULTS_CLOSE)
|
||||
prompt = prompt.replace(_ORPHAN_THINK_END, _FIXED_THINK_BLOCK)
|
||||
for empty in _EMPTY_THINK_BLOCKS:
|
||||
prompt = prompt.replace(empty, "")
|
||||
return prompt
|
||||
return prompt.replace(_ORPHAN_THINK_END, _FIXED_THINK_BLOCK)
|
||||
|
||||
|
||||
_INVOKE_PATTERN = re.compile(
|
||||
|
||||
@@ -4,7 +4,6 @@ import re
|
||||
import sys
|
||||
import tempfile
|
||||
import time
|
||||
from collections.abc import Generator
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any, cast
|
||||
|
||||
@@ -52,7 +51,6 @@ from exo.shared.types.worker.instances import (
|
||||
MlxJacclInstance,
|
||||
MlxRingInstance,
|
||||
)
|
||||
from exo.shared.types.worker.runner_response import ModelLoadingResponse
|
||||
from exo.shared.types.worker.shards import (
|
||||
CfgShardMetadata,
|
||||
PipelineShardMetadata,
|
||||
@@ -60,6 +58,7 @@ from exo.shared.types.worker.shards import (
|
||||
TensorShardMetadata,
|
||||
)
|
||||
from exo.worker.engines.mlx.auto_parallel import (
|
||||
LayerLoadedCallback,
|
||||
get_inner_model,
|
||||
get_layers,
|
||||
pipeline_auto_parallel,
|
||||
@@ -67,6 +66,8 @@ from exo.worker.engines.mlx.auto_parallel import (
|
||||
)
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
|
||||
Group = mx.distributed.Group
|
||||
|
||||
|
||||
def get_weights_size(model_shard_meta: ShardMetadata) -> Memory:
|
||||
return Memory.from_float_kb(
|
||||
@@ -89,7 +90,7 @@ class HostList(RootModel[list[str]]):
|
||||
|
||||
def mlx_distributed_init(
|
||||
bound_instance: BoundInstance,
|
||||
) -> mx.distributed.Group:
|
||||
) -> Group:
|
||||
"""
|
||||
Initialize MLX distributed.
|
||||
"""
|
||||
@@ -148,7 +149,7 @@ def mlx_distributed_init(
|
||||
|
||||
def initialize_mlx(
|
||||
bound_instance: BoundInstance,
|
||||
) -> mx.distributed.Group:
|
||||
) -> Group:
|
||||
# should we unseed it?
|
||||
# TODO: pass in seed from params
|
||||
mx.random.seed(42)
|
||||
@@ -161,10 +162,9 @@ def initialize_mlx(
|
||||
|
||||
def load_mlx_items(
|
||||
bound_instance: BoundInstance,
|
||||
group: mx.distributed.Group | None,
|
||||
) -> Generator[
|
||||
ModelLoadingResponse, None, tuple[Model, TokenizerWrapper, "VisionProcessor | None"]
|
||||
]:
|
||||
group: Group | None,
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> "tuple[Model, TokenizerWrapper, VisionProcessor | None]":
|
||||
if group is None:
|
||||
logger.info(f"Single device used for {bound_instance.instance}")
|
||||
model_path = build_model_path(bound_instance.bound_shard.model_card.model_id)
|
||||
@@ -177,7 +177,8 @@ def load_mlx_items(
|
||||
total = len(layers)
|
||||
for i, layer in enumerate(layers):
|
||||
mx.eval(layer) # type: ignore
|
||||
yield ModelLoadingResponse(layers_loaded=i, total=total)
|
||||
if on_layer_loaded is not None:
|
||||
on_layer_loaded(i, total)
|
||||
except ValueError as e:
|
||||
logger.opt(exception=e).debug(
|
||||
"Model architecture doesn't support layer-by-layer progress tracking",
|
||||
@@ -190,9 +191,10 @@ def load_mlx_items(
|
||||
else:
|
||||
logger.info("Starting distributed init")
|
||||
start_time = time.perf_counter()
|
||||
model, tokenizer = yield from shard_and_load(
|
||||
model, tokenizer = shard_and_load(
|
||||
bound_instance.bound_shard,
|
||||
group=group,
|
||||
on_layer_loaded=on_layer_loaded,
|
||||
)
|
||||
end_time = time.perf_counter()
|
||||
logger.info(
|
||||
@@ -208,20 +210,9 @@ def load_mlx_items(
|
||||
if vision_config is not None:
|
||||
from exo.worker.engines.mlx.vision import VisionProcessor
|
||||
|
||||
vision_start_time = time.perf_counter()
|
||||
try:
|
||||
vision_processor: VisionProcessor | None = VisionProcessor(
|
||||
vision_config, bound_instance.bound_shard.model_card.model_id
|
||||
)
|
||||
vision_processor.load()
|
||||
logger.info(
|
||||
f"Time taken to load vision weights: {(time.perf_counter() - vision_start_time):.2f}s"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.opt(exception=e).error(
|
||||
"Failed to load vision weights — disabling vision for this runner"
|
||||
)
|
||||
vision_processor = None
|
||||
vision_processor: VisionProcessor | None = VisionProcessor(
|
||||
vision_config, bound_instance.bound_shard.model_card.model_id
|
||||
)
|
||||
else:
|
||||
vision_processor = None
|
||||
|
||||
@@ -230,8 +221,9 @@ def load_mlx_items(
|
||||
|
||||
def shard_and_load(
|
||||
shard_metadata: ShardMetadata,
|
||||
group: mx.distributed.Group,
|
||||
) -> Generator[ModelLoadingResponse, None, tuple[nn.Module, TokenizerWrapper]]:
|
||||
group: Group,
|
||||
on_layer_loaded: LayerLoadedCallback | None,
|
||||
) -> tuple[nn.Module, TokenizerWrapper]:
|
||||
model_path = build_model_path(shard_metadata.model_card.model_id)
|
||||
|
||||
model, _ = load_model(model_path, lazy=True, strict=False)
|
||||
@@ -262,10 +254,12 @@ def shard_and_load(
|
||||
match shard_metadata:
|
||||
case TensorShardMetadata():
|
||||
logger.info(f"loading model from {model_path} with tensor parallelism")
|
||||
model = yield from tensor_auto_parallel(model, group)
|
||||
model = tensor_auto_parallel(model, group, on_layer_loaded)
|
||||
case PipelineShardMetadata():
|
||||
logger.info(f"loading model from {model_path} with pipeline parallelism")
|
||||
model = yield from pipeline_auto_parallel(model, group, shard_metadata)
|
||||
model = pipeline_auto_parallel(
|
||||
model, group, shard_metadata, on_layer_loaded=on_layer_loaded
|
||||
)
|
||||
mx.eval(model.parameters())
|
||||
case CfgShardMetadata():
|
||||
raise ValueError(
|
||||
@@ -547,6 +541,7 @@ def render_chat_template(
|
||||
)
|
||||
if partial_assistant_content:
|
||||
prompt += partial_assistant_content
|
||||
logger.info(prompt)
|
||||
return prompt
|
||||
|
||||
for msg in formatted_messages:
|
||||
@@ -753,9 +748,7 @@ def set_wired_limit_for_model(model_size: Memory):
|
||||
|
||||
|
||||
def mlx_cleanup(
|
||||
model: Model | None,
|
||||
tokenizer: TokenizerWrapper | None,
|
||||
group: mx.distributed.Group | None,
|
||||
model: Model | None, tokenizer: TokenizerWrapper | None, group: Group | None
|
||||
) -> None:
|
||||
del model, tokenizer, group
|
||||
mx.clear_cache()
|
||||
@@ -764,7 +757,7 @@ def mlx_cleanup(
|
||||
gc.collect()
|
||||
|
||||
|
||||
def mx_any(bool_: bool, group: mx.distributed.Group | None) -> bool:
|
||||
def mx_any(bool_: bool, group: Group | None) -> bool:
|
||||
if group is None:
|
||||
return bool_
|
||||
num_true = mx.distributed.all_sum(
|
||||
@@ -774,7 +767,7 @@ def mx_any(bool_: bool, group: mx.distributed.Group | None) -> bool:
|
||||
return num_true.item() > 0
|
||||
|
||||
|
||||
def mx_barrier(group: mx.distributed.Group | None):
|
||||
def mx_barrier(group: Group | None):
|
||||
if group is None:
|
||||
return
|
||||
mx.eval(
|
||||
|
||||
@@ -36,19 +36,6 @@ from exo.worker.runner.bootstrap import logger
|
||||
|
||||
_video_processor_patched = False
|
||||
|
||||
_MLX_VLM_MODEL_TYPE_ALIASES: dict[str, str] = {
|
||||
"kimi_k25": "kimi_vl",
|
||||
"kimi_k26": "kimi_vl",
|
||||
}
|
||||
|
||||
|
||||
def _torch_tensor_to_mx(
|
||||
tensor: Any, # pyright: ignore[reportAny]
|
||||
) -> mx.array:
|
||||
if str(tensor.dtype) == "torch.bfloat16": # type: ignore
|
||||
return mx.array(tensor.float().numpy(), dtype=mx.bfloat16) # type: ignore
|
||||
return mx.array(tensor.numpy()) # type: ignore
|
||||
|
||||
|
||||
def _filter_config(cls: type, d: dict[str, Any]) -> dict[str, Any]:
|
||||
valid = set(inspect.signature(cls.__init__).parameters.keys()) - {"self"}
|
||||
@@ -98,8 +85,6 @@ def _instantiate_projector(
|
||||
params = {n: p for n, p in init_sig.parameters.items() if n != "self"}
|
||||
kwargs: dict[str, Any] = {}
|
||||
|
||||
if "config" in params:
|
||||
kwargs["config"] = model_config
|
||||
if "embedding_dim" in params:
|
||||
kwargs["embedding_dim"] = vision_config.hidden_size # pyright: ignore[reportAny]
|
||||
if "text_hidden_size" in params:
|
||||
@@ -220,9 +205,7 @@ class VisionEncoder:
|
||||
return {}
|
||||
|
||||
def _import_mlx_vlm(self, *submodules: str) -> Any: # type: ignore
|
||||
mt = _MLX_VLM_MODEL_TYPE_ALIASES.get(
|
||||
self._config.model_type, self._config.model_type
|
||||
)
|
||||
mt = self._config.model_type
|
||||
results: list[Any] = []
|
||||
for sub in submodules:
|
||||
name = f"mlx_vlm.models.{mt}.{sub}"
|
||||
@@ -255,7 +238,7 @@ class VisionEncoder:
|
||||
def _load_image_processor_from_module(self, repo: str) -> "ImageProcessor | None":
|
||||
# mlx_vlm.utils.load_image_processor only works for models that set
|
||||
# `Model.ImageProcessor = <cls>`, but Gemma4 just uses
|
||||
# `Gemma4ImageProcessor` from the package `__init__.py`.
|
||||
# `Gemma4ImageProcessor` from the package `__init__.py`
|
||||
try:
|
||||
pkg: Any = importlib.import_module(
|
||||
f"mlx_vlm.models.{self._config.model_type}"
|
||||
@@ -336,16 +319,10 @@ class VisionEncoder:
|
||||
else:
|
||||
self._load_weights_from_model_repo()
|
||||
|
||||
if processor_repo:
|
||||
repo = str(build_model_path(ModelId(processor_repo)))
|
||||
else:
|
||||
repo = str(self._model_path)
|
||||
try:
|
||||
image_proc = load_image_processor(repo)
|
||||
except ValueError:
|
||||
image_proc = None
|
||||
if image_proc is None:
|
||||
image_proc = self._load_image_processor_from_module(repo)
|
||||
repo = processor_repo or str(self._model_path)
|
||||
image_proc = load_image_processor(
|
||||
repo
|
||||
) or self._load_image_processor_from_module(repo)
|
||||
if image_proc is not None:
|
||||
self._processor = image_proc
|
||||
else:
|
||||
@@ -362,42 +339,39 @@ class VisionEncoder:
|
||||
if not safetensors_files:
|
||||
raise FileNotFoundError(f"No safetensors files found in {self._model_path}")
|
||||
|
||||
vision_weights: dict[str, mx.array] = {}
|
||||
projector_weights: dict[str, mx.array] = {}
|
||||
|
||||
weights: dict[str, mx.array] = {}
|
||||
for sf_path in safetensors_files:
|
||||
with safe_open(str(sf_path), framework="pt") as f:
|
||||
keys = cast(list[str], list(f.keys())) # type: ignore
|
||||
keys = f.keys()
|
||||
for key in keys:
|
||||
if key.startswith("vision_tower."):
|
||||
short_key = key[len("vision_tower.") :]
|
||||
if short_key.startswith("encoder."):
|
||||
short_key = short_key[len("encoder.") :]
|
||||
m = re.match(
|
||||
r"^(blocks\.\d+)\.(wqkv|wo)\.(weight|bias)$", short_key
|
||||
)
|
||||
if m:
|
||||
short_key = f"{m.group(1)}.attn.{m.group(2)}.{m.group(3)}"
|
||||
tensor = f.get_tensor(key) # type: ignore
|
||||
val = mx.array(tensor.float().numpy(), dtype=mx.bfloat16) # type: ignore
|
||||
if short_key == "patch_embed.proj.weight" and val.ndim == 4:
|
||||
val = val.transpose(0, 2, 3, 1)
|
||||
vision_weights[short_key] = val
|
||||
elif key.startswith(("mm_projector.", "multi_modal_projector.")):
|
||||
if key.startswith("multi_modal_projector."):
|
||||
short_key = key[len("multi_modal_projector.") :]
|
||||
if short_key.startswith("mm_projector."):
|
||||
short_key = short_key[len("mm_projector.") :]
|
||||
else:
|
||||
short_key = key[len("mm_projector.") :]
|
||||
short_key = short_key.replace("proj.0.", "linear_1.").replace(
|
||||
"proj.2.", "linear_2."
|
||||
)
|
||||
tensor = f.get_tensor(key) # type: ignore
|
||||
projector_weights[short_key] = mx.array(
|
||||
tensor.float().numpy(), # type: ignore
|
||||
dtype=mx.bfloat16,
|
||||
)
|
||||
tensor = f.get_tensor(key) # type: ignore
|
||||
np_tensor = tensor.float().numpy() # type: ignore
|
||||
weights[key] = mx.array(np_tensor, dtype=mx.bfloat16) # type: ignore
|
||||
|
||||
vision_weights: dict[str, mx.array] = {}
|
||||
projector_weights: dict[str, mx.array] = {}
|
||||
for key, val in weights.items():
|
||||
if key.startswith("vision_tower."):
|
||||
short_key = key[len("vision_tower.") :]
|
||||
if short_key.startswith("encoder."):
|
||||
short_key = short_key[len("encoder.") :]
|
||||
m = re.match(r"^(blocks\.\d+)\.(wqkv|wo)\.(weight|bias)$", short_key)
|
||||
if m:
|
||||
short_key = f"{m.group(1)}.attn.{m.group(2)}.{m.group(3)}"
|
||||
if short_key == "patch_embed.proj.weight" and val.ndim == 4:
|
||||
val = val.transpose(0, 2, 3, 1)
|
||||
vision_weights[short_key] = val
|
||||
elif key.startswith(("mm_projector.", "multi_modal_projector.")):
|
||||
if key.startswith("multi_modal_projector."):
|
||||
short_key = key[len("multi_modal_projector.") :]
|
||||
if short_key.startswith("mm_projector."):
|
||||
short_key = short_key[len("mm_projector.") :]
|
||||
else:
|
||||
short_key = key[len("mm_projector.") :]
|
||||
short_key = short_key.replace("proj.0.", "linear_1.").replace(
|
||||
"proj.2.", "linear_2."
|
||||
)
|
||||
projector_weights[short_key] = val
|
||||
|
||||
assert self._vision_tower is not None
|
||||
self._vision_tower.load_weights(list(vision_weights.items()))
|
||||
@@ -433,26 +407,18 @@ class VisionEncoder:
|
||||
needs_sanitize = False
|
||||
|
||||
for sf_path in safetensors_files:
|
||||
with safe_open(str(sf_path), framework="pt") as f:
|
||||
keys = cast(list[str], list(f.keys())) # type: ignore
|
||||
for key in keys:
|
||||
matched = False
|
||||
for prefix in vision_prefixes:
|
||||
if key.startswith(prefix):
|
||||
vision_weights[key[len(prefix) :]] = _torch_tensor_to_mx(
|
||||
f.get_tensor(key)
|
||||
)
|
||||
if prefix == "model.visual.":
|
||||
needs_sanitize = True
|
||||
matched = True
|
||||
break
|
||||
if matched:
|
||||
continue
|
||||
file_weights: dict[str, mx.array] = mx.load(str(sf_path)) # type: ignore
|
||||
for key, val in file_weights.items():
|
||||
for prefix in vision_prefixes:
|
||||
if key.startswith(prefix):
|
||||
vision_weights[key[len(prefix) :]] = val
|
||||
if prefix == "model.visual.":
|
||||
needs_sanitize = True
|
||||
break
|
||||
else:
|
||||
for prefix in projector_prefixes:
|
||||
if key.startswith(prefix):
|
||||
projector_weights[key[len(prefix) :]] = _torch_tensor_to_mx(
|
||||
f.get_tensor(key)
|
||||
)
|
||||
projector_weights[key[len(prefix) :]] = val
|
||||
break
|
||||
|
||||
if not vision_weights:
|
||||
@@ -497,12 +463,7 @@ class VisionEncoder:
|
||||
grid_thw: mx.array | None
|
||||
n_tokens_per_image: list[int]
|
||||
|
||||
is_kimi_vl_processor = any(
|
||||
"mlx_vlm.models.kimi_vl" in cls.__module__
|
||||
for cls in type(self._processor).__mro__
|
||||
)
|
||||
|
||||
if self._config.processor_repo and not is_kimi_vl_processor:
|
||||
if self._config.processor_repo:
|
||||
processed = self._processor.preprocess(
|
||||
[{"type": "image", "image": img} for img in pil_images],
|
||||
return_tensors="np",
|
||||
@@ -520,24 +481,6 @@ class VisionEncoder:
|
||||
int(mx.prod(grid_thw[i]).item()) // merge_length
|
||||
for i in range(grid_thw.shape[0])
|
||||
]
|
||||
elif is_kimi_vl_processor:
|
||||
proc: Any = self._processor
|
||||
raw_processed = proc.preprocess(pil_images, return_tensors="np") # type: ignore
|
||||
stacked_pixels = mx.array(raw_processed["pixel_values"]) # type: ignore
|
||||
if stacked_pixels.ndim == 3:
|
||||
stacked_pixels = stacked_pixels[None]
|
||||
per_image_pixels = [
|
||||
stacked_pixels[i : i + 1] for i in range(stacked_pixels.shape[0])
|
||||
]
|
||||
grid_raw = raw_processed.get("image_grid_hws") # type: ignore
|
||||
if grid_raw is None:
|
||||
grid_raw = raw_processed["grid_thws"] # type: ignore
|
||||
grid_thw = mx.array(grid_raw) # type: ignore
|
||||
merge_length = int(np.prod(self._merge_kernel_size or [2, 2]))
|
||||
n_tokens_per_image = [
|
||||
int(mx.prod(grid_thw[i]).item()) // merge_length
|
||||
for i in range(grid_thw.shape[0])
|
||||
]
|
||||
else:
|
||||
batch, tokens_override = _run_processor(self._processor, pil_images)
|
||||
# `Gemma4ImageProcessor` returns pixel_values as a plain ndarray
|
||||
|
||||
+35
-25
@@ -152,26 +152,6 @@ class Worker:
|
||||
event.chunk
|
||||
)
|
||||
|
||||
if (
|
||||
len(self.input_chunk_buffer[cmd_id])
|
||||
== self.input_chunk_counts[cmd_id]
|
||||
):
|
||||
per_image: defaultdict[int, list[InputImageChunk]] = (
|
||||
defaultdict(list)
|
||||
)
|
||||
for chunk in self.input_chunk_buffer[cmd_id].values():
|
||||
per_image[chunk.image_index].append(chunk)
|
||||
for chunks_for_image in per_image.values():
|
||||
sorted_chunks = sorted(
|
||||
chunks_for_image, key=lambda c: c.chunk_index
|
||||
)
|
||||
img = Base64Image("".join(c.data for c in sorted_chunks))
|
||||
self.image_cache[
|
||||
Base64ImageHash(
|
||||
hashlib.sha256(img.encode("ascii")).hexdigest()
|
||||
)
|
||||
] = img
|
||||
|
||||
if isinstance(event, CustomModelCardAdded):
|
||||
await event.model_card.save_to_custom_dir()
|
||||
add_to_card_cache(event.model_card)
|
||||
@@ -190,7 +170,6 @@ class Worker:
|
||||
self.state.runners,
|
||||
self.state.tasks,
|
||||
self.input_chunk_buffer,
|
||||
self.image_cache,
|
||||
self._instance_backoff,
|
||||
self._download_backoff,
|
||||
)
|
||||
@@ -230,7 +209,7 @@ class Worker:
|
||||
self._download_backoff.record_attempt(model_id)
|
||||
|
||||
found_path = await to_thread.run_sync(
|
||||
resolve_existing_model, model_id, shard.model_card
|
||||
resolve_existing_model, model_id
|
||||
)
|
||||
if found_path is not None:
|
||||
logger.info(f"Model {model_id} found at {found_path}")
|
||||
@@ -328,11 +307,42 @@ class Worker:
|
||||
del self.input_chunk_counts[cmd_id]
|
||||
await self._start_runner_task(modified_task)
|
||||
|
||||
case TextGeneration() if task.task_params.image_hashes:
|
||||
case TextGeneration() if (
|
||||
task.task_params.image_hashes
|
||||
or task.task_params.total_input_chunks > 0
|
||||
):
|
||||
cmd_id = task.command_id
|
||||
by_index: dict[int, Base64Image] = {}
|
||||
|
||||
for idx, h in task.task_params.image_hashes.items():
|
||||
assert h in self.image_cache
|
||||
by_index[idx] = self.image_cache[h]
|
||||
|
||||
if task.task_params.total_input_chunks > 0:
|
||||
chunk_buffer = self.input_chunk_buffer.get(cmd_id, {})
|
||||
per_image: defaultdict[int, list[InputImageChunk]] = (
|
||||
defaultdict(list)
|
||||
)
|
||||
for chunk in chunk_buffer.values():
|
||||
per_image[chunk.image_index].append(chunk)
|
||||
for img_idx in sorted(per_image):
|
||||
sorted_chunks = sorted(
|
||||
per_image[img_idx], key=lambda c: c.chunk_index
|
||||
)
|
||||
img = Base64Image("".join(c.data for c in sorted_chunks))
|
||||
self.image_cache[
|
||||
Base64ImageHash(
|
||||
hashlib.sha256(img.encode("ascii")).hexdigest()
|
||||
)
|
||||
] = img
|
||||
by_index[img_idx] = img
|
||||
logger.info(
|
||||
f"Assembled {len(per_image)} VLM image(s) "
|
||||
f"from {len(chunk_buffer)} chunks"
|
||||
)
|
||||
|
||||
resolved_images = [
|
||||
self.image_cache[h]
|
||||
for _, h in sorted(task.task_params.image_hashes.items())
|
||||
Base64Image(by_index[i]) for i in sorted(by_index)
|
||||
]
|
||||
modified_task = task.model_copy(
|
||||
update={
|
||||
|
||||
+9
-16
@@ -19,7 +19,6 @@ from exo.shared.types.tasks import (
|
||||
TaskStatus,
|
||||
TextGeneration,
|
||||
)
|
||||
from exo.shared.types.text_generation import Base64Image, Base64ImageHash
|
||||
from exo.shared.types.worker.downloads import (
|
||||
DownloadCompleted,
|
||||
DownloadFailed,
|
||||
@@ -53,7 +52,6 @@ def plan(
|
||||
all_runners: Mapping[RunnerId, RunnerStatus], # all global
|
||||
tasks: Mapping[TaskId, Task],
|
||||
input_chunk_buffer: Mapping[CommandId, Mapping[int, InputImageChunk]],
|
||||
image_cache: Mapping[Base64ImageHash, Base64Image],
|
||||
instance_backoff: KeyedBackoff[InstanceId],
|
||||
download_backoff: KeyedBackoff[ModelId],
|
||||
) -> Task | None:
|
||||
@@ -68,7 +66,7 @@ def plan(
|
||||
or _init_distributed_backend(runners, all_runners)
|
||||
or _load_model(runners, all_runners, global_download_status)
|
||||
or _ready_to_warmup(runners, all_runners)
|
||||
or _pending_tasks(runners, tasks, all_runners, input_chunk_buffer, image_cache)
|
||||
or _pending_tasks(runners, tasks, all_runners, input_chunk_buffer)
|
||||
)
|
||||
|
||||
|
||||
@@ -302,7 +300,6 @@ def _pending_tasks(
|
||||
tasks: Mapping[TaskId, Task],
|
||||
all_runners: Mapping[RunnerId, RunnerStatus],
|
||||
input_chunk_buffer: Mapping[CommandId, Mapping[int, InputImageChunk]],
|
||||
image_cache: Mapping[Base64ImageHash, Base64Image],
|
||||
) -> Task | None:
|
||||
for task in tasks.values():
|
||||
# for now, just forward chat completions
|
||||
@@ -312,20 +309,16 @@ def _pending_tasks(
|
||||
if task.task_status not in (TaskStatus.Pending, TaskStatus.Running):
|
||||
continue
|
||||
|
||||
if isinstance(task, ImageEdits) and task.task_params.total_input_chunks > 0:
|
||||
received = len(input_chunk_buffer.get(task.command_id, {}))
|
||||
if received < task.task_params.total_input_chunks:
|
||||
# For tasks with images, verify all input chunks have been received
|
||||
expected_image_chunks = 0
|
||||
if isinstance(task, (ImageEdits, TextGeneration)):
|
||||
expected_image_chunks = task.task_params.total_input_chunks
|
||||
if expected_image_chunks > 0:
|
||||
cmd_id = task.command_id
|
||||
received = len(input_chunk_buffer.get(cmd_id, {}))
|
||||
if received < expected_image_chunks:
|
||||
continue # Wait for all chunks to arrive
|
||||
|
||||
if (
|
||||
isinstance(task, TextGeneration)
|
||||
and task.task_params.image_hashes
|
||||
and not all(
|
||||
h in image_cache for h in task.task_params.image_hashes.values()
|
||||
)
|
||||
):
|
||||
continue # Wait for all images to be assembled into the cache
|
||||
|
||||
for runner in runners.values():
|
||||
if task.instance_id != runner.bound_instance.instance.instance_id:
|
||||
continue
|
||||
|
||||
@@ -1,17 +1,19 @@
|
||||
import base64
|
||||
import time
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Literal
|
||||
|
||||
import mlx.core as mx
|
||||
|
||||
from exo.api.types import (
|
||||
ImageEditsTaskParams,
|
||||
ImageGenerationStats,
|
||||
ImageGenerationTaskParams,
|
||||
)
|
||||
from exo.shared.constants import EXO_TRACING_ENABLED
|
||||
from exo.shared.constants import EXO_MAX_CHUNK_SIZE, EXO_TRACING_ENABLED
|
||||
from exo.shared.models.model_cards import ModelTask
|
||||
from exo.shared.tracing import clear_trace_buffer, get_trace_buffer
|
||||
from exo.shared.types.chunks import ErrorChunk
|
||||
from exo.shared.types.common import CommandId
|
||||
from exo.shared.types.chunks import ErrorChunk, ImageChunk
|
||||
from exo.shared.types.common import CommandId, ModelId
|
||||
from exo.shared.types.events import (
|
||||
ChunkGenerated,
|
||||
Event,
|
||||
@@ -34,6 +36,10 @@ from exo.shared.types.tasks import (
|
||||
TaskStatus,
|
||||
)
|
||||
from exo.shared.types.worker.instances import BoundInstance
|
||||
from exo.shared.types.worker.runner_response import (
|
||||
ImageGenerationResponse,
|
||||
PartialImageResponse,
|
||||
)
|
||||
from exo.shared.types.worker.runners import (
|
||||
RunnerConnected,
|
||||
RunnerConnecting,
|
||||
@@ -81,6 +87,32 @@ def _is_primary_output_node(shard_metadata: ShardMetadata) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _process_image_response(
|
||||
response: ImageGenerationResponse | PartialImageResponse,
|
||||
command_id: CommandId,
|
||||
shard_metadata: ShardMetadata,
|
||||
event_sender: MpSender[Event],
|
||||
image_index: int,
|
||||
) -> None:
|
||||
"""Process a single image response and send chunks."""
|
||||
encoded_data = base64.b64encode(response.image_data).decode("utf-8")
|
||||
is_partial = isinstance(response, PartialImageResponse)
|
||||
# Extract stats from final ImageGenerationResponse if available
|
||||
stats = response.stats if isinstance(response, ImageGenerationResponse) else None
|
||||
_send_image_chunk(
|
||||
encoded_data=encoded_data,
|
||||
command_id=command_id,
|
||||
model_id=shard_metadata.model_card.model_id,
|
||||
event_sender=event_sender,
|
||||
image_index=response.image_index,
|
||||
is_partial=is_partial,
|
||||
partial_index=response.partial_index if is_partial else None,
|
||||
total_partials=response.total_partials if is_partial else None,
|
||||
stats=stats,
|
||||
image_format=response.format,
|
||||
)
|
||||
|
||||
|
||||
def _send_traces_if_enabled(
|
||||
event_sender: MpSender[Event],
|
||||
task_id: TaskId,
|
||||
@@ -111,6 +143,48 @@ def _send_traces_if_enabled(
|
||||
clear_trace_buffer()
|
||||
|
||||
|
||||
def _send_image_chunk(
|
||||
encoded_data: str,
|
||||
command_id: CommandId,
|
||||
model_id: ModelId,
|
||||
event_sender: MpSender[Event],
|
||||
image_index: int,
|
||||
is_partial: bool,
|
||||
partial_index: int | None = None,
|
||||
total_partials: int | None = None,
|
||||
stats: ImageGenerationStats | None = None,
|
||||
image_format: Literal["png", "jpeg", "webp"] | None = None,
|
||||
) -> None:
|
||||
"""Send base64-encoded image data as chunks via events."""
|
||||
data_chunks = [
|
||||
encoded_data[i : i + EXO_MAX_CHUNK_SIZE]
|
||||
for i in range(0, len(encoded_data), EXO_MAX_CHUNK_SIZE)
|
||||
]
|
||||
total_chunks = len(data_chunks)
|
||||
for chunk_index, chunk_data in enumerate(data_chunks):
|
||||
# Only include stats on the last chunk of the final image
|
||||
chunk_stats = (
|
||||
stats if chunk_index == total_chunks - 1 and not is_partial else None
|
||||
)
|
||||
event_sender.send(
|
||||
ChunkGenerated(
|
||||
command_id=command_id,
|
||||
chunk=ImageChunk(
|
||||
model=model_id,
|
||||
data=chunk_data,
|
||||
chunk_index=chunk_index,
|
||||
total_chunks=total_chunks,
|
||||
image_index=image_index,
|
||||
is_partial=is_partial,
|
||||
partial_index=partial_index,
|
||||
total_partials=total_partials,
|
||||
stats=chunk_stats,
|
||||
format=image_format,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
class Runner:
|
||||
def __init__(
|
||||
self,
|
||||
@@ -187,21 +261,35 @@ class Runner:
|
||||
return self._check_cancelled(task.task_id)
|
||||
|
||||
try:
|
||||
for chunk in generate_image(
|
||||
image_index = 0
|
||||
for response in generate_image(
|
||||
model=self.image_model,
|
||||
task=task_params,
|
||||
cancel_checker=cancel_checker,
|
||||
):
|
||||
if _is_primary_output_node(self.shard_metadata):
|
||||
if chunk.is_partial:
|
||||
logger.info(
|
||||
f"sending partial ImageChunk {chunk.partial_index}/{chunk.total_partials}"
|
||||
)
|
||||
else:
|
||||
logger.info("sending final ImageChunk")
|
||||
self.event_sender.send(
|
||||
ChunkGenerated(command_id=command_id, chunk=chunk)
|
||||
)
|
||||
match response:
|
||||
case PartialImageResponse():
|
||||
logger.info(
|
||||
f"sending partial ImageChunk {response.partial_index}/{response.total_partials}"
|
||||
)
|
||||
_process_image_response(
|
||||
response,
|
||||
command_id,
|
||||
self.shard_metadata,
|
||||
self.event_sender,
|
||||
image_index,
|
||||
)
|
||||
case ImageGenerationResponse():
|
||||
logger.info("sending final ImageChunk")
|
||||
_process_image_response(
|
||||
response,
|
||||
command_id,
|
||||
self.shard_metadata,
|
||||
self.event_sender,
|
||||
image_index,
|
||||
)
|
||||
image_index += 1
|
||||
except Exception as e:
|
||||
if _is_primary_output_node(self.shard_metadata):
|
||||
self.event_sender.send(
|
||||
|
||||
@@ -2,20 +2,20 @@ import itertools
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import deque
|
||||
from collections.abc import Generator, Iterator
|
||||
from collections.abc import Generator, Iterable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import mlx.core as mx
|
||||
from mlx_lm.tokenizer_utils import TokenizerWrapper
|
||||
|
||||
from exo.shared.constants import EXO_MAX_CONCURRENT_REQUESTS
|
||||
from exo.shared.types.chunks import ErrorChunk, GenerationChunk, PrefillProgressChunk
|
||||
from exo.shared.types.chunks import ErrorChunk, PrefillProgressChunk
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.events import ChunkGenerated, Event
|
||||
from exo.shared.types.mlx import Model
|
||||
from exo.shared.types.tasks import CANCEL_ALL_TASKS, TaskId, TextGeneration
|
||||
from exo.shared.types.text_generation import TextGenerationTaskParams
|
||||
from exo.shared.types.worker.runner_response import GenerationResponse
|
||||
from exo.shared.types.worker.runner_response import GenerationResponse, ToolCallResponse
|
||||
from exo.utils.channels import MpReceiver, MpSender
|
||||
from exo.worker.engines.mlx.cache import KVPrefixCache
|
||||
from exo.worker.engines.mlx.generator.batch_generate import ExoBatchGenerator
|
||||
@@ -32,7 +32,7 @@ from exo.worker.engines.mlx.utils_mlx import (
|
||||
from exo.worker.engines.mlx.vision import VisionProcessor
|
||||
from exo.worker.runner.bootstrap import logger
|
||||
|
||||
from .model_output_parsers import apply_all_parsers, map_responses_to_chunks
|
||||
from .model_output_parsers import apply_all_parsers
|
||||
from .tool_parsers import ToolParser
|
||||
|
||||
|
||||
@@ -80,7 +80,9 @@ class InferenceGenerator(ABC):
|
||||
@abstractmethod
|
||||
def step(
|
||||
self,
|
||||
) -> Iterator[tuple[TaskId, GenerationChunk | Cancelled | Finished]]: ...
|
||||
) -> Iterable[
|
||||
tuple[TaskId, ToolCallResponse | GenerationResponse | Cancelled | Finished]
|
||||
]: ...
|
||||
|
||||
@abstractmethod
|
||||
def close(self) -> None: ...
|
||||
@@ -135,7 +137,7 @@ class SequentialGenerator(InferenceGenerator):
|
||||
# queue that the 1st generator should push to and 3rd generator should pull from
|
||||
GeneratorQueue[GenerationResponse],
|
||||
# generator to get parsed outputs
|
||||
Iterator[GenerationChunk | None],
|
||||
Generator[GenerationResponse | ToolCallResponse | None],
|
||||
]
|
||||
| None
|
||||
) = field(default=None, init=False)
|
||||
@@ -181,7 +183,9 @@ class SequentialGenerator(InferenceGenerator):
|
||||
|
||||
def step(
|
||||
self,
|
||||
) -> Iterator[tuple[TaskId, GenerationChunk | Cancelled | Finished]]:
|
||||
) -> Iterable[
|
||||
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
|
||||
]:
|
||||
if self._active is None:
|
||||
self.agree_on_tasks()
|
||||
|
||||
@@ -193,7 +197,9 @@ class SequentialGenerator(InferenceGenerator):
|
||||
assert self._active is not None
|
||||
|
||||
task, mlx_gen, queue, output_generator = self._active
|
||||
output: list[tuple[TaskId, GenerationChunk | Cancelled | Finished]] = []
|
||||
output: list[
|
||||
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
|
||||
] = []
|
||||
try:
|
||||
response = next(mlx_gen)
|
||||
queue.push(response)
|
||||
@@ -227,9 +233,7 @@ class SequentialGenerator(InferenceGenerator):
|
||||
queue = GeneratorQueue[GenerationResponse]()
|
||||
|
||||
if task.task_params.bench:
|
||||
output_generator: Iterator[GenerationChunk | None] = map(
|
||||
lambda r: map_responses_to_chunks(r, self.model_id), queue.gen()
|
||||
)
|
||||
output_generator = queue.gen()
|
||||
else:
|
||||
output_generator = apply_all_parsers(
|
||||
queue.gen(),
|
||||
@@ -334,7 +338,7 @@ class BatchGenerator(InferenceGenerator):
|
||||
tuple[
|
||||
TextGeneration,
|
||||
GeneratorQueue[GenerationResponse],
|
||||
Iterator[GenerationChunk | None],
|
||||
Generator[GenerationResponse | ToolCallResponse | None],
|
||||
],
|
||||
] = field(default_factory=dict, init=False)
|
||||
|
||||
@@ -388,7 +392,9 @@ class BatchGenerator(InferenceGenerator):
|
||||
|
||||
def step(
|
||||
self,
|
||||
) -> Iterator[tuple[TaskId, GenerationChunk | Cancelled | Finished]]:
|
||||
) -> Iterable[
|
||||
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
|
||||
]:
|
||||
if not self._queue:
|
||||
self.agree_on_tasks()
|
||||
|
||||
@@ -405,9 +411,7 @@ class BatchGenerator(InferenceGenerator):
|
||||
|
||||
queue = GeneratorQueue[GenerationResponse]()
|
||||
if task.task_params.bench:
|
||||
output_generator: Iterator[GenerationChunk | None] = map(
|
||||
lambda r: map_responses_to_chunks(r, self.model_id), queue.gen()
|
||||
)
|
||||
output_generator = queue.gen()
|
||||
else:
|
||||
output_generator = apply_all_parsers(
|
||||
queue.gen(),
|
||||
@@ -425,7 +429,9 @@ class BatchGenerator(InferenceGenerator):
|
||||
|
||||
results = self._mlx_gen.step()
|
||||
|
||||
output: list[tuple[TaskId, GenerationChunk | Cancelled | Finished]] = []
|
||||
output: list[
|
||||
tuple[TaskId, GenerationResponse | ToolCallResponse | Cancelled | Finished]
|
||||
] = []
|
||||
for uid, response in results:
|
||||
if uid not in self._active_tasks:
|
||||
# should we error here?
|
||||
@@ -447,9 +453,9 @@ class BatchGenerator(InferenceGenerator):
|
||||
|
||||
def _apply_cancellations(
|
||||
self,
|
||||
) -> Iterator[tuple[TaskId, Cancelled]]:
|
||||
) -> list[tuple[TaskId, Cancelled]]:
|
||||
if not self._cancelled_tasks:
|
||||
return iter([])
|
||||
return []
|
||||
|
||||
cancel_all = CANCEL_ALL_TASKS in self._cancelled_tasks
|
||||
|
||||
@@ -471,7 +477,7 @@ class BatchGenerator(InferenceGenerator):
|
||||
results.append((tid, Cancelled()))
|
||||
|
||||
self._cancelled_tasks.clear()
|
||||
return iter(results)
|
||||
return results
|
||||
|
||||
def _send_error(self, task: TextGeneration, e: Exception) -> None:
|
||||
if self.device_rank == 0:
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from collections.abc import Generator, Iterator
|
||||
from collections.abc import Generator
|
||||
from functools import cache
|
||||
from typing import Any
|
||||
|
||||
@@ -14,12 +14,6 @@ from openai_harmony import ( # pyright: ignore[reportMissingTypeStubs]
|
||||
)
|
||||
|
||||
from exo.api.types import ToolCallItem
|
||||
from exo.shared.types.chunks import (
|
||||
ErrorChunk,
|
||||
GenerationChunk,
|
||||
TokenChunk,
|
||||
ToolCallChunk,
|
||||
)
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.mlx import Model
|
||||
from exo.shared.types.worker.runner_response import GenerationResponse, ToolCallResponse
|
||||
@@ -70,77 +64,29 @@ def apply_all_parsers(
|
||||
model_type: type[Model],
|
||||
model_id: ModelId,
|
||||
tools: list[dict[str, Any]] | None,
|
||||
) -> Iterator[GenerationChunk | None]:
|
||||
generator = receiver
|
||||
) -> Generator[GenerationResponse | ToolCallResponse | None]:
|
||||
mlx_generator = receiver
|
||||
|
||||
if issubclass(model_type, GptOssModel):
|
||||
generator = parse_gpt_oss(generator)
|
||||
mlx_generator = parse_gpt_oss(mlx_generator)
|
||||
elif (
|
||||
issubclass(model_type, DeepseekV32Model)
|
||||
and "deepseek" in model_id.normalize().lower()
|
||||
):
|
||||
if tokenizer.has_thinking:
|
||||
generator = parse_thinking_models(
|
||||
generator,
|
||||
tokenizer.think_start,
|
||||
tokenizer.think_end,
|
||||
starts_in_thinking=detect_thinking_prompt_suffix(prompt, tokenizer),
|
||||
)
|
||||
generator = parse_deepseek_v32(generator)
|
||||
mlx_generator = parse_deepseek_v32(mlx_generator)
|
||||
else:
|
||||
if tokenizer.has_thinking:
|
||||
generator = parse_thinking_models(
|
||||
generator,
|
||||
mlx_generator = parse_thinking_models(
|
||||
mlx_generator,
|
||||
tokenizer.think_start,
|
||||
tokenizer.think_end,
|
||||
starts_in_thinking=detect_thinking_prompt_suffix(prompt, tokenizer),
|
||||
)
|
||||
|
||||
if tool_parser:
|
||||
generator = parse_tool_calls(generator, tool_parser, tools)
|
||||
mlx_generator = parse_tool_calls(mlx_generator, tool_parser, tools)
|
||||
|
||||
generator = count_reasoning_tokens(generator)
|
||||
|
||||
return map(lambda r: map_responses_to_chunks(r, model_id), generator)
|
||||
|
||||
|
||||
def map_responses_to_chunks(
|
||||
response: GenerationResponse | ToolCallResponse | None, model_id: ModelId
|
||||
) -> GenerationChunk | None:
|
||||
match response:
|
||||
case None:
|
||||
return None
|
||||
case GenerationResponse():
|
||||
if response.finish_reason == "error":
|
||||
return ErrorChunk(
|
||||
error_message=response.text,
|
||||
model=model_id,
|
||||
)
|
||||
else:
|
||||
finish_reason = response.finish_reason
|
||||
assert finish_reason not in (
|
||||
"error",
|
||||
"tool_calls",
|
||||
"function_call",
|
||||
)
|
||||
return TokenChunk(
|
||||
model=model_id,
|
||||
text=response.text,
|
||||
token_id=response.token,
|
||||
usage=response.usage,
|
||||
finish_reason=finish_reason,
|
||||
stats=response.stats,
|
||||
logprob=response.logprob,
|
||||
top_logprobs=response.top_logprobs,
|
||||
is_thinking=response.is_thinking,
|
||||
)
|
||||
case ToolCallResponse():
|
||||
return ToolCallChunk(
|
||||
tool_calls=response.tool_calls,
|
||||
model=model_id,
|
||||
usage=response.usage,
|
||||
stats=response.stats,
|
||||
)
|
||||
return count_reasoning_tokens(mlx_generator)
|
||||
|
||||
|
||||
def parse_gpt_oss(
|
||||
@@ -217,10 +163,11 @@ def parse_deepseek_v32(
|
||||
|
||||
Uses accumulated-text matching (not per-token marker checks) because
|
||||
DSML markers like <|DSML|function_calls> may span multiple tokens.
|
||||
Thinking tag handling is delegated to parse_thinking_models, which
|
||||
wraps this parser in apply_all_parsers.
|
||||
Also handles <think>...</think> blocks for thinking mode.
|
||||
"""
|
||||
from exo.worker.engines.mlx.dsml_encoding import (
|
||||
THINKING_END,
|
||||
THINKING_START,
|
||||
TOOL_CALLS_END,
|
||||
TOOL_CALLS_START,
|
||||
parse_dsml_output,
|
||||
@@ -228,6 +175,7 @@ def parse_deepseek_v32(
|
||||
|
||||
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
|
||||
@@ -269,6 +217,29 @@ def parse_deepseek_v32(
|
||||
yield response
|
||||
break
|
||||
|
||||
# ── 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
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import os
|
||||
import time
|
||||
from collections.abc import Generator
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
@@ -9,7 +8,11 @@ from anyio import WouldBlock
|
||||
from mlx_lm.tokenizer_utils import TokenizerWrapper
|
||||
|
||||
from exo.shared.models.model_cards import ModelTask
|
||||
from exo.shared.types.chunks import GenerationChunk
|
||||
from exo.shared.types.chunks import (
|
||||
ErrorChunk,
|
||||
TokenChunk,
|
||||
ToolCallChunk,
|
||||
)
|
||||
from exo.shared.types.common import CommandId, ModelId
|
||||
from exo.shared.types.events import (
|
||||
ChunkGenerated,
|
||||
@@ -31,7 +34,8 @@ from exo.shared.types.tasks import (
|
||||
)
|
||||
from exo.shared.types.worker.instances import BoundInstance
|
||||
from exo.shared.types.worker.runner_response import (
|
||||
ModelLoadingResponse,
|
||||
GenerationResponse,
|
||||
ToolCallResponse,
|
||||
)
|
||||
from exo.shared.types.worker.runners import (
|
||||
RunnerConnected,
|
||||
@@ -177,28 +181,23 @@ class Runner:
|
||||
)
|
||||
self.acknowledge_task(task)
|
||||
|
||||
def on_layer_loaded(layers_loaded: int, total: int) -> None:
|
||||
self.update_status(
|
||||
RunnerLoading(layers_loaded=layers_loaded, total_layers=total)
|
||||
)
|
||||
|
||||
assert (
|
||||
ModelTask.TextGeneration in self.shard_metadata.model_card.tasks
|
||||
), f"Incorrect model task(s): {self.shard_metadata.model_card.tasks}"
|
||||
|
||||
def load_model() -> Generator[ModelLoadingResponse]:
|
||||
assert isinstance(self.generator, Builder)
|
||||
(
|
||||
self.generator.inference_model,
|
||||
self.generator.tokenizer,
|
||||
self.generator.vision_processor,
|
||||
) = yield from load_mlx_items(
|
||||
self.bound_instance,
|
||||
self.generator.group,
|
||||
)
|
||||
|
||||
for load_resp in load_model():
|
||||
self.update_status(
|
||||
RunnerLoading(
|
||||
layers_loaded=load_resp.layers_loaded,
|
||||
total_layers=load_resp.total,
|
||||
)
|
||||
)
|
||||
(
|
||||
self.generator.inference_model,
|
||||
self.generator.tokenizer,
|
||||
self.generator.vision_processor,
|
||||
) = load_mlx_items(
|
||||
self.bound_instance,
|
||||
self.generator.group,
|
||||
on_layer_loaded=on_layer_loaded,
|
||||
)
|
||||
|
||||
self.generator = self.generator.build()
|
||||
|
||||
@@ -279,7 +278,9 @@ class Runner:
|
||||
self.send_task_status(task_id, TaskStatus.Complete)
|
||||
finished.append(task_id)
|
||||
case _:
|
||||
self.send_chunk(result, self.active_tasks[task_id].command_id)
|
||||
self.send_response(
|
||||
result, self.active_tasks[task_id].command_id
|
||||
)
|
||||
|
||||
for task_id in finished:
|
||||
self.active_tasks.pop(task_id, None)
|
||||
@@ -312,13 +313,59 @@ class Runner:
|
||||
|
||||
return ExitCode.AllTasksComplete
|
||||
|
||||
def send_chunk(
|
||||
def send_response(
|
||||
self,
|
||||
chunk: GenerationChunk,
|
||||
response: GenerationResponse | ToolCallResponse,
|
||||
command_id: CommandId,
|
||||
):
|
||||
if self.device_rank == 0:
|
||||
self.event_sender.send(ChunkGenerated(command_id=command_id, chunk=chunk))
|
||||
match response:
|
||||
case GenerationResponse():
|
||||
if self.device_rank == 0 and response.finish_reason == "error":
|
||||
self.event_sender.send(
|
||||
ChunkGenerated(
|
||||
command_id=command_id,
|
||||
chunk=ErrorChunk(
|
||||
error_message=response.text,
|
||||
model=self.model_id,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
elif self.device_rank == 0:
|
||||
assert response.finish_reason not in (
|
||||
"error",
|
||||
"tool_calls",
|
||||
"function_call",
|
||||
)
|
||||
self.event_sender.send(
|
||||
ChunkGenerated(
|
||||
command_id=command_id,
|
||||
chunk=TokenChunk(
|
||||
model=self.model_id,
|
||||
text=response.text,
|
||||
token_id=response.token,
|
||||
usage=response.usage,
|
||||
finish_reason=response.finish_reason,
|
||||
stats=response.stats,
|
||||
logprob=response.logprob,
|
||||
top_logprobs=response.top_logprobs,
|
||||
is_thinking=response.is_thinking,
|
||||
),
|
||||
)
|
||||
)
|
||||
case ToolCallResponse():
|
||||
if self.device_rank == 0:
|
||||
self.event_sender.send(
|
||||
ChunkGenerated(
|
||||
command_id=command_id,
|
||||
chunk=ToolCallChunk(
|
||||
tool_calls=response.tool_calls,
|
||||
model=self.model_id,
|
||||
usage=response.usage,
|
||||
stats=response.stats,
|
||||
),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -96,12 +96,7 @@ def run_gpt_oss_pipeline_device(
|
||||
n_layers=24,
|
||||
)
|
||||
|
||||
gen = shard_and_load(shard_meta, group)
|
||||
try:
|
||||
while True:
|
||||
next(gen)
|
||||
except StopIteration as stop:
|
||||
model, tokenizer = stop.value
|
||||
model, tokenizer = shard_and_load(shard_meta, group, on_layer_loaded=None)
|
||||
model = cast(Model, model)
|
||||
|
||||
# Generate a prompt of exact token length
|
||||
@@ -177,12 +172,7 @@ def run_gpt_oss_tensor_parallel_device(
|
||||
n_layers=24,
|
||||
)
|
||||
|
||||
gen = shard_and_load(shard_meta, group)
|
||||
try:
|
||||
while True:
|
||||
next(gen)
|
||||
except StopIteration as stop:
|
||||
model, tokenizer = stop.value
|
||||
model, tokenizer = shard_and_load(shard_meta, group, on_layer_loaded=None)
|
||||
model = cast(Model, model)
|
||||
|
||||
base_text = "The quick brown fox jumps over the lazy dog. "
|
||||
|
||||
@@ -343,7 +343,7 @@ class TestKVPrefixCacheWithModel:
|
||||
)
|
||||
|
||||
def test_mlx_generate_populates_cache(self, model_and_tokenizer):
|
||||
"""mlx_generate should save the post-prefill cache (before the decode loop)."""
|
||||
"""mlx_generate should save the cache after generation completes."""
|
||||
model, tokenizer = model_and_tokenizer
|
||||
|
||||
kv_prefix_cache = KVPrefixCache(None)
|
||||
@@ -356,6 +356,7 @@ class TestKVPrefixCacheWithModel:
|
||||
prompt_tokens = encode_prompt(tokenizer, prompt)
|
||||
|
||||
# Consume the entire generator so the cache-saving code after yield runs
|
||||
generated_tokens = 0
|
||||
for _response in mlx_generate(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
@@ -364,14 +365,13 @@ class TestKVPrefixCacheWithModel:
|
||||
kv_prefix_cache=kv_prefix_cache,
|
||||
group=None,
|
||||
):
|
||||
pass
|
||||
generated_tokens += 1
|
||||
|
||||
assert len(kv_prefix_cache.prompts) == 1
|
||||
assert len(kv_prefix_cache.caches) == 1
|
||||
# add_kv_cache is called before the decode loop and stores a deepcopy of
|
||||
# the cache as it is just after prefill + trim(2). Generation tokens are
|
||||
# never written into the stored entry.
|
||||
assert cache_length(kv_prefix_cache.caches[0]) == len(prompt_tokens) - 2
|
||||
# Cache should contain prompt + generated tokens
|
||||
expected_length = len(prompt_tokens) + generated_tokens
|
||||
assert cache_length(kv_prefix_cache.caches[0]) == expected_length
|
||||
|
||||
def test_mlx_generate_second_call_gets_prefix_hit(self, model_and_tokenizer):
|
||||
"""Second mlx_generate call with same prompt should get a prefix hit from stored cache."""
|
||||
|
||||
@@ -174,12 +174,7 @@ def _run_pipeline_device(
|
||||
n_layers=TOTAL_LAYERS,
|
||||
)
|
||||
|
||||
gen = shard_and_load(shard_meta, group)
|
||||
try:
|
||||
while True:
|
||||
next(gen)
|
||||
except StopIteration as stop:
|
||||
model, tokenizer = stop.value
|
||||
model, tokenizer = shard_and_load(shard_meta, group, on_layer_loaded=None)
|
||||
model = cast(Any, model)
|
||||
|
||||
prompt, task = _build_prompt(tokenizer, prompt_tokens)
|
||||
|
||||
@@ -14,8 +14,6 @@ import pytest
|
||||
from mlx.utils import tree_flatten, tree_unflatten
|
||||
from mlx_lm.tokenizer_utils import TokenizerWrapper
|
||||
|
||||
from exo.download.download_utils import resolve_existing_model
|
||||
from exo.shared.constants import EXO_MODELS_DIRS, EXO_MODELS_READ_ONLY_DIRS
|
||||
from exo.shared.types.common import ModelId
|
||||
from exo.shared.types.mlx import Model
|
||||
from exo.shared.types.text_generation import (
|
||||
@@ -30,6 +28,8 @@ from exo.worker.engines.mlx.utils_mlx import (
|
||||
load_tokenizer_for_model_id,
|
||||
)
|
||||
|
||||
HF_CACHE = Path.home() / ".cache" / "huggingface" / "hub"
|
||||
|
||||
# ── Config reduction ──────────────────────────────────────────────────────── #
|
||||
|
||||
_REDUCE = {
|
||||
@@ -100,21 +100,12 @@ def _reduce_config(cfg: dict[str, Any]) -> dict[str, Any]:
|
||||
|
||||
|
||||
def _find_snapshot(hub_name: str) -> Path | None:
|
||||
"""Locate a model directory under exo's models dirs.
|
||||
|
||||
Uses resolve_existing_model for fully-downloaded models; falls back to any
|
||||
existing directory (even partial) so that tokenizer-only copies still work.
|
||||
"""
|
||||
model_id = ModelId(f"mlx-community/{hub_name}")
|
||||
found = resolve_existing_model(model_id)
|
||||
if found is not None:
|
||||
return found
|
||||
normalized = model_id.normalize()
|
||||
for search_dir in (*EXO_MODELS_READ_ONLY_DIRS, *EXO_MODELS_DIRS):
|
||||
candidate = search_dir / normalized
|
||||
if candidate.is_dir():
|
||||
return candidate
|
||||
return None
|
||||
model_dir = HF_CACHE / f"models--mlx-community--{hub_name}"
|
||||
snaps = model_dir / "snapshots"
|
||||
if not snaps.exists():
|
||||
return None
|
||||
children = sorted(snaps.iterdir())
|
||||
return children[0] if children else None
|
||||
|
||||
|
||||
def _copy_tokenizer(src: Path, dst: Path) -> None:
|
||||
@@ -201,31 +192,13 @@ ARCHITECTURES: list[ArchSpec] = [
|
||||
]
|
||||
|
||||
|
||||
def _has_chat_template(model_dir: Path) -> bool:
|
||||
"""Check if a model dir has a usable chat template (inline or separate)."""
|
||||
if (model_dir / "chat_template.jinja").exists():
|
||||
return True
|
||||
cfg = model_dir / "tokenizer_config.json"
|
||||
if not cfg.exists():
|
||||
return False
|
||||
try:
|
||||
data = cast(dict[str, Any], json.loads(cfg.read_text()))
|
||||
except (OSError, json.JSONDecodeError):
|
||||
return False
|
||||
return bool(data.get("chat_template"))
|
||||
|
||||
|
||||
def _arch_available(spec: ArchSpec) -> bool:
|
||||
snap = _find_snapshot(spec.hub_name)
|
||||
if snap is None or not (snap / "config.json").exists():
|
||||
return False
|
||||
tokenizer_snap = snap
|
||||
if spec.tokenizer_hub is not None:
|
||||
alt = _find_snapshot(spec.tokenizer_hub)
|
||||
if alt is None:
|
||||
return False
|
||||
tokenizer_snap = alt
|
||||
return _has_chat_template(tokenizer_snap)
|
||||
return _find_snapshot(spec.tokenizer_hub) is not None
|
||||
return True
|
||||
|
||||
|
||||
def _make_task() -> TextGenerationTaskParams:
|
||||
|
||||
@@ -1,395 +0,0 @@
|
||||
# type: ignore
|
||||
"""uv run pytest -v -m "" src/exo/worker/tests/unittests/test_mlx/test_tp_bit_exact.py"""
|
||||
|
||||
import importlib
|
||||
import json
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import traceback
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
|
||||
MODEL_CONFIGS = {
|
||||
"llama": dict(
|
||||
module="mlx_lm.models.llama",
|
||||
args=dict(
|
||||
model_type="llama",
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_hidden_layers=2,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=4,
|
||||
rms_norm_eps=1e-6,
|
||||
vocab_size=512,
|
||||
max_position_embeddings=128,
|
||||
head_dim=32,
|
||||
rope_theta=10000.0,
|
||||
),
|
||||
),
|
||||
"qwen3_5_moe": dict(
|
||||
module="mlx_lm.models.qwen3_5_moe",
|
||||
args=dict(
|
||||
model_type="qwen3_5_moe",
|
||||
text_config=dict(
|
||||
model_type="qwen3_5_moe",
|
||||
vocab_size=512,
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_hidden_layers=4,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=4,
|
||||
head_dim=32,
|
||||
max_position_embeddings=128,
|
||||
rms_norm_eps=1e-6,
|
||||
tie_word_embeddings=False,
|
||||
attention_bias=False,
|
||||
full_attention_interval=2,
|
||||
linear_num_value_heads=32,
|
||||
linear_num_key_heads=16,
|
||||
linear_key_head_dim=32,
|
||||
linear_value_head_dim=32,
|
||||
linear_conv_kernel_dim=4,
|
||||
num_experts=16,
|
||||
num_experts_per_tok=2,
|
||||
decoder_sparse_step=1,
|
||||
shared_expert_intermediate_size=256,
|
||||
moe_intermediate_size=256,
|
||||
norm_topk_prob=True,
|
||||
rope_parameters={
|
||||
"type": "default",
|
||||
"rope_theta": 10000.0,
|
||||
"partial_rotary_factor": 0.25,
|
||||
"mrope_section": [11, 11, 10],
|
||||
},
|
||||
),
|
||||
),
|
||||
),
|
||||
"qwen3_next": dict(
|
||||
module="mlx_lm.models.qwen3_next",
|
||||
args=dict(
|
||||
model_type="qwen3_next",
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_hidden_layers=4,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=4,
|
||||
head_dim=32,
|
||||
max_position_embeddings=128,
|
||||
rms_norm_eps=1e-6,
|
||||
vocab_size=512,
|
||||
attention_bias=False,
|
||||
full_attention_interval=2,
|
||||
linear_num_value_heads=32,
|
||||
linear_num_key_heads=16,
|
||||
linear_key_head_dim=32,
|
||||
linear_value_head_dim=32,
|
||||
linear_conv_kernel_dim=4,
|
||||
num_experts=16,
|
||||
num_experts_per_tok=2,
|
||||
decoder_sparse_step=1,
|
||||
shared_expert_intermediate_size=256,
|
||||
moe_intermediate_size=256,
|
||||
norm_topk_prob=True,
|
||||
mlp_only_layers=[],
|
||||
rope_theta=10000.0,
|
||||
partial_rotary_factor=0.25,
|
||||
),
|
||||
),
|
||||
"deepseek_v3": dict(
|
||||
module="mlx_lm.models.deepseek_v3",
|
||||
args=dict(
|
||||
model_type="deepseek_v3",
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_hidden_layers=2,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=16,
|
||||
vocab_size=512,
|
||||
max_position_embeddings=128,
|
||||
rms_norm_eps=1e-6,
|
||||
n_routed_experts=8,
|
||||
n_shared_experts=1,
|
||||
num_experts_per_tok=2,
|
||||
moe_intermediate_size=256,
|
||||
moe_layer_freq=1,
|
||||
first_k_dense_replace=0,
|
||||
n_group=1,
|
||||
topk_group=1,
|
||||
routed_scaling_factor=1.0,
|
||||
q_lora_rank=None,
|
||||
kv_lora_rank=16,
|
||||
qk_nope_head_dim=16,
|
||||
qk_rope_head_dim=16,
|
||||
v_head_dim=32,
|
||||
rope_theta=10000.0,
|
||||
rope_scaling={},
|
||||
attention_bias=False,
|
||||
norm_topk_prob=True,
|
||||
scoring_func="sigmoid",
|
||||
topk_method="noaux_tc",
|
||||
),
|
||||
),
|
||||
"deepseek_v3_q4": dict(
|
||||
module="mlx_lm.models.deepseek_v3",
|
||||
quantize=dict(group_size=32, bits=4, mode="affine"),
|
||||
args=dict(
|
||||
model_type="deepseek_v3",
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_hidden_layers=2,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=16,
|
||||
vocab_size=512,
|
||||
max_position_embeddings=128,
|
||||
rms_norm_eps=1e-6,
|
||||
n_routed_experts=8,
|
||||
n_shared_experts=1,
|
||||
num_experts_per_tok=2,
|
||||
moe_intermediate_size=256,
|
||||
moe_layer_freq=1,
|
||||
first_k_dense_replace=0,
|
||||
n_group=1,
|
||||
topk_group=1,
|
||||
routed_scaling_factor=1.0,
|
||||
q_lora_rank=None,
|
||||
kv_lora_rank=64,
|
||||
qk_nope_head_dim=32,
|
||||
qk_rope_head_dim=32,
|
||||
v_head_dim=32,
|
||||
rope_theta=10000.0,
|
||||
rope_scaling={},
|
||||
attention_bias=False,
|
||||
norm_topk_prob=True,
|
||||
scoring_func="sigmoid",
|
||||
topk_method="noaux_tc",
|
||||
),
|
||||
),
|
||||
"glm4_moe_lite": dict(
|
||||
module="mlx_lm.models.glm4_moe_lite",
|
||||
args=dict(
|
||||
model_type="glm4_moe_lite",
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_hidden_layers=2,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=16,
|
||||
vocab_size=512,
|
||||
max_position_embeddings=128,
|
||||
rms_norm_eps=1e-6,
|
||||
n_routed_experts=8,
|
||||
n_shared_experts=1,
|
||||
num_experts_per_tok=2,
|
||||
moe_intermediate_size=256,
|
||||
first_k_dense_replace=1,
|
||||
n_group=1,
|
||||
topk_group=1,
|
||||
routed_scaling_factor=1.0,
|
||||
rope_theta=10000.0,
|
||||
attention_bias=False,
|
||||
q_lora_rank=None,
|
||||
kv_lora_rank=16,
|
||||
qk_rope_head_dim=16,
|
||||
qk_nope_head_dim=16,
|
||||
v_head_dim=32,
|
||||
),
|
||||
),
|
||||
"minimax": dict(
|
||||
module="mlx_lm.models.minimax",
|
||||
args=dict(
|
||||
model_type="minimax",
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=4,
|
||||
max_position_embeddings=128,
|
||||
num_experts_per_tok=2,
|
||||
num_local_experts=8,
|
||||
shared_intermediate_size=256,
|
||||
num_hidden_layers=2,
|
||||
rms_norm_eps=1e-6,
|
||||
rope_theta=10000.0,
|
||||
rotary_dim=32,
|
||||
vocab_size=512,
|
||||
),
|
||||
),
|
||||
"gpt_oss": dict(
|
||||
module="mlx_lm.models.gpt_oss",
|
||||
args=dict(
|
||||
model_type="gpt_oss",
|
||||
hidden_size=512,
|
||||
intermediate_size=256,
|
||||
num_hidden_layers=2,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=4,
|
||||
vocab_size=512,
|
||||
head_dim=32,
|
||||
rms_norm_eps=1e-6,
|
||||
num_local_experts=8,
|
||||
num_experts_per_tok=2,
|
||||
layer_types=["sliding_attention", "full_attention"],
|
||||
sliding_window=64,
|
||||
rope_theta=10000.0,
|
||||
),
|
||||
),
|
||||
"gemma4": dict(
|
||||
module="mlx_lm.models.gemma4",
|
||||
args=dict(
|
||||
model_type="gemma4",
|
||||
vocab_size=512,
|
||||
text_config=dict(
|
||||
vocab_size=512,
|
||||
hidden_size=512,
|
||||
intermediate_size=1024,
|
||||
num_hidden_layers=4,
|
||||
num_attention_heads=16,
|
||||
num_key_value_heads=4,
|
||||
head_dim=32,
|
||||
global_head_dim=32,
|
||||
num_kv_shared_layers=0,
|
||||
vocab_size_per_layer_input=512,
|
||||
hidden_size_per_layer_input=512,
|
||||
rms_norm_eps=1e-6,
|
||||
max_position_embeddings=128,
|
||||
sliding_window=64,
|
||||
sliding_window_pattern=2,
|
||||
layer_types=[
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
"sliding_attention",
|
||||
"full_attention",
|
||||
],
|
||||
enable_moe_block=True,
|
||||
num_experts=8,
|
||||
top_k_experts=2,
|
||||
moe_intermediate_size=256,
|
||||
),
|
||||
),
|
||||
),
|
||||
}
|
||||
|
||||
_PROMPT = [[1, 23, 45, 67, 89, 12, 34, 56]]
|
||||
|
||||
|
||||
def _build(name):
|
||||
import mlx.core as mx
|
||||
import mlx.nn as nn
|
||||
from mlx.utils import tree_map_with_path
|
||||
|
||||
import exo.worker.engines.mlx.auto_parallel # noqa: F401
|
||||
|
||||
cfg = MODEL_CONFIGS[name]
|
||||
module = importlib.import_module(cfg["module"])
|
||||
model_cls = module.Model
|
||||
model_args_cls = module.ModelArgs
|
||||
|
||||
mx.random.seed(0)
|
||||
args = model_args_cls(**cfg["args"])
|
||||
m = model_cls(args)
|
||||
|
||||
def _to_bf16(_p, v):
|
||||
if hasattr(v, "dtype") and v.dtype in (mx.float16, mx.float32, mx.bfloat16):
|
||||
return v.astype(mx.bfloat16)
|
||||
return v
|
||||
|
||||
m.update(tree_map_with_path(_to_bf16, m.parameters()))
|
||||
if "quantize" in cfg:
|
||||
nn.quantize(m, **cfg["quantize"])
|
||||
mx.eval(m.parameters())
|
||||
return mx, m
|
||||
|
||||
|
||||
def _run(name, out_path, shard):
|
||||
import mlx.core as mx
|
||||
|
||||
if shard:
|
||||
g = mx.distributed.init(backend="ring", strict=True)
|
||||
mx_, m = _build(name)
|
||||
if shard:
|
||||
from exo.worker.engines.mlx.auto_parallel import tensor_auto_parallel
|
||||
|
||||
m = tensor_auto_parallel(m, g, on_layer_loaded=None)
|
||||
mx_.eval(m.parameters())
|
||||
inputs = mx_.array(_PROMPT, dtype=mx_.int32)
|
||||
logits = m(inputs)
|
||||
mx_.eval(logits)
|
||||
np.savez(out_path, logits=np.asarray(logits.astype(mx_.float32)))
|
||||
|
||||
|
||||
def _ref_worker(name, out_path, q):
|
||||
try:
|
||||
_run(name, out_path, shard=False)
|
||||
q.put(True)
|
||||
except BaseException as e:
|
||||
q.put(f"{e}\n{traceback.format_exc()}")
|
||||
|
||||
|
||||
def _tp_worker(name, rank, hf, out_path, q):
|
||||
os.environ["MLX_HOSTFILE"] = hf
|
||||
os.environ["MLX_RANK"] = str(rank)
|
||||
try:
|
||||
path = out_path if rank == 0 else out_path + f".r{rank}"
|
||||
_run(name, path, shard=True)
|
||||
q.put((rank, True, None))
|
||||
except BaseException as e:
|
||||
q.put((rank, False, f"{e}\n{traceback.format_exc()}"))
|
||||
|
||||
|
||||
def _run_compare(name, world_size, port_base):
|
||||
d = tempfile.mkdtemp()
|
||||
ref_path = f"{d}/ref.npz"
|
||||
tp_path = f"{d}/tp.npz"
|
||||
ctx = mp.get_context("spawn")
|
||||
q = ctx.Queue()
|
||||
|
||||
p = ctx.Process(target=_ref_worker, args=(name, ref_path, q))
|
||||
p.start()
|
||||
p.join(300)
|
||||
r = q.get(timeout=10)
|
||||
if r is not True:
|
||||
pytest.fail(f"[{name}] ref FAIL: {str(r)[:500]}")
|
||||
|
||||
hosts = [f"127.0.0.1:{port_base + i}" for i in range(world_size)]
|
||||
with tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) as f:
|
||||
json.dump(hosts, f)
|
||||
hf = f.name
|
||||
ps = [
|
||||
ctx.Process(target=_tp_worker, args=(name, rank, hf, tp_path, q))
|
||||
for rank in range(world_size)
|
||||
]
|
||||
for pp in ps:
|
||||
pp.start()
|
||||
results = [q.get(timeout=300) for _ in range(world_size)]
|
||||
for pp in ps:
|
||||
pp.join(60)
|
||||
for rank, ok, payload in results:
|
||||
if not ok:
|
||||
pytest.fail(f"[{name}] rank {rank} FAIL: {payload[:500]}")
|
||||
|
||||
ref = np.load(ref_path)["logits"]
|
||||
tp = np.load(tp_path)["logits"]
|
||||
diff = np.abs(ref - tp)
|
||||
max_diff = float(diff.max())
|
||||
mean_diff = float(diff.mean())
|
||||
assert max_diff == 0.0, (
|
||||
f"[{name} TP={world_size}] not bit-exact: max={max_diff} mean={mean_diff}"
|
||||
)
|
||||
|
||||
|
||||
pytestmark = [
|
||||
pytest.mark.slow,
|
||||
pytest.mark.skipif(
|
||||
sys.platform != "darwin", reason="MLX distributed requires Metal"
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.skip("TP=2 is currently very different to TP=1. This test will not pass")
|
||||
@pytest.mark.parametrize("world_size", [2, 4])
|
||||
@pytest.mark.parametrize("name", list(MODEL_CONFIGS))
|
||||
def test_tp_bit_exact(name, world_size):
|
||||
name_idx = list(MODEL_CONFIGS).index(name)
|
||||
port = 32000 + name_idx * 20 + world_size
|
||||
_run_compare(name, world_size, port)
|
||||
@@ -54,7 +54,6 @@ def test_plan_requests_download_when_waiting_and_shard_not_downloaded():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -110,7 +109,6 @@ def test_plan_loads_model_when_all_shards_downloaded_and_waiting():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -156,7 +154,6 @@ def test_plan_does_not_request_download_when_shard_already_downloaded():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -207,7 +204,6 @@ def test_plan_does_not_load_model_until_all_shards_downloaded_globally():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -231,7 +227,6 @@ def test_plan_does_not_load_model_until_all_shards_downloaded_globally():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
|
||||
@@ -54,7 +54,6 @@ def test_plan_kills_runner_when_instance_missing():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -97,7 +96,6 @@ def test_plan_kills_runner_when_sibling_failed():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -132,7 +130,6 @@ def test_plan_creates_runner_when_missing_for_node():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -174,7 +171,6 @@ def test_plan_does_not_create_runner_when_supervisor_already_present():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -207,7 +203,6 @@ def test_plan_does_not_create_runner_for_unassigned_node():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
|
||||
@@ -78,7 +78,6 @@ def test_plan_forwards_pending_chat_completion_when_runner_ready():
|
||||
all_runners=all_runners,
|
||||
tasks={TASK_1_ID: task},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -132,7 +131,6 @@ def test_plan_does_not_forward_chat_completion_if_any_runner_not_ready():
|
||||
all_runners=all_runners,
|
||||
tasks={TASK_1_ID: task},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -183,7 +181,6 @@ def test_plan_does_not_forward_tasks_for_other_instances():
|
||||
all_runners=all_runners,
|
||||
tasks={foreign_task.task_id: foreign_task},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -252,7 +249,6 @@ def test_plan_ignores_non_pending_or_non_chat_tasks():
|
||||
all_runners=all_runners,
|
||||
tasks={TASK_1_ID: completed_task, other_task_id: other_task},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -295,7 +291,6 @@ def test_plan_returns_none_when_nothing_to_do():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
|
||||
@@ -63,7 +63,6 @@ def test_plan_starts_warmup_for_accepting_rank_when_all_loaded_or_warming():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -108,7 +107,6 @@ def test_plan_starts_warmup_for_rank_zero_after_others_warming():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -152,7 +150,6 @@ def test_plan_does_not_start_warmup_for_non_zero_rank_until_all_loaded_or_warmin
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -200,7 +197,6 @@ def test_plan_does_not_start_warmup_for_rank_zero_until_others_warming():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -220,7 +216,6 @@ def test_plan_does_not_start_warmup_for_rank_zero_until_others_warming():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -267,7 +262,6 @@ def test_plan_starts_warmup_for_connecting_rank_after_others_warming():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -313,7 +307,6 @@ def test_plan_does_not_start_warmup_for_accepting_rank_until_all_loaded_or_warmi
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
@@ -358,7 +351,6 @@ def test_plan_does_not_start_warmup_for_connecting_rank_until_others_warming():
|
||||
all_runners=all_runners,
|
||||
tasks={},
|
||||
input_chunk_buffer={},
|
||||
image_cache={},
|
||||
instance_backoff=KeyedBackoff(),
|
||||
download_backoff=KeyedBackoff(),
|
||||
)
|
||||
|
||||
@@ -20,25 +20,7 @@ from exo.worker.engines.mlx.dsml_encoding import (
|
||||
encode_messages,
|
||||
parse_dsml_output,
|
||||
)
|
||||
from exo.worker.runner.llm_inference.model_output_parsers import (
|
||||
parse_deepseek_v32,
|
||||
parse_thinking_models,
|
||||
)
|
||||
|
||||
|
||||
def _parse_deepseek_with_thinking(
|
||||
source: Generator[GenerationResponse | None],
|
||||
starts_in_thinking: bool = False,
|
||||
) -> Generator[GenerationResponse | ToolCallResponse | None]:
|
||||
return parse_deepseek_v32(
|
||||
parse_thinking_models(
|
||||
source,
|
||||
think_start=THINKING_START,
|
||||
think_end=THINKING_END,
|
||||
starts_in_thinking=starts_in_thinking,
|
||||
)
|
||||
)
|
||||
|
||||
from exo.worker.runner.llm_inference.model_output_parsers import parse_deepseek_v32
|
||||
|
||||
# ── Shared fixtures ──────────────────────────────────────────────
|
||||
|
||||
@@ -351,7 +333,9 @@ class TestE2EThinkingAndToolCall:
|
||||
assert prompt.endswith(THINKING_START)
|
||||
|
||||
# Simulate: model outputs <think>, thinks, closes thinking, then tool call.
|
||||
# Use the full production chain (parse_thinking_models → parse_deepseek_v32).
|
||||
# In the full pipeline, parse_thinking_models handles the case where
|
||||
# <think> is in the prompt. Here we test parse_deepseek_v32 directly,
|
||||
# which detects <think>/<think> markers in the stream.
|
||||
model_tokens = [
|
||||
THINKING_START,
|
||||
"The user wants weather",
|
||||
@@ -369,7 +353,7 @@ class TestE2EThinkingAndToolCall:
|
||||
TOOL_CALLS_END,
|
||||
]
|
||||
|
||||
results = list(_parse_deepseek_with_thinking(_simulate_tokens(model_tokens)))
|
||||
results = list(parse_deepseek_v32(_simulate_tokens(model_tokens)))
|
||||
|
||||
gen_results = [r for r in results if isinstance(r, GenerationResponse)]
|
||||
tool_results = [r for r in results if isinstance(r, ToolCallResponse)]
|
||||
@@ -403,7 +387,7 @@ class TestE2EThinkingAndToolCall:
|
||||
prompt_no_think = encode_messages(
|
||||
messages, tools=_WEATHER_TOOLS, thinking_mode="chat"
|
||||
)
|
||||
assert not prompt_no_think.endswith(THINKING_START)
|
||||
assert prompt_no_think.endswith(THINKING_END)
|
||||
|
||||
# Both should have the same tool definitions
|
||||
assert "get_weather" in prompt_think
|
||||
@@ -613,9 +597,7 @@ class TestE2EFullRoundTrip:
|
||||
f"</{DSML_TOKEN}invoke>\n",
|
||||
TOOL_CALLS_END,
|
||||
]
|
||||
results_1 = list(
|
||||
_parse_deepseek_with_thinking(_simulate_tokens(model_tokens_1))
|
||||
)
|
||||
results_1 = list(parse_deepseek_v32(_simulate_tokens(model_tokens_1)))
|
||||
|
||||
# Verify: thinking tokens + tool call
|
||||
gen_1 = [r for r in results_1 if isinstance(r, GenerationResponse)]
|
||||
@@ -678,9 +660,7 @@ class TestE2EFullRoundTrip:
|
||||
THINKING_END,
|
||||
"The weather in Hangzhou is currently cloudy with temperatures between 7°C and 13°C.",
|
||||
]
|
||||
results_2 = list(
|
||||
_parse_deepseek_with_thinking(_simulate_tokens(model_tokens_2))
|
||||
)
|
||||
results_2 = list(parse_deepseek_v32(_simulate_tokens(model_tokens_2)))
|
||||
|
||||
gen_2 = [r for r in results_2 if isinstance(r, GenerationResponse)]
|
||||
tool_2 = [r for r in results_2 if isinstance(r, ToolCallResponse)]
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
# Check tasks are complete before runner is ever ready.
|
||||
import unittest.mock
|
||||
from collections.abc import Iterable
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable
|
||||
|
||||
import mlx.core as mx
|
||||
@@ -116,22 +115,13 @@ def assert_events_equal(test_events: Iterable[Event], true_events: Iterable[Even
|
||||
assert test_event == true_event, f"{test_event} != {true_event}"
|
||||
|
||||
|
||||
@dataclass
|
||||
class MockLoadOutput:
|
||||
layers_loaded: int
|
||||
total: int
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patch_out_mlx(monkeypatch: pytest.MonkeyPatch):
|
||||
# initialize_mlx returns a mock group
|
||||
monkeypatch.setattr(mlx_runner, "initialize_mlx", make_nothin(MockGroup()))
|
||||
|
||||
def lmi_gen():
|
||||
yield MockLoadOutput(1, 1)
|
||||
return (1, MockTokenizer, None)
|
||||
|
||||
monkeypatch.setattr(mlx_runner, "load_mlx_items", make_nothin(lmi_gen()))
|
||||
monkeypatch.setattr(
|
||||
mlx_runner, "load_mlx_items", make_nothin((1, MockTokenizer, None))
|
||||
)
|
||||
monkeypatch.setattr(mlx_batch_generator, "warmup_inference", make_nothin(1))
|
||||
monkeypatch.setattr(mlx_batch_generator, "_check_for_debug_prompts", nothin)
|
||||
monkeypatch.setattr(mlx_batch_generator, "mx_any", make_nothin(False))
|
||||
@@ -328,10 +318,6 @@ def test_events_processed_in_correct_order(patch_out_mlx: pytest.MonkeyPatch):
|
||||
runner_status=RunnerLoading(layers_loaded=0, total_layers=32),
|
||||
),
|
||||
TaskAcknowledged(task_id=LOAD_TASK_ID),
|
||||
RunnerStatusUpdated(
|
||||
runner_id=RUNNER_1_ID,
|
||||
runner_status=RunnerLoading(layers_loaded=1, total_layers=1),
|
||||
),
|
||||
TaskStatusUpdated(task_id=LOAD_TASK_ID, task_status=TaskStatus.Complete),
|
||||
RunnerStatusUpdated(runner_id=RUNNER_1_ID, runner_status=RunnerLoaded()),
|
||||
TaskStatusUpdated(task_id=WARMUP_TASK_ID, task_status=TaskStatus.Running),
|
||||
|
||||
@@ -380,110 +380,6 @@ class TestGenericToolCallsFinishReason:
|
||||
# ── Double parser chain (parse_thinking_models → parse_deepseek_v32) ──
|
||||
|
||||
|
||||
class TestDeepSeekV32StartsInThinking:
|
||||
"""Regression tests for deepseek v3.2 where the chat template appends
|
||||
<think> to the prompt so the model starts already inside a thinking block.
|
||||
"""
|
||||
|
||||
def test_reasoning_tagged_when_starts_in_thinking(self):
|
||||
tokens = [
|
||||
_make_response("let me", 0),
|
||||
_make_response(" think", 1),
|
||||
_make_response(THINKING_END, 2),
|
||||
_make_response("\n", 3),
|
||||
_make_response("42", 4, finish_reason="stop"),
|
||||
]
|
||||
thinking = parse_thinking_models(
|
||||
_queue_source(tokens),
|
||||
think_start=THINKING_START,
|
||||
think_end=THINKING_END,
|
||||
starts_in_thinking=True,
|
||||
)
|
||||
results = _step_until_finish(parse_deepseek_v32(thinking))
|
||||
gens = [
|
||||
r
|
||||
for r in results
|
||||
if isinstance(r, GenerationResponse) and r.finish_reason is None
|
||||
]
|
||||
texts = [(r.text, r.is_thinking) for r in gens]
|
||||
assert texts == [("let me", True), (" think", True), ("\n", False)]
|
||||
final = [
|
||||
r
|
||||
for r in results
|
||||
if isinstance(r, GenerationResponse) and r.finish_reason is not None
|
||||
]
|
||||
assert len(final) == 1
|
||||
assert final[0].text == "42"
|
||||
assert final[0].is_thinking is False
|
||||
|
||||
def test_starts_in_thinking_then_tool_call(self):
|
||||
tokens = [
|
||||
_make_response("need weather", 0),
|
||||
_make_response(THINKING_END, 1),
|
||||
_make_response("\n\n", 2),
|
||||
_make_response(TOOL_CALLS_START, 3),
|
||||
_make_response("\n", 4),
|
||||
_make_response(f'<{DSML_TOKEN}invoke name="get_weather">\n', 5),
|
||||
_make_response(
|
||||
f'<{DSML_TOKEN}parameter name="city" string="true">NYC</{DSML_TOKEN}parameter>\n',
|
||||
6,
|
||||
),
|
||||
_make_response(f"</{DSML_TOKEN}invoke>\n", 7),
|
||||
_make_response(TOOL_CALLS_END, 8, finish_reason="stop"),
|
||||
]
|
||||
thinking = parse_thinking_models(
|
||||
_queue_source(tokens),
|
||||
think_start=THINKING_START,
|
||||
think_end=THINKING_END,
|
||||
starts_in_thinking=True,
|
||||
)
|
||||
results = _step_until_finish(parse_deepseek_v32(thinking))
|
||||
reasoning_gens = [
|
||||
r
|
||||
for r in results
|
||||
if isinstance(r, GenerationResponse)
|
||||
and r.finish_reason is None
|
||||
and r.is_thinking
|
||||
]
|
||||
assert [r.text for r in reasoning_gens] == ["need weather"]
|
||||
tool_results = [r for r in results if isinstance(r, ToolCallResponse)]
|
||||
assert len(tool_results) == 1
|
||||
assert tool_results[0].tool_calls[0].name == "get_weather"
|
||||
|
||||
def test_reasoning_tokens_counted_starts_in_thinking(self):
|
||||
usage = Usage(
|
||||
prompt_tokens=10,
|
||||
completion_tokens=5,
|
||||
total_tokens=15,
|
||||
prompt_tokens_details=PromptTokensDetails(cached_tokens=0),
|
||||
completion_tokens_details=CompletionTokensDetails(reasoning_tokens=0),
|
||||
)
|
||||
tokens = [
|
||||
_make_response("reasoning", 0),
|
||||
_make_response(" more", 1),
|
||||
_make_response(THINKING_END, 2),
|
||||
_make_response("\n", 3),
|
||||
GenerationResponse(text="42", token=4, finish_reason="stop", usage=usage),
|
||||
]
|
||||
thinking = parse_thinking_models(
|
||||
_queue_source(tokens),
|
||||
think_start=THINKING_START,
|
||||
think_end=THINKING_END,
|
||||
starts_in_thinking=True,
|
||||
)
|
||||
results = _step_until_finish(
|
||||
count_reasoning_tokens(parse_deepseek_v32(thinking))
|
||||
)
|
||||
final = [
|
||||
r
|
||||
for r in results
|
||||
if isinstance(r, GenerationResponse) and r.finish_reason is not None
|
||||
]
|
||||
assert len(final) == 1
|
||||
assert final[0].usage is not None
|
||||
assert final[0].usage.completion_tokens_details.reasoning_tokens == 2
|
||||
|
||||
|
||||
class TestBatchGeneratorSingleNext:
|
||||
def test_finish_reason_with_buffered_tokens_drain_loop(self):
|
||||
from exo.worker.runner.llm_inference.batch_generator import GeneratorQueue
|
||||
|
||||
Reference in new issue
Block a user