From c55cbf6739664bf1d3d5b0cd652f9b36638997dd Mon Sep 17 00:00:00 2001 From: rltakashige Date: Tue, 27 Jan 2026 15:29:06 +0000 Subject: [PATCH 1/3] Add mlx lm style tensor sharding for Minimax (#1299) ## Motivation Broken right now. We'll potentially add a better one later ## Changes ## Why It Works ## Test Plan ### Manual Testing Used for evals without any issue. ### Automated Testing --- src/exo/worker/engines/mlx/auto_parallel.py | 11 +++++++++++ 1 file changed, 11 insertions(+) diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index f2ecb698..89eb4185 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -622,6 +622,7 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy): on_timeout: TimeoutCallback | None, ) -> nn.Module: model = cast(MiniMaxModel, model) + rank = self.group.rank() for layer in model.layers: eval_with_timeout( layer.parameters(), timeout_seconds / len(model.layers), on_timeout @@ -631,6 +632,16 @@ class MiniMaxShardingStrategy(TensorParallelShardingStrategy): layer.self_attn.k_proj = self.all_to_sharded_linear(layer.self_attn.k_proj) layer.self_attn.v_proj = self.all_to_sharded_linear(layer.self_attn.v_proj) layer.self_attn.o_proj = self.sharded_to_all_linear(layer.self_attn.o_proj) + + # Shard qk_norm weights if present (must match sharded head count) + if getattr(layer.self_attn, "use_qk_norm", False): + layer.self_attn.q_norm.weight = layer.self_attn.q_norm.weight.split( # type: ignore + self.N, axis=-1 + )[rank] + layer.self_attn.k_norm.weight = layer.self_attn.k_norm.weight.split( # type: ignore + self.N, axis=-1 + )[rank] + layer.self_attn.num_attention_heads //= self.N layer.self_attn.num_key_value_heads //= self.N From 991d2781197a4c2a06bbde5aa4bde6520fca8bfa Mon Sep 17 00:00:00 2001 From: Evan Quiney Date: Tue, 27 Jan 2026 17:03:01 +0000 Subject: [PATCH 2/3] replace nix fmt with treefmt in just lint (#1301) man evaluating the nix flake is so slow. treefmt speeeedy --- justfile | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/justfile b/justfile index d8da7690..dabb4fea 100644 --- a/justfile +++ b/justfile @@ -1,7 +1,7 @@ export NIX_CONFIG := "extra-experimental-features = nix-command flakes" fmt: - nix fmt + treefmt || nix fmt lint: uv run ruff check --fix From a562114ba531929eff1a6a3f4db881b613c6f542 Mon Sep 17 00:00:00 2001 From: rltakashige Date: Wed, 28 Jan 2026 05:44:19 +0000 Subject: [PATCH 3/3] Add Kimi K2.5 support (#1302) ## Motivation ## Changes ## Why It Works ## Test Plan ### Manual Testing ### Automated Testing --------- Co-authored-by: Alex Cheema <41707476+AlexCheema@users.noreply.github.com> --- pyproject.toml | 3 +- src/exo/shared/models/model_cards.py | 8 +++++ src/exo/worker/engines/mlx/auto_parallel.py | 5 +-- src/exo/worker/engines/mlx/utils_mlx.py | 35 ++++++++++++++++++--- uv.lock | 15 +++------ 5 files changed, 49 insertions(+), 17 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 319b3155..96958224 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,7 +19,7 @@ dependencies = [ "anyio==4.11.0", "mlx==0.30.3; sys_platform == 'darwin'", "mlx[cpu]==0.30.3; sys_platform == 'linux'", - "mlx-lm==0.30.5", + "mlx-lm", "tiktoken>=0.12.0", # required for kimi k2 tokenizer "hypercorn>=0.18.0", "openai-harmony>=0.0.8", @@ -63,6 +63,7 @@ members = [ [tool.uv.sources] exo_pyo3_bindings = { workspace = true } +mlx-lm = { git = "https://github.com/ml-explore/mlx-lm", branch = "main" } # Uncomment to use local mlx/mlx-lm development versions: # mlx = { path = "/Users/Shared/mlx", editable=true } # mlx-lm = { path = "/Users/Shared/mlx-lm", editable=true } diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index bcc5c73f..63d7c2a4 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -121,6 +121,14 @@ MODEL_CARDS: dict[str, ModelCard] = { supports_tensor=True, tasks=[ModelTask.TextGeneration], ), + "kimi-k2.5": ModelCard( + model_id=ModelId("mlx-community/Kimi-K2.5"), + storage_size=Memory.from_gb(617), + n_layers=61, + hidden_size=7168, + supports_tensor=True, + tasks=[ModelTask.TextGeneration], + ), # llama-3.1 "llama-3.1-8b": ModelCard( model_id=ModelId("mlx-community/Meta-Llama-3.1-8B-Instruct-4bit"), diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py index 89eb4185..ff2052fb 100644 --- a/src/exo/worker/engines/mlx/auto_parallel.py +++ b/src/exo/worker/engines/mlx/auto_parallel.py @@ -23,6 +23,7 @@ from mlx_lm.models.glm4_moe_lite import Glm4MoeLiteDecoderLayer, Glm4MoeLiteMLP from mlx_lm.models.glm4_moe_lite import Model as GLM4MoeLiteModel from mlx_lm.models.gpt_oss import GptOssMoeModel from mlx_lm.models.gpt_oss import Model as GptOssModel +from mlx_lm.models.kimi_k25 import Model as KimiK25Model from mlx_lm.models.llama import Model as LlamaModel from mlx_lm.models.minimax import Model as MiniMaxModel from mlx_lm.models.ministral3 import Model as Ministral3Model @@ -344,7 +345,7 @@ def tensor_auto_parallel( all_to_sharded_linear_in_place, sharded_to_all_linear_in_place, ) - elif isinstance(model, (DeepseekV3Model, DeepseekV32Model)): + elif isinstance(model, (DeepseekV3Model, DeepseekV32Model, KimiK25Model)): tensor_parallel_sharding_strategy = DeepSeekShardingStrategy( group, all_to_sharded_linear, @@ -453,7 +454,7 @@ def _set_layers(model: nn.Module, layers: list[_LayerCallable]) -> None: # Update DeepSeek V3 specific parameters when layers are shrunk if isinstance( - model, (DeepseekV3Model, DeepseekV32Model, Glm4MoeModel) + model, (DeepseekV3Model, DeepseekV32Model, Glm4MoeModel, KimiK25Model) ) and hasattr(inner_model_instance, "num_layers"): logger.info( f"Setting num_layers to {len(layers)} for model {model.model.__class__.__name__}" diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py index 5a2fda9d..ccc54a78 100644 --- a/src/exo/worker/engines/mlx/utils_mlx.py +++ b/src/exo/worker/engines/mlx/utils_mlx.py @@ -259,10 +259,10 @@ def shard_and_load( logger.info(f"Group size: {group.size()}, group rank: {group.rank()}") - # Estimate timeout based on model size - base_timeout = float(os.environ.get("EXO_MODEL_LOAD_TIMEOUT", "60")) + # Estimate timeout based on model size (5x default for large queued workloads) + base_timeout = float(os.environ.get("EXO_MODEL_LOAD_TIMEOUT", "300")) model_size_gb = get_weights_size(shard_metadata).in_bytes / (1024**3) - timeout_seconds = base_timeout + model_size_gb / 5 + timeout_seconds = base_timeout + model_size_gb logger.info( f"Evaluating model parameters with timeout of {timeout_seconds:.0f}s " f"(model size: {model_size_gb:.1f}GB)" @@ -339,8 +339,35 @@ def load_tokenizer_for_model_id( # Kimi uses a custom TikTokenTokenizer that transformers 5.x can't load via AutoTokenizer if "kimi-k2" in model_id_lower: + import importlib.util + import types + sys.path.insert(0, str(model_path)) - from tokenization_kimi import TikTokenTokenizer # type: ignore[import-not-found] # noqa: I001 + + # Load tool_declaration_ts first (tokenization_kimi imports it with relative import) + tool_decl_path = model_path / "tool_declaration_ts.py" + if tool_decl_path.exists(): + spec = importlib.util.spec_from_file_location( + "tool_declaration_ts", tool_decl_path + ) + if spec and spec.loader: + tool_decl_module = importlib.util.module_from_spec(spec) + sys.modules["tool_declaration_ts"] = tool_decl_module + spec.loader.exec_module(tool_decl_module) + + # Load tokenization_kimi with patched source (convert relative to absolute import) + tok_path = model_path / "tokenization_kimi.py" + source = tok_path.read_text() + source = source.replace("from .tool_declaration_ts", "from tool_declaration_ts") + spec = importlib.util.spec_from_file_location("tokenization_kimi", tok_path) + if spec: + tok_module = types.ModuleType("tokenization_kimi") + tok_module.__file__ = str(tok_path) + sys.modules["tokenization_kimi"] = tok_module + exec(compile(source, tok_path, "exec"), tok_module.__dict__) # noqa: S102 + TikTokenTokenizer = tok_module.TikTokenTokenizer # type: ignore[attr-defined] # noqa: N806 + else: + from tokenization_kimi import TikTokenTokenizer # type: ignore[import-not-found] # noqa: I001 hf_tokenizer: Any = TikTokenTokenizer.from_pretrained(model_path) # pyright: ignore[reportUnknownVariableType,reportUnknownMemberType] diff --git a/uv.lock b/uv.lock index 79e8b5f1..75dbd5c1 100644 --- a/uv.lock +++ b/uv.lock @@ -415,7 +415,7 @@ requires-dist = [ { name = "mflux", specifier = "==0.15.4" }, { name = "mlx", marker = "sys_platform == 'darwin'", specifier = "==0.30.3" }, { name = "mlx", extras = ["cpu"], marker = "sys_platform == 'linux'", specifier = "==0.30.3" }, - { name = "mlx-lm", specifier = "==0.30.5" }, + { name = "mlx-lm", git = "https://github.com/ml-explore/mlx-lm?branch=main" }, { name = "openai-harmony", specifier = ">=0.0.8" }, { name = "pillow", specifier = ">=11.0,<12.0" }, { name = "psutil", specifier = ">=7.0.0" }, @@ -1073,7 +1073,7 @@ wheels = [ [[package]] name = "mlx-lm" version = "0.30.5" -source = { registry = "https://pypi.org/simple" } +source = { git = "https://github.com/ml-explore/mlx-lm?branch=main#96699e6dadb13b82b28285bb131a0741997d19ae" } dependencies = [ { name = "jinja2", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "mlx", marker = "sys_platform == 'darwin'" }, @@ -1083,10 +1083,6 @@ dependencies = [ { name = "sentencepiece", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "transformers", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/0b/90/4469d9f75f196e6255f59a89441abe0079925d30a001462e1c1c4bc4e6a1/mlx_lm-0.30.5.tar.gz", hash = "sha256:9e6cb258c65b766c6af25cb90958aef40acab67139f05839eef19864cb3154f6", size = 262367, upload-time = "2026-01-25T15:29:30.125Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/89/ba/66db6e1e5f1ef506655b562932f6bd8f72600116d5f31f92d71c1f200b3f/mlx_lm-0.30.5-py3-none-any.whl", hash = "sha256:a80bc8e3efdebe81813b0f6eb403fb66a7a15071e256f4e7102ada986acb75bb", size = 366716, upload-time = "2026-01-25T15:29:28.29Z" }, -] [[package]] name = "mlx-metal" @@ -2285,7 +2281,7 @@ wheels = [ [[package]] name = "transformers" -version = "5.0.0rc3" +version = "5.0.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "filelock", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, @@ -2294,15 +2290,14 @@ dependencies = [ { name = "packaging", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "pyyaml", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "regex", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, - { name = "requests", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "safetensors", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "tokenizers", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "tqdm", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, { name = "typer-slim", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/3f/a3/7c116a8d85f69ea7749cf4c2df79e64c35d028e5fc7ea0168f299d03b8c7/transformers-5.0.0rc3.tar.gz", hash = "sha256:a0315b92b7e087617ade42ec9e6e92ee7620541cc5d6a3331886c52cbe306f5c", size = 8388520, upload-time = "2026-01-14T16:49:02.952Z" } +sdist = { url = "https://files.pythonhosted.org/packages/bc/79/845941711811789c85fb7e2599cea425a14a07eda40f50896b9d3fda7492/transformers-5.0.0.tar.gz", hash = "sha256:5f5634efed6cf76ad068cc5834c7adbc32db78bbd6211fb70df2325a9c37dec8", size = 8424830, upload-time = "2026-01-26T10:46:46.813Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/1e/f2/ae2b8968764253bdf38a48dee3c299b8d0bedf7c8ffbe3449fca9bd95338/transformers-5.0.0rc3-py3-none-any.whl", hash = "sha256:383fad27f4f73092d330e45fae384681e5c8521e1dc1cf6cb1a297780e68bf2d", size = 10107087, upload-time = "2026-01-14T16:48:59.393Z" }, + { url = "https://files.pythonhosted.org/packages/52/f3/ac976fa8e305c9e49772527e09fbdc27cc6831b8a2f6b6063406626be5dd/transformers-5.0.0-py3-none-any.whl", hash = "sha256:587086f249ce64c817213cf36afdb318d087f790723e9b3d4500b97832afd52d", size = 10142091, upload-time = "2026-01-26T10:46:43.88Z" }, ] [[package]]