From 2994b41089b8a1eb76310bb3b9e5748c9c17bd2c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Mustafa=20Alp=20Y=C4=B1lmaz?= <96022931+mustafalpyilmaz@users.noreply.github.com> Date: Wed, 11 Mar 2026 16:46:33 +0300 Subject: [PATCH] fix: validate num_key_value_heads in tensor sharding placement (#1669) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## Problem Models with fewer KV heads than nodes crash during tensor parallelism. For example, Qwen3.5 MoE models have only 2 KV heads — trying to shard across 4 nodes produces empty tensors and a reshape error at runtime. The placement system already validates `hidden_size % num_nodes == 0` but doesn't check KV heads, so it creates configurations that look valid but blow up when the worker tries to split the attention heads. Affected models include Qwen3.5-35B-A3B, Qwen3.5-122B-A10B, Qwen3.5-397B-A17B, Qwen3-Next-80B-A3B, and Qwen3-Coder-Next (all have 2 KV heads). ## Changes **Placement validation** (`src/exo/master/placement.py`): - Combined KV heads divisibility check with the existing hidden_size filter in a single pass - Cycles where `num_key_value_heads % len(cycle) != 0` are now excluded for tensor sharding - Error message includes both constraints when no valid cycle is found **Model card schema** (`src/exo/shared/models/model_cards.py`): - Added optional `num_key_value_heads` field to `ModelCard` and `ConfigData` - Extracted from HuggingFace `config.json` (handles both top-level and `text_config` nesting) - Passed through in `fetch_from_hf()` for dynamically fetched cards **All 68 inference model cards** (`resources/inference_model_cards/*.toml`): - Populated `num_key_value_heads` from each model's HuggingFace config **Utility script** (`scripts/fetch_kv_heads.py`): - Fetches `num_key_value_heads` from HuggingFace and updates TOML cards - `--missing`: only fills in cards that don't have the field yet - `--all`: re-fetches and overwrites everything - Uses tomlkit for safe TOML editing and ThreadPoolExecutor for parallel fetches ## Behavior - Instance previews no longer show tensor options for models that can't split their KV heads across the cluster size - `place_instance()` rejects with a clear error instead of crash-looping - Pipeline parallelism is unaffected - 2-node tensor still works for 2-KV-head models (2 ÷ 2 = 1) - Field is optional — existing custom cards without it continue to work (validation is skipped when `None`) --- .../mlx-community--DeepSeek-V3.1-4bit.toml | 1 + .../mlx-community--DeepSeek-V3.1-8bit.toml | 1 + .../mlx-community--GLM-4.5-Air-8bit.toml | 1 + .../mlx-community--GLM-4.5-Air-bf16.toml | 1 + .../mlx-community--GLM-4.7-4bit.toml | 1 + .../mlx-community--GLM-4.7-6bit.toml | 1 + .../mlx-community--GLM-4.7-8bit-gs32.toml | 1 + .../mlx-community--GLM-4.7-Flash-4bit.toml | 1 + .../mlx-community--GLM-4.7-Flash-5bit.toml | 1 + .../mlx-community--GLM-4.7-Flash-6bit.toml | 1 + .../mlx-community--GLM-4.7-Flash-8bit.toml | 1 + .../mlx-community--GLM-5-8bit.toml | 1 + .../mlx-community--GLM-5-MXFP4-Q8.toml | 1 + .../mlx-community--GLM-5-bf16.toml | 1 + .../mlx-community--Kimi-K2-Instruct-4bit.toml | 1 + .../mlx-community--Kimi-K2-Thinking.toml | 1 + .../mlx-community--Kimi-K2.5.toml | 1 + ...community--Llama-3.2-1B-Instruct-4bit.toml | 1 + ...community--Llama-3.2-3B-Instruct-4bit.toml | 1 + ...community--Llama-3.2-3B-Instruct-8bit.toml | 1 + ...ommunity--Llama-3.3-70B-Instruct-4bit.toml | 1 + ...ommunity--Llama-3.3-70B-Instruct-8bit.toml | 1 + ...ity--Meta-Llama-3.1-70B-Instruct-4bit.toml | 1 + ...nity--Meta-Llama-3.1-8B-Instruct-4bit.toml | 1 + ...nity--Meta-Llama-3.1-8B-Instruct-8bit.toml | 1 + ...nity--Meta-Llama-3.1-8B-Instruct-bf16.toml | 1 + .../mlx-community--MiniMax-M2.1-3bit.toml | 1 + .../mlx-community--MiniMax-M2.1-8bit.toml | 1 + .../mlx-community--MiniMax-M2.5-4bit.toml | 1 + .../mlx-community--MiniMax-M2.5-6bit.toml | 1 + .../mlx-community--MiniMax-M2.5-8bit.toml | 1 + .../mlx-community--Qwen3-0.6B-4bit.toml | 1 + .../mlx-community--Qwen3-0.6B-8bit.toml | 1 + ...y--Qwen3-235B-A22B-Instruct-2507-4bit.toml | 1 + ...y--Qwen3-235B-A22B-Instruct-2507-8bit.toml | 1 + .../mlx-community--Qwen3-30B-A3B-4bit.toml | 1 + .../mlx-community--Qwen3-30B-A3B-8bit.toml | 1 + ...--Qwen3-Coder-480B-A35B-Instruct-4bit.toml | 1 + ...--Qwen3-Coder-480B-A35B-Instruct-8bit.toml | 1 + .../mlx-community--Qwen3-Coder-Next-4bit.toml | 1 + .../mlx-community--Qwen3-Coder-Next-5bit.toml | 1 + .../mlx-community--Qwen3-Coder-Next-6bit.toml | 1 + .../mlx-community--Qwen3-Coder-Next-8bit.toml | 1 + .../mlx-community--Qwen3-Coder-Next-bf16.toml | 1 + ...ity--Qwen3-Next-80B-A3B-Instruct-4bit.toml | 1 + ...ity--Qwen3-Next-80B-A3B-Instruct-8bit.toml | 1 + ...ity--Qwen3-Next-80B-A3B-Thinking-4bit.toml | 1 + ...ity--Qwen3-Next-80B-A3B-Thinking-8bit.toml | 1 + ...mlx-community--Qwen3.5-122B-A10B-4bit.toml | 1 + ...mlx-community--Qwen3.5-122B-A10B-6bit.toml | 1 + ...mlx-community--Qwen3.5-122B-A10B-8bit.toml | 1 + ...mlx-community--Qwen3.5-122B-A10B-bf16.toml | 1 + .../mlx-community--Qwen3.5-27B-4bit.toml | 1 + .../mlx-community--Qwen3.5-27B-8bit.toml | 1 + .../mlx-community--Qwen3.5-2B-MLX-8bit.toml | 1 + .../mlx-community--Qwen3.5-35B-A3B-4bit.toml | 1 + .../mlx-community--Qwen3.5-35B-A3B-8bit.toml | 1 + ...mlx-community--Qwen3.5-397B-A17B-4bit.toml | 1 + ...mlx-community--Qwen3.5-397B-A17B-6bit.toml | 1 + ...mlx-community--Qwen3.5-397B-A17B-8bit.toml | 1 + .../mlx-community--Qwen3.5-9B-4bit.toml | 1 + .../mlx-community--Qwen3.5-9B-8bit.toml | 1 + .../mlx-community--Step-3.5-Flash-4bit.toml | 1 + .../mlx-community--Step-3.5-Flash-6bit.toml | 1 + .../mlx-community--Step-3.5-Flash-8Bit.toml | 1 + .../mlx-community--gpt-oss-120b-MXFP4-Q8.toml | 1 + .../mlx-community--gpt-oss-20b-MXFP4-Q8.toml | 1 + ...ommunity--llama-3.3-70b-instruct-fp16.toml | 1 + scripts/fetch_kv_heads.py | 133 ++++++++++++++++++ src/exo/master/placement.py | 7 +- src/exo/shared/models/model_cards.py | 4 + 71 files changed, 211 insertions(+), 1 deletion(-) create mode 100755 scripts/fetch_kv_heads.py diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml index 2c982882..8c415095 100644 --- a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml +++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/DeepSeek-V3.1-4bit" n_layers = 61 hidden_size = 7168 +num_key_value_heads = 128 supports_tensor = true tasks = ["TextGeneration"] family = "deepseek" diff --git a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml index 4cf99bec..a1889395 100644 --- a/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml +++ b/resources/inference_model_cards/mlx-community--DeepSeek-V3.1-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/DeepSeek-V3.1-8bit" n_layers = 61 hidden_size = 7168 +num_key_value_heads = 128 supports_tensor = true tasks = ["TextGeneration"] family = "deepseek" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml index 0f8708bc..efa2c1f6 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.5-Air-8bit" n_layers = 46 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = false tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml index 764372fd..42727eed 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.5-Air-bf16.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.5-Air-bf16" n_layers = 46 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml index 559cea65..22fba6a9 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.7-4bit" n_layers = 91 hidden_size = 5120 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml index bdcd0b4e..b83c34e5 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-6bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.7-6bit" n_layers = 91 hidden_size = 5120 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml index 1d942440..3f6b2391 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-8bit-gs32.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.7-8bit-gs32" n_layers = 91 hidden_size = 5120 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml index 2aad9b5b..ca183ded 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.7-Flash-4bit" n_layers = 47 hidden_size = 2048 +num_key_value_heads = 20 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml index 1efed44f..3b0c06c9 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-5bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.7-Flash-5bit" n_layers = 47 hidden_size = 2048 +num_key_value_heads = 20 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml index b5ec0fa5..8669298f 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-6bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.7-Flash-6bit" n_layers = 47 hidden_size = 2048 +num_key_value_heads = 20 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml index 4b400aeb..3ffcacb2 100644 --- a/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-4.7-Flash-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-4.7-Flash-8bit" n_layers = 47 hidden_size = 2048 +num_key_value_heads = 20 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-5-8bit.toml b/resources/inference_model_cards/mlx-community--GLM-5-8bit.toml index 24beb79f..5affe179 100644 --- a/resources/inference_model_cards/mlx-community--GLM-5-8bit.toml +++ b/resources/inference_model_cards/mlx-community--GLM-5-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-5-8bit-MXFP8" n_layers = 78 hidden_size = 6144 +num_key_value_heads = 64 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-5-MXFP4-Q8.toml b/resources/inference_model_cards/mlx-community--GLM-5-MXFP4-Q8.toml index ceb1f74c..8e942ed3 100644 --- a/resources/inference_model_cards/mlx-community--GLM-5-MXFP4-Q8.toml +++ b/resources/inference_model_cards/mlx-community--GLM-5-MXFP4-Q8.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-5-MXFP4-Q8" n_layers = 78 hidden_size = 6144 +num_key_value_heads = 64 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--GLM-5-bf16.toml b/resources/inference_model_cards/mlx-community--GLM-5-bf16.toml index 18a7aec6..2803710e 100644 --- a/resources/inference_model_cards/mlx-community--GLM-5-bf16.toml +++ b/resources/inference_model_cards/mlx-community--GLM-5-bf16.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/GLM-5" n_layers = 78 hidden_size = 6144 +num_key_value_heads = 64 supports_tensor = true tasks = ["TextGeneration"] family = "glm" diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml index 3f21d4c0..da82c4d8 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Kimi-K2-Instruct-4bit" n_layers = 61 hidden_size = 7168 +num_key_value_heads = 64 supports_tensor = true tasks = ["TextGeneration"] family = "kimi" diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml b/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml index 3e7dedd2..d088b615 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2-Thinking.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Kimi-K2-Thinking" n_layers = 61 hidden_size = 7168 +num_key_value_heads = 64 supports_tensor = true tasks = ["TextGeneration"] family = "kimi" diff --git a/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml b/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml index eb73ea09..7bb689cc 100644 --- a/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml +++ b/resources/inference_model_cards/mlx-community--Kimi-K2.5.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Kimi-K2.5" n_layers = 61 hidden_size = 7168 +num_key_value_heads = 64 supports_tensor = true tasks = ["TextGeneration"] family = "kimi" diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml index b38ec20f..888b9f97 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-1B-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Llama-3.2-1B-Instruct-4bit" n_layers = 16 hidden_size = 2048 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml index 81ce4567..60486e64 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Llama-3.2-3B-Instruct-4bit" n_layers = 28 hidden_size = 3072 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml index ac9a203b..d252a2e5 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.2-3B-Instruct-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Llama-3.2-3B-Instruct-8bit" n_layers = 28 hidden_size = 3072 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml index 24c7cbaa..03f9c12b 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Llama-3.3-70B-Instruct-4bit" n_layers = 80 hidden_size = 8192 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml index 3bfc97dc..ef0254a8 100644 --- a/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Llama-3.3-70B-Instruct-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Llama-3.3-70B-Instruct-8bit" n_layers = 80 hidden_size = 8192 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml index 27d0b724..d4c57ae1 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-70B-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Meta-Llama-3.1-70B-Instruct-4bit" n_layers = 80 hidden_size = 8192 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml index 1fe34ba8..816564a0 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Meta-Llama-3.1-8B-Instruct-4bit" n_layers = 32 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml index 5310a2a0..ff82803a 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Meta-Llama-3.1-8B-Instruct-8bit" n_layers = 32 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml index eb6405e0..db19a097 100644 --- a/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml +++ b/resources/inference_model_cards/mlx-community--Meta-Llama-3.1-8B-Instruct-bf16.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Meta-Llama-3.1-8B-Instruct-bf16" n_layers = 32 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml index f740e0d5..ec1e4f50 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-3bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/MiniMax-M2.1-3bit" n_layers = 61 hidden_size = 3072 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "minimax" diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml index 6cf55637..c3bb8177 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.1-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/MiniMax-M2.1-8bit" n_layers = 61 hidden_size = 3072 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "minimax" diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.5-4bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.5-4bit.toml index d64f29f9..963b3b9b 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.5-4bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.5-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/MiniMax-M2.5-4bit" n_layers = 62 hidden_size = 3072 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "minimax" diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.5-6bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.5-6bit.toml index c92dcd0d..af329569 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.5-6bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.5-6bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/MiniMax-M2.5-6bit" n_layers = 62 hidden_size = 3072 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "minimax" diff --git a/resources/inference_model_cards/mlx-community--MiniMax-M2.5-8bit.toml b/resources/inference_model_cards/mlx-community--MiniMax-M2.5-8bit.toml index b1abb744..2a39701d 100644 --- a/resources/inference_model_cards/mlx-community--MiniMax-M2.5-8bit.toml +++ b/resources/inference_model_cards/mlx-community--MiniMax-M2.5-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/MiniMax-M2.5-8bit" n_layers = 62 hidden_size = 3072 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "minimax" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml index 03c6c28b..5b4bd368 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-0.6B-4bit" n_layers = 28 hidden_size = 1024 +num_key_value_heads = 8 supports_tensor = false tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml index 44177385..0dad33c0 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-0.6B-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-0.6B-8bit" n_layers = 28 hidden_size = 1024 +num_key_value_heads = 8 supports_tensor = false tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml index e8ef3494..4a4fb10f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-235B-A22B-Instruct-2507-4bit" n_layers = 94 hidden_size = 4096 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml index 1ce1fd62..ed70dca4 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-235B-A22B-Instruct-2507-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-235B-A22B-Instruct-2507-8bit" n_layers = 94 hidden_size = 4096 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml index bcae53f9..c0f7b577 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-30B-A3B-4bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml index ed08db63..4018ec2d 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-30B-A3B-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-30B-A3B-8bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml index b390bd21..e4838f9c 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Coder-480B-A35B-Instruct-4bit" n_layers = 62 hidden_size = 6144 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml index 1c21307c..44558b78 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-480B-A35B-Instruct-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Coder-480B-A35B-Instruct-8bit" n_layers = 62 hidden_size = 6144 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-4bit.toml index 5f494594..a81e0059 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Coder-Next-4bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-5bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-5bit.toml index 60a5f131..535b62eb 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-5bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-5bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Coder-Next-5bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-6bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-6bit.toml index 59ede911..4ad7289d 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-6bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-6bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Coder-Next-6bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-8bit.toml index 1e8b0f12..51f13a63 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Coder-Next-8bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-bf16.toml b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-bf16.toml index 5db8fae9..76d77693 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-bf16.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Coder-Next-bf16.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Coder-Next-bf16" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml index 386a3fa1..d9eb5845 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Next-80B-A3B-Instruct-4bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml index 0e2bf2a5..09352d3f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Instruct-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Next-80B-A3B-Instruct-8bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml index 48caa3fc..12f341fb 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Next-80B-A3B-Thinking-4bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml index bd4ae549..7a52dbf7 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3-Next-80B-A3B-Thinking-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3-Next-80B-A3B-Thinking-8bit" n_layers = 48 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-4bit.toml index cb4d1b24..452b0a44 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-122B-A10B-4bit" n_layers = 48 hidden_size = 3072 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-6bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-6bit.toml index caa4dcca..7c01c831 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-6bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-6bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-122B-A10B-6bit" n_layers = 48 hidden_size = 3072 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-8bit.toml index cf11daa8..12c6a460 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-122B-A10B-8bit" n_layers = 48 hidden_size = 3072 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-bf16.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-bf16.toml index 419572fe..a4b22916 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-bf16.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-bf16.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-122B-A10B-bf16" n_layers = 48 hidden_size = 3072 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-27B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-27B-4bit.toml index 2fe284cb..ee42738f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-27B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-27B-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-27B-4bit" n_layers = 64 hidden_size = 5120 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-27B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-27B-8bit.toml index fb701acc..42b1184c 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-27B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-27B-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-27B-8bit" n_layers = 64 hidden_size = 5120 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-2B-MLX-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-2B-MLX-8bit.toml index 867545e8..b1821192 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-2B-MLX-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-2B-MLX-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-2B-MLX-8bit" n_layers = 24 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-4bit.toml index 11158c38..1c0a6fe9 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-35B-A3B-4bit" n_layers = 40 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-8bit.toml index c8b282cd..64b7a61e 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-35B-A3B-8bit" n_layers = 40 hidden_size = 2048 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-4bit.toml index 75251cf3..10af0e43 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-397B-A17B-4bit" n_layers = 60 hidden_size = 4096 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-6bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-6bit.toml index d0512f8f..b8427aa1 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-6bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-6bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-397B-A17B-6bit" n_layers = 60 hidden_size = 4096 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-8bit.toml index 7aa235cb..73fdbcd3 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-397B-A17B-8bit" n_layers = 60 hidden_size = 4096 +num_key_value_heads = 2 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-9B-4bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-9B-4bit.toml index a1d9f828..67a4970f 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-9B-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-9B-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-9B-4bit" n_layers = 32 hidden_size = 4096 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Qwen3.5-9B-8bit.toml b/resources/inference_model_cards/mlx-community--Qwen3.5-9B-8bit.toml index 289f1899..3dcf5b0d 100644 --- a/resources/inference_model_cards/mlx-community--Qwen3.5-9B-8bit.toml +++ b/resources/inference_model_cards/mlx-community--Qwen3.5-9B-8bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Qwen3.5-9B-8bit" n_layers = 32 hidden_size = 4096 +num_key_value_heads = 4 supports_tensor = true tasks = ["TextGeneration"] family = "qwen" diff --git a/resources/inference_model_cards/mlx-community--Step-3.5-Flash-4bit.toml b/resources/inference_model_cards/mlx-community--Step-3.5-Flash-4bit.toml index 78385661..9dd73125 100644 --- a/resources/inference_model_cards/mlx-community--Step-3.5-Flash-4bit.toml +++ b/resources/inference_model_cards/mlx-community--Step-3.5-Flash-4bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Step-3.5-Flash-4bit" n_layers = 45 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "step" diff --git a/resources/inference_model_cards/mlx-community--Step-3.5-Flash-6bit.toml b/resources/inference_model_cards/mlx-community--Step-3.5-Flash-6bit.toml index ab336bc2..f910f25f 100644 --- a/resources/inference_model_cards/mlx-community--Step-3.5-Flash-6bit.toml +++ b/resources/inference_model_cards/mlx-community--Step-3.5-Flash-6bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Step-3.5-Flash-6bit" n_layers = 45 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "step" diff --git a/resources/inference_model_cards/mlx-community--Step-3.5-Flash-8Bit.toml b/resources/inference_model_cards/mlx-community--Step-3.5-Flash-8Bit.toml index ef2a5074..eb4d635c 100644 --- a/resources/inference_model_cards/mlx-community--Step-3.5-Flash-8Bit.toml +++ b/resources/inference_model_cards/mlx-community--Step-3.5-Flash-8Bit.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/Step-3.5-Flash-8Bit" n_layers = 45 hidden_size = 4096 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "step" diff --git a/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml b/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml index f579c618..900c18f5 100644 --- a/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml +++ b/resources/inference_model_cards/mlx-community--gpt-oss-120b-MXFP4-Q8.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/gpt-oss-120b-MXFP4-Q8" n_layers = 36 hidden_size = 2880 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "gpt-oss" diff --git a/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml b/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml index af1e04ad..98748fd1 100644 --- a/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml +++ b/resources/inference_model_cards/mlx-community--gpt-oss-20b-MXFP4-Q8.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/gpt-oss-20b-MXFP4-Q8" n_layers = 24 hidden_size = 2880 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "gpt-oss" diff --git a/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml b/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml index e61660c2..a46dbb88 100644 --- a/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml +++ b/resources/inference_model_cards/mlx-community--llama-3.3-70b-instruct-fp16.toml @@ -1,6 +1,7 @@ model_id = "mlx-community/llama-3.3-70b-instruct-fp16" n_layers = 80 hidden_size = 8192 +num_key_value_heads = 8 supports_tensor = true tasks = ["TextGeneration"] family = "llama" diff --git a/scripts/fetch_kv_heads.py b/scripts/fetch_kv_heads.py new file mode 100755 index 00000000..58188a96 --- /dev/null +++ b/scripts/fetch_kv_heads.py @@ -0,0 +1,133 @@ +#!/usr/bin/env python3 +"""Fetch num_key_value_heads from HuggingFace config.json and update TOML model cards. + +Usage: + # Update only cards missing num_key_value_heads + uv run python scripts/fetch_kv_heads.py --missing + + # Update all cards (overwrite existing values) + uv run python scripts/fetch_kv_heads.py --all +""" + +from __future__ import annotations + +import argparse +import json +import sys +import urllib.request +from concurrent.futures import ThreadPoolExecutor, as_completed +from pathlib import Path + +import tomlkit + +CARDS_DIR = ( + Path(__file__).resolve().parent.parent / "resources" / "inference_model_cards" +) +MAX_WORKERS = 5 + + +def fetch_kv_heads(model_id: str) -> int | None: + """Fetch num_key_value_heads from HuggingFace config.json.""" + url = f"https://huggingface.co/{model_id}/raw/main/config.json" + try: + with urllib.request.urlopen(url, timeout=15) as resp: + config = json.loads(resp.read()) + except Exception as e: + print(f" ERROR fetching {url}: {e}", file=sys.stderr) + return None + + for source in [config, config.get("text_config", {})]: + if "num_key_value_heads" in source: + return int(source["num_key_value_heads"]) + + return None + + +def update_toml(path: Path, kv_heads: int) -> bool: + """Insert or update num_key_value_heads in a TOML file. Returns True if changed.""" + content = path.read_text() + doc = tomlkit.parse(content) + + if doc.get("num_key_value_heads") == kv_heads: + return False + + # Insert after hidden_size if adding for the first time + if "num_key_value_heads" not in doc: + new_doc = tomlkit.document() + for key, value in doc.items(): + new_doc[key] = value + if key == "hidden_size": + new_doc["num_key_value_heads"] = kv_heads + path.write_text(tomlkit.dumps(new_doc)) + else: + doc["num_key_value_heads"] = kv_heads + path.write_text(tomlkit.dumps(doc)) + + return True + + +def process_card(path: Path) -> tuple[str, str]: + """Fetch and update a single card. Returns (filename, status).""" + content = path.read_text() + doc = tomlkit.parse(content) + model_id = doc.get("model_id") + if not model_id: + return path.name, "SKIP (no model_id)" + + kv_heads = fetch_kv_heads(str(model_id)) + if kv_heads is None: + return path.name, "FAILED" + + changed = update_toml(path, kv_heads) + return path.name, f"{kv_heads} ({'UPDATED' if changed else 'UNCHANGED'})" + + +def main(): + parser = argparse.ArgumentParser( + description="Fetch num_key_value_heads from HuggingFace and update TOML cards." + ) + group = parser.add_mutually_exclusive_group(required=True) + group.add_argument( + "--all", + action="store_true", + help="Update all model cards (overwrite existing values)", + ) + group.add_argument( + "--missing", + action="store_true", + help="Only update cards missing num_key_value_heads", + ) + args = parser.parse_args() + + toml_files = sorted(CARDS_DIR.glob("*.toml")) + if not toml_files: + print(f"No TOML files found in {CARDS_DIR}", file=sys.stderr) + sys.exit(1) + + to_process = [] + skipped = 0 + + for path in toml_files: + if args.missing and "num_key_value_heads" in path.read_text(): + skipped += 1 + continue + to_process.append(path) + + updated = 0 + failed = 0 + + with ThreadPoolExecutor(max_workers=MAX_WORKERS) as pool: + futures = {pool.submit(process_card, path): path for path in to_process} + for future in as_completed(futures): + name, status = future.result() + print(f" {name}: {status}") + if "UPDATED" in status: + updated += 1 + elif "FAILED" in status: + failed += 1 + + print(f"\nDone: {updated} updated, {skipped} skipped, {failed} failed") + + +if __name__ == "__main__": + main() diff --git a/src/exo/master/placement.py b/src/exo/master/placement.py index c7fc95e5..f9e9de73 100644 --- a/src/exo/master/placement.py +++ b/src/exo/master/placement.py @@ -90,14 +90,19 @@ def place_instance( f"Requested Tensor sharding but this model does not support tensor parallelism: {command.model_card.model_id}" ) # TODO: the condition here for tensor parallel is not correct, but it works good enough for now. + kv_heads = command.model_card.num_key_value_heads cycles_with_sufficient_memory = [ cycle for cycle in cycles_with_sufficient_memory if command.model_card.hidden_size % len(cycle) == 0 + and (kv_heads is None or kv_heads % len(cycle) == 0) ] if not cycles_with_sufficient_memory: raise ValueError( - f"No tensor sharding found for model with hidden_size {command.model_card.hidden_size} candidate cycles" + f"No tensor sharding found for model with " + f"hidden_size={command.model_card.hidden_size}" + f"{f', num_key_value_heads={kv_heads}' if kv_heads is not None else ''}" + f" across candidate cycles" ) if command.sharding == Sharding.Pipeline and command.model_card.model_id == ModelId( "mlx-community/DeepSeek-V3.1-8bit" diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py index f62363a0..048c469f 100644 --- a/src/exo/shared/models/model_cards.py +++ b/src/exo/shared/models/model_cards.py @@ -83,6 +83,7 @@ class ModelCard(CamelCaseModel): n_layers: PositiveInt hidden_size: PositiveInt supports_tensor: bool + num_key_value_heads: PositiveInt | None = None tasks: list[ModelTask] components: list[ComponentInfo] | None = None family: str = "" @@ -137,6 +138,7 @@ class ModelCard(CamelCaseModel): n_layers=num_layers, hidden_size=config_data.hidden_size or 0, supports_tensor=config_data.supports_tensor, + num_key_value_heads=config_data.num_key_value_heads, tasks=[ModelTask.TextGeneration], trust_remote_code=False, ) @@ -170,6 +172,7 @@ class ConfigData(BaseModel): architectures: list[str] | None = None hidden_size: Annotated[int, Field(ge=0)] | None = None + num_key_value_heads: PositiveInt | None = None layer_count: int = Field( validation_alias=AliasChoices( "num_hidden_layers", @@ -209,6 +212,7 @@ class ConfigData(BaseModel): for field in [ "architectures", "hidden_size", + "num_key_value_heads", "num_hidden_layers", "num_layers", "n_layer",