← back to Exo
add glm-47, minimax-m21 (#1147)
82ba42bae9fc111e25478a0a8c8f45fd95163b65 · 2026-01-14 16:33:17 +0000 · Evan Quiney
Adds support glm 4.7 and MiniMax M2.1
Manual testing:
Tensor + Pipeline execution of both models.
Closes #1141 and #1142
Files touched
M src/exo/shared/models/model_cards.pyM src/exo/worker/engines/mlx/auto_parallel.py
Diff
commit 82ba42bae9fc111e25478a0a8c8f45fd95163b65
Author: Evan Quiney <evanev7@gmail.com>
Date: Wed Jan 14 16:33:17 2026 +0000
add glm-47, minimax-m21 (#1147)
Adds support glm 4.7 and MiniMax M2.1
Manual testing:
Tensor + Pipeline execution of both models.
Closes #1141 and #1142
---
src/exo/shared/models/model_cards.py | 77 +++++++++++++++++++++++++++++
src/exo/worker/engines/mlx/auto_parallel.py | 40 ++++++++++++++-
2 files changed, 116 insertions(+), 1 deletion(-)
diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py
index fd4f9003..d46a8115 100644
--- a/src/exo/shared/models/model_cards.py
+++ b/src/exo/shared/models/model_cards.py
@@ -82,6 +82,7 @@ MODEL_CARDS: dict[str, ModelCard] = {
# storage_size=Memory.from_kb(754706307),
# n_layers=61,
# hidden_size=7168,
+ # supports_tensor=True,
# ),
# ),
# "deepseek-v3.2-4bit": ModelCard(
@@ -96,6 +97,7 @@ MODEL_CARDS: dict[str, ModelCard] = {
# storage_size=Memory.from_kb(754706307 // 2), # TODO !!!!!
# n_layers=61,
# hidden_size=7168,
+ # supports_tensor=True,
# ),
# ),
# deepseek r1
@@ -554,6 +556,81 @@ MODEL_CARDS: dict[str, ModelCard] = {
supports_tensor=True,
),
),
+ "glm-4.7-4bit": ModelCard(
+ short_id="glm-4.7-4bit",
+ model_id=ModelId("mlx-community/GLM-4.7-4bit"),
+ name="GLM 4.7 4bit",
+ description="GLM 4.7 4bit",
+ tags=[],
+ metadata=ModelMetadata(
+ model_id=ModelId("mlx-community/GLM-4.7-4bit"),
+ pretty_name="GLM 4.7 4bit",
+ storage_size=Memory.from_bytes(198556925568),
+ n_layers=91,
+ hidden_size=5120,
+ supports_tensor=True,
+ ),
+ ),
+ "glm-4.7-6bit": ModelCard(
+ short_id="glm-4.7-6bit",
+ model_id=ModelId("mlx-community/GLM-4.7-6bit"),
+ name="GLM 4.7 6bit",
+ description="GLM 4.7 6bit",
+ tags=[],
+ metadata=ModelMetadata(
+ model_id=ModelId("mlx-community/GLM-4.7-6bit"),
+ pretty_name="GLM 4.7 6bit",
+ storage_size=Memory.from_bytes(286737579648),
+ n_layers=91,
+ hidden_size=5120,
+ supports_tensor=True,
+ ),
+ ),
+ "glm-4.7-8bit-gs32": ModelCard(
+ short_id="glm-4.7-8bit-gs32",
+ model_id=ModelId("mlx-community/GLM-4.7-8bit-gs32"),
+ name="GLM 4.7 8bit (gs32)",
+ description="GLM 4.7 8bit (gs32)",
+ tags=[],
+ metadata=ModelMetadata(
+ model_id=ModelId("mlx-community/GLM-4.7-8bit-gs32"),
+ pretty_name="GLM 4.7 8bit (gs32)",
+ storage_size=Memory.from_bytes(396963397248),
+ n_layers=91,
+ hidden_size=5120,
+ supports_tensor=True,
+ ),
+ ),
+ "minimax-m2.1-8bit": ModelCard(
+ short_id="minimax-m2.1-8bit",
+ model_id=ModelId("mlx-community/MiniMax-M2.1-8bit"),
+ name="MiniMax M2.1 8bit",
+ description="MiniMax M2.1 8bit",
+ tags=[],
+ metadata=ModelMetadata(
+ model_id=ModelId("mlx-community/MiniMax-M2.1-8bit"),
+ pretty_name="MiniMax M2.1 8bit",
+ storage_size=Memory.from_bytes(242986745856),
+ n_layers=61,
+ hidden_size=3072,
+ supports_tensor=True,
+ ),
+ ),
+ "minimax-m2.1-3bit": ModelCard(
+ short_id="minimax-m2.1-3bit",
+ model_id=ModelId("mlx-community/MiniMax-M2.1-3bit"),
+ name="MiniMax M2.1 3bit",
+ description="MiniMax M2.1 3bit",
+ tags=[],
+ metadata=ModelMetadata(
+ model_id=ModelId("mlx-community/MiniMax-M2.1-3bit"),
+ pretty_name="MiniMax M2.1 3bit",
+ storage_size=Memory.from_bytes(100086644736),
+ n_layers=61,
+ hidden_size=3072,
+ supports_tensor=True,
+ ),
+ ),
# "devstral-2-123b-instruct-2512-8bit": ModelCard(
# short_id="devstral-2-123b-instruct-2512-8bit",
# model_id=ModelId("mlx-community/Devstral-2-123B-Instruct-2512-8bit"),
diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py
index 55a549fd..46d812e3 100644
--- a/src/exo/worker/engines/mlx/auto_parallel.py
+++ b/src/exo/worker/engines/mlx/auto_parallel.py
@@ -19,6 +19,7 @@ from mlx_lm.models.glm4_moe import MoE
from mlx_lm.models.gpt_oss import GptOssMoeModel
from mlx_lm.models.gpt_oss import Model as GptOssModel
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
from mlx_lm.models.qwen3_moe import Model as Qwen3MoeModel
from mlx_lm.models.qwen3_moe import Qwen3MoeSparseMoeBlock
@@ -252,6 +253,14 @@ def tensor_auto_parallel(
all_to_sharded_linear_in_place,
sharded_to_all_linear_in_place,
)
+ elif isinstance(model, MiniMaxModel):
+ tensor_parallel_sharding_strategy = MiniMaxShardingStrategy(
+ group,
+ all_to_sharded_linear,
+ sharded_to_all_linear,
+ all_to_sharded_linear_in_place,
+ sharded_to_all_linear_in_place,
+ )
elif isinstance(model, (Qwen3MoeModel, Glm4MoeModel, Qwen3NextModel)):
tensor_parallel_sharding_strategy = QwenShardingStrategy(
group,
@@ -394,6 +403,35 @@ class ShardedDeepseekV3MoE(CustomMlxLayer):
return y
+class MiniMaxShardingStrategy(TensorParallelShardingStrategy):
+ def shard_model(self, model: nn.Module) -> nn.Module:
+ model = cast(MiniMaxModel, model)
+ for layer in model.layers:
+ # Shard the self attention
+ layer.self_attn.q_proj = self.all_to_sharded_linear(layer.self_attn.q_proj)
+ 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)
+ layer.self_attn.num_attention_heads //= self.N
+ layer.self_attn.num_key_value_heads //= self.N
+
+ # Shard the MoE. Shard in place since the MoE should be responsible
+ # for aggregating the results.
+ self.all_to_sharded_linear_in_place(
+ layer.block_sparse_moe.switch_mlp.gate_proj
+ )
+ self.sharded_to_all_linear_in_place(
+ layer.block_sparse_moe.switch_mlp.down_proj
+ )
+ self.all_to_sharded_linear_in_place(
+ layer.block_sparse_moe.switch_mlp.up_proj
+ )
+ layer.block_sparse_moe = ShardedQwenMoE(layer.block_sparse_moe) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
+ layer.block_sparse_moe.sharding_group = self.group
+
+ return model
+
+
class QwenShardingStrategy(TensorParallelShardingStrategy):
def shard_model(self, model: nn.Module) -> nn.Module:
model = cast(Qwen3MoeModel, model)
@@ -414,7 +452,7 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.gate_proj)
self.sharded_to_all_linear_in_place(layer.mlp.switch_mlp.down_proj)
self.all_to_sharded_linear_in_place(layer.mlp.switch_mlp.up_proj)
- layer.mlp = ShardedQwenMoE(layer.mlp) # type: ignore
+ layer.mlp = ShardedQwenMoE(layer.mlp) # pyright: ignore[reportAttributeAccessIssue, reportArgumentType]
layer.mlp.sharding_group = self.group
# Shard the MLP
← 3671528f nix: add dashboard build with dream2nix
·
back to Exo
·
model_cards.py: clean up commented out code e0aab46f →