[object Object]

← back to Exo

Add support for Qwen3.5 (#1644)

28817d3ee30282aba1ea1fd0ebc16babaf851141 · 2026-03-03 23:31:57 +0900 · Daiz

## Motivation

Qwen3.5 MoE models (e.g., `Qwen3.5-397B-A17B-6bit`) are now supported by
`mlx-lm` via `qwen3_5_moe` model type, but exo lacks tensor parallel
sharding support for this architecture. This prevents running large
Qwen3.5 models across multiple nodes.

Qwen3.5 uses a GatedDeltaNet hybrid attention mechanism similar to
Qwen3-Next, but with a different projection layout — separate
`in_proj_qkv`, `in_proj_z`, `in_proj_b`, `in_proj_a` instead of
Qwen3-Next's combined `in_proj_qkvz` and `in_proj_ba`. This requires
architecture-aware sharding logic.

## Changes (evan summary)

- enable qwen3_5 dense + moe tensor parallelism from config
- defensively skip evalling _cache.keys if it doesn't exist
- ignore kwargs in qwen35 pipeline masking and ensure pipeline segments match global model parameters for mask creation
- add sharding for qwen3_5 moe linear attention
- added another 6 million model cards

## Why It Works

Qwen3.5's GatedDeltaNet has an `in_proj_qkv` linear layer with three
concatenated sections: `[q(key_dim), k(key_dim), v(value_dim)]`. A naive
contiguous split (`segments=1`) would slice across section boundaries,
corrupting q/k/v values and producing garbled output.

By passing `segments=[key_dim, key_dim + key_dim]` to `shard_linear()`,
each section is split independently before distributing across devices.
This ensures every rank receives correctly aligned q, k, and v
components.

The remaining separate projections (`in_proj_z`, `in_proj_b`,
`in_proj_a`) and the MoE layers follow the same `all_to_sharded` /
`sharded_to_all` pattern already used for Qwen3-Next.

Some pipeline splits didn't include an ssm layer or a linear layer resulting in a subset of the model acting like it shouldn't create the appropriate masks for the next layer - we patch the model to manually create such masks.

## Test Plan

tensor sharded 2,3,4 models & pipeline sharded 2,3,4 with simple eval.

---------

Co-authored-by: hw <hw@hwStudio1.local>
Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
Co-authored-by: Evan <evanev7@gmail.com>

Files touched

Diff

commit 28817d3ee30282aba1ea1fd0ebc16babaf851141
Author: Daiz <90135051+0xDaizz@users.noreply.github.com>
Date:   Tue Mar 3 23:31:57 2026 +0900

    Add support for Qwen3.5 (#1644)
    
    ## Motivation
    
    Qwen3.5 MoE models (e.g., `Qwen3.5-397B-A17B-6bit`) are now supported by
    `mlx-lm` via `qwen3_5_moe` model type, but exo lacks tensor parallel
    sharding support for this architecture. This prevents running large
    Qwen3.5 models across multiple nodes.
    
    Qwen3.5 uses a GatedDeltaNet hybrid attention mechanism similar to
    Qwen3-Next, but with a different projection layout — separate
    `in_proj_qkv`, `in_proj_z`, `in_proj_b`, `in_proj_a` instead of
    Qwen3-Next's combined `in_proj_qkvz` and `in_proj_ba`. This requires
    architecture-aware sharding logic.
    
    ## Changes (evan summary)
    
    - enable qwen3_5 dense + moe tensor parallelism from config
    - defensively skip evalling _cache.keys if it doesn't exist
    - ignore kwargs in qwen35 pipeline masking and ensure pipeline segments match global model parameters for mask creation
    - add sharding for qwen3_5 moe linear attention
    - added another 6 million model cards
    
    ## Why It Works
    
    Qwen3.5's GatedDeltaNet has an `in_proj_qkv` linear layer with three
    concatenated sections: `[q(key_dim), k(key_dim), v(value_dim)]`. A naive
    contiguous split (`segments=1`) would slice across section boundaries,
    corrupting q/k/v values and producing garbled output.
    
    By passing `segments=[key_dim, key_dim + key_dim]` to `shard_linear()`,
    each section is split independently before distributing across devices.
    This ensures every rank receives correctly aligned q, k, and v
    components.
    
    The remaining separate projections (`in_proj_z`, `in_proj_b`,
    `in_proj_a`) and the MoE layers follow the same `all_to_sharded` /
    `sharded_to_all` pattern already used for Qwen3-Next.
    
    Some pipeline splits didn't include an ssm layer or a linear layer resulting in a subset of the model acting like it shouldn't create the appropriate masks for the next layer - we patch the model to manually create such masks.
    
    ## Test Plan
    
    tensor sharded 2,3,4 models & pipeline sharded 2,3,4 with simple eval.
    
    ---------
    
    Co-authored-by: hw <hw@hwStudio1.local>
    Co-authored-by: Ryuichi Leo Takashige <leo@exolabs.net>
    Co-authored-by: Evan <evanev7@gmail.com>
---
 .mlx_typings/mlx_lm/models/cache.pyi               |   8 +-
 .mlx_typings/mlx_lm/models/qwen3_5.pyi             | 153 +++++++++++++++++++++
 .mlx_typings/mlx_lm/models/qwen3_5_moe.pyi         |  19 +++
 .mlx_typings/mlx_lm/models/qwen3_next.pyi          |   9 ++
 pyproject.toml                                     |   4 +-
 .../mlx-community--Qwen3.5-122B-A10B-4bit.toml     |  12 ++
 .../mlx-community--Qwen3.5-122B-A10B-6bit.toml     |  12 ++
 .../mlx-community--Qwen3.5-122B-A10B-8bit.toml     |  12 ++
 .../mlx-community--Qwen3.5-122B-A10B-bf16.toml     |  12 ++
 .../mlx-community--Qwen3.5-27B-4bit.toml           |  12 ++
 .../mlx-community--Qwen3.5-27B-8bit.toml           |  12 ++
 .../mlx-community--Qwen3.5-2B-MLX-8bit.toml        |  12 ++
 .../mlx-community--Qwen3.5-35B-A3B-4bit.toml       |  12 ++
 .../mlx-community--Qwen3.5-35B-A3B-8bit.toml       |  12 ++
 .../mlx-community--Qwen3.5-397B-A17B-4bit.toml     |  12 ++
 .../mlx-community--Qwen3.5-397B-A17B-6bit.toml     |  12 ++
 .../mlx-community--Qwen3.5-397B-A17B-8bit.toml     |  12 ++
 .../mlx-community--Qwen3.5-9B-4bit.toml            |  12 ++
 .../mlx-community--Qwen3.5-9B-8bit.toml            |  12 ++
 src/exo/shared/models/model_cards.py               |   2 +
 src/exo/worker/engines/mlx/auto_parallel.py        | 122 ++++++++++++++--
 src/exo/worker/engines/mlx/utils_mlx.py            |   3 +
 tmp/gen_card.py                                    |  33 +++++
 uv.lock                                            |  10 +-
 24 files changed, 503 insertions(+), 28 deletions(-)

diff --git a/.mlx_typings/mlx_lm/models/cache.pyi b/.mlx_typings/mlx_lm/models/cache.pyi
index efe31ca3..57f9aa7e 100644
--- a/.mlx_typings/mlx_lm/models/cache.pyi
+++ b/.mlx_typings/mlx_lm/models/cache.pyi
@@ -164,8 +164,9 @@ class KVCache(_BaseCache):
     def to_quantized(
         self, group_size: int = ..., bits: int = ...
     ) -> QuantizedKVCache: ...
-    def make_mask(self, *args, **kwargs):  # -> array | Literal['causal'] | None:
-        ...
+    def make_mask(
+        self, *args: Any, **kwargs: Any
+    ) -> mx.array | Literal["causal"] | None: ...
 
 class RotatingKVCache(_BaseCache):
     step = ...
@@ -218,8 +219,7 @@ class ArraysCache(_BaseCache):
         In-place extend this cache with the other cache.
         """
 
-    def make_mask(self, N: int):  # -> array | None:
-        ...
+    def make_mask(self, N: int) -> mx.array | None: ...
 
 class MambaCache(ArraysCache):
     def __init__(self, left_padding: Optional[List[int]] = ...) -> None: ...
diff --git a/.mlx_typings/mlx_lm/models/qwen3_5.pyi b/.mlx_typings/mlx_lm/models/qwen3_5.pyi
new file mode 100644
index 00000000..890450c7
--- /dev/null
+++ b/.mlx_typings/mlx_lm/models/qwen3_5.pyi
@@ -0,0 +1,153 @@
+from dataclasses import dataclass
+from typing import Any, Optional
+
+import mlx.core as mx
+import mlx.nn as nn
+
+from .cache import ArraysCache, KVCache
+from .qwen3_next import (
+    Qwen3NextAttention as Attention,
+    Qwen3NextMLP as MLP,
+    Qwen3NextRMSNormGated as RMSNormGated,
+    Qwen3NextSparseMoeBlock,
+)
+
+SparseMoeBlock = Qwen3NextSparseMoeBlock
+from .switch_layers import SwitchGLU
+
+@dataclass
+class TextModelArgs:
+    model_type: str
+    hidden_size: int
+    intermediate_size: int
+    num_hidden_layers: int
+    num_attention_heads: int
+    rms_norm_eps: float
+    vocab_size: int
+    num_key_value_heads: int
+    max_position_embeddings: int
+    linear_num_value_heads: int
+    linear_num_key_heads: int
+    linear_key_head_dim: int
+    linear_value_head_dim: int
+    linear_conv_kernel_dim: int
+    tie_word_embeddings: bool
+    attention_bias: bool
+    head_dim: Optional[int]
+    full_attention_interval: int
+    num_experts: int
+    num_experts_per_tok: int
+    decoder_sparse_step: int
+    shared_expert_intermediate_size: int
+    moe_intermediate_size: int
+    norm_topk_prob: bool
+    rope_parameters: Optional[dict[str, Any]]
+    partial_rotary_factor: float
+    rope_theta: float
+    rope_scaling: Optional[dict[str, Any]]
+
+    @classmethod
+    def from_dict(cls, params: dict[str, Any]) -> TextModelArgs: ...
+    def __post_init__(self) -> None: ...
+
+class GatedDeltaNet(nn.Module):
+    hidden_size: int
+    num_v_heads: int
+    num_k_heads: int
+    head_k_dim: int
+    head_v_dim: int
+    key_dim: int
+    value_dim: int
+    conv_kernel_size: int
+    conv_dim: int
+    conv1d: nn.Conv1d
+    in_proj_qkv: nn.Linear
+    in_proj_z: nn.Linear
+    in_proj_b: nn.Linear
+    in_proj_a: nn.Linear
+    dt_bias: mx.array
+    A_log: mx.array
+    norm: RMSNormGated
+    out_proj: nn.Linear
+
+    def __init__(self, config: TextModelArgs) -> None: ...
+    def __call__(
+        self,
+        inputs: mx.array,
+        mask: Optional[mx.array] = None,
+        cache: Optional[Any] = None,
+    ) -> mx.array: ...
+
+class DecoderLayer(nn.Module):
+    is_linear: bool
+    linear_attn: GatedDeltaNet
+    self_attn: Attention
+    input_layernorm: nn.RMSNorm
+    post_attention_layernorm: nn.RMSNorm
+    mlp: MLP | SparseMoeBlock
+
+    def __init__(self, args: TextModelArgs, layer_idx: int) -> None: ...
+    def __call__(
+        self,
+        x: mx.array,
+        mask: Optional[mx.array] = None,
+        cache: Optional[Any] = None,
+    ) -> mx.array: ...
+
+class Qwen3_5TextModel(nn.Module):
+    embed_tokens: nn.Embedding
+    layers: list[DecoderLayer]
+    norm: nn.RMSNorm
+    ssm_idx: int
+    fa_idx: int
+
+    def __init__(self, args: TextModelArgs) -> None: ...
+    def __call__(
+        self,
+        inputs: mx.array,
+        cache: Optional[Any] = None,
+        input_embeddings: Optional[mx.array] = None,
+    ) -> mx.array: ...
+
+class TextModel(nn.Module):
+    args: TextModelArgs
+    model_type: str
+    model: Qwen3_5TextModel
+    lm_head: nn.Linear
+
+    def __init__(self, args: TextModelArgs) -> None: ...
+    def __call__(
+        self,
+        inputs: mx.array,
+        cache: Optional[Any] = None,
+        input_embeddings: Optional[mx.array] = None,
+    ) -> mx.array: ...
+    @property
+    def layers(self) -> list[DecoderLayer]: ...
+    def make_cache(self) -> list[ArraysCache | KVCache]: ...
+    def sanitize(self, weights: dict[str, Any]) -> dict[str, Any]: ...
+
+@dataclass
+class ModelArgs:
+    model_type: str
+    text_config: dict[str, Any]
+
+    @classmethod
+    def from_dict(cls, params: dict[str, Any]) -> ModelArgs: ...
+
+class Model(nn.Module):
+    args: ModelArgs
+    model_type: str
+    language_model: TextModel
+
+    def __init__(self, args: ModelArgs) -> None: ...
+    def __call__(
+        self,
+        inputs: mx.array,
+        cache: Optional[Any] = None,
+        input_embeddings: Optional[mx.array] = None,
+    ) -> mx.array: ...
+    def sanitize(self, weights: dict[str, Any]) -> dict[str, Any]: ...
+    @property
+    def layers(self) -> list[DecoderLayer]: ...
+    def make_cache(self) -> list[ArraysCache | KVCache]: ...
diff --git a/.mlx_typings/mlx_lm/models/qwen3_5_moe.pyi b/.mlx_typings/mlx_lm/models/qwen3_5_moe.pyi
new file mode 100644
index 00000000..51aae5c9
--- /dev/null
+++ b/.mlx_typings/mlx_lm/models/qwen3_5_moe.pyi
@@ -0,0 +1,19 @@
+from dataclasses import dataclass
+from typing import Any, Optional
+
+import mlx.core as mx
+import mlx.nn as nn
+
+from .cache import ArraysCache, KVCache
+from .qwen3_5 import DecoderLayer, Model as Qwen3_5Model, TextModel
+
+@dataclass
+class ModelArgs:
+    model_type: str
+    text_config: dict[str, Any]
+
+    @classmethod
+    def from_dict(cls, params: dict[str, Any]) -> ModelArgs: ...
+
+class Model(Qwen3_5Model):
+    def sanitize(self, weights: dict[str, Any]) -> dict[str, Any]: ...
diff --git a/.mlx_typings/mlx_lm/models/qwen3_next.pyi b/.mlx_typings/mlx_lm/models/qwen3_next.pyi
index 9f0b940c..10a508c0 100644
--- a/.mlx_typings/mlx_lm/models/qwen3_next.pyi
+++ b/.mlx_typings/mlx_lm/models/qwen3_next.pyi
@@ -7,6 +7,15 @@ import mlx.nn as nn
 
 from .switch_layers import SwitchGLU
 
+class Qwen3NextRMSNormGated(nn.Module):
+    eps: float
+    weight: mx.array
+
+    def __init__(self, hidden_size: int, eps: float = ...) -> None: ...
+    def __call__(
+        self, hidden_states: mx.array, gate: mx.array | None = None
+    ) -> mx.array: ...
+
 class Qwen3NextMLP(nn.Module):
     gate_proj: nn.Linear
     down_proj: nn.Linear
diff --git a/pyproject.toml b/pyproject.toml
index 6f28de11..86293059 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -19,7 +19,7 @@ dependencies = [
   "anyio==4.11.0",
   "mlx; sys_platform == 'darwin'",
   "mlx[cpu]==0.30.6; sys_platform == 'linux'",
-  "mlx-lm==0.30.7",
+  "mlx-lm",
   "tiktoken>=0.12.0",                          # required for kimi k2 tokenizer
   "hypercorn>=0.18.0",
   "openai-harmony>=0.0.8",
@@ -62,7 +62,7 @@ members = ["rust/exo_pyo3_bindings", "bench"]
 [tool.uv.sources]
 exo_pyo3_bindings = { workspace = true }
 mlx = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git", branch = "address-rdma-gpu-locks", marker = "sys_platform == 'darwin'" }
-#mlx-lm = { git = "https://github.com/davidmcc73/mlx-lm", branch = "stable" }
+mlx-lm = { git = "https://github.com/ml-explore/mlx-lm", rev = "834fac934c4e04de9b3d723e2b9287a2c60cfd4a" }
 # 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/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
new file mode 100644
index 00000000..cb4d1b24
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-4bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-122B-A10B-4bit"
+n_layers = 48
+hidden_size = 3072
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "4bit"
+base_model = "Qwen3.5 122B A10B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 69593314272
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
new file mode 100644
index 00000000..caa4dcca
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-6bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-122B-A10B-6bit"
+n_layers = 48
+hidden_size = 3072
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "6bit"
+base_model = "Qwen3.5 122B A10B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 100120675296
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
new file mode 100644
index 00000000..cf11daa8
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-122B-A10B-8bit"
+n_layers = 48
+hidden_size = 3072
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "8bit"
+base_model = "Qwen3.5 122B A10B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 130648036320
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
new file mode 100644
index 00000000..419572fe
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-122B-A10B-bf16.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-122B-A10B-bf16"
+n_layers = 48
+hidden_size = 3072
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "bf16"
+base_model = "Qwen3.5 122B A10B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 245125640160
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
new file mode 100644
index 00000000..2fe284cb
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-27B-4bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-27B-4bit"
+n_layers = 64
+hidden_size = 5120
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "4bit"
+base_model = "Qwen3.5 27B"
+capabilities = ["text", "thinking"]
+
+[storage_size]
+in_bytes = 16054266848
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
new file mode 100644
index 00000000..fb701acc
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-27B-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-27B-8bit"
+n_layers = 64
+hidden_size = 5120
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "8bit"
+base_model = "Qwen3.5 27B"
+capabilities = ["text", "thinking"]
+
+[storage_size]
+in_bytes = 29500943328
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
new file mode 100644
index 00000000..867545e8
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-2B-MLX-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-2B-MLX-8bit"
+n_layers = 24
+hidden_size = 2048
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "8bit"
+base_model = "Qwen3.5 2B"
+capabilities = ["text", "thinking"]
+
+[storage_size]
+in_bytes = 2662787264
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
new file mode 100644
index 00000000..11158c38
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-4bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-35B-A3B-4bit"
+n_layers = 40
+hidden_size = 2048
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "4bit"
+base_model = "Qwen3.5 35B A3B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 20391405152
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
new file mode 100644
index 00000000..c8b282cd
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-35B-A3B-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-35B-A3B-8bit"
+n_layers = 40
+hidden_size = 2048
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "8bit"
+base_model = "Qwen3.5 35B A3B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 37721130592
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
new file mode 100644
index 00000000..75251cf3
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-4bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-397B-A17B-4bit"
+n_layers = 60
+hidden_size = 4096
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "4bit"
+base_model = "Qwen3.5 397B A17B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 223860768352
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
new file mode 100644
index 00000000..d0512f8f
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-6bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-397B-A17B-6bit"
+n_layers = 60
+hidden_size = 4096
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "6bit"
+base_model = "Qwen3.5 397B A17B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 322946674272
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
new file mode 100644
index 00000000..7aa235cb
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-397B-A17B-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-397B-A17B-8bit"
+n_layers = 60
+hidden_size = 4096
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "8bit"
+base_model = "Qwen3.5 397B A17B"
+capabilities = ["text", "thinking", "thinking_toggle"]
+
+[storage_size]
+in_bytes = 422032580192
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
new file mode 100644
index 00000000..a1d9f828
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-9B-4bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-9B-4bit"
+n_layers = 32
+hidden_size = 4096
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "4bit"
+base_model = "Qwen3.5 9B"
+capabilities = ["text", "thinking"]
+
+[storage_size]
+in_bytes = 5950062560
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
new file mode 100644
index 00000000..289f1899
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Qwen3.5-9B-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Qwen3.5-9B-8bit"
+n_layers = 32
+hidden_size = 4096
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "qwen"
+quantization = "8bit"
+base_model = "Qwen3.5 9B"
+capabilities = ["text", "thinking"]
+
+[storage_size]
+in_bytes = 10426433504
diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py
index 83420541..b6d94e74 100644
--- a/src/exo/shared/models/model_cards.py
+++ b/src/exo/shared/models/model_cards.py
@@ -190,6 +190,8 @@ class ConfigData(BaseModel):
             ["DeepseekV3ForCausalLM"],
             ["Qwen3NextForCausalLM"],
             ["Qwen3MoeForCausalLM"],
+            ["Qwen3_5MoeForConditionalGeneration"],
+            ["Qwen3_5ForConditionalGeneration"],
             ["MiniMaxM2ForCausalLM"],
             ["LlamaForCausalLM"],
             ["GptOssForCausalLM"],
diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py
index 9ada458f..daf010ff 100644
--- a/src/exo/worker/engines/mlx/auto_parallel.py
+++ b/src/exo/worker/engines/mlx/auto_parallel.py
@@ -16,6 +16,7 @@ from mlx.nn.layers.distributed import (
 from mlx_lm.models.base import (
     scaled_dot_product_attention,  # pyright: ignore[reportUnknownVariableType]
 )
+from mlx_lm.models.cache import ArraysCache, KVCache
 from mlx_lm.models.deepseek_v3 import DeepseekV3MLP
 from mlx_lm.models.deepseek_v3 import Model as DeepseekV3Model
 from mlx_lm.models.deepseek_v32 import DeepseekV32MLP
@@ -31,10 +32,19 @@ from mlx_lm.models.llama import Model as LlamaModel
 from mlx_lm.models.minimax import MiniMaxAttention
 from mlx_lm.models.minimax import Model as MiniMaxModel
 from mlx_lm.models.ministral3 import Model as Ministral3Model
+from mlx_lm.models.qwen3_5 import DecoderLayer as Qwen3_5DecoderLayer
+from mlx_lm.models.qwen3_5 import Model as Qwen3_5TextModel
+from mlx_lm.models.qwen3_5 import Qwen3_5TextModel as Qwen3_5TextModelInner
+from mlx_lm.models.qwen3_5 import SparseMoeBlock as Qwen3_5SparseMoeBlock
+from mlx_lm.models.qwen3_5_moe import Model as Qwen3_5MoeModel
 from mlx_lm.models.qwen3_moe import Model as Qwen3MoeModel
 from mlx_lm.models.qwen3_moe import Qwen3MoeDecoderLayer, Qwen3MoeSparseMoeBlock
 from mlx_lm.models.qwen3_next import Model as Qwen3NextModel
-from mlx_lm.models.qwen3_next import Qwen3NextDecoderLayer, Qwen3NextSparseMoeBlock
+from mlx_lm.models.qwen3_next import (
+    Qwen3NextDecoderLayer,
+    Qwen3NextGatedDeltaNet,
+    Qwen3NextSparseMoeBlock,
+)
 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
@@ -191,9 +201,10 @@ class PipelineLastLayer(CustomMlxLayer):
                 # CacheList (used by MLA models like DeepSeekV32, GLM MoE DSA)
                 # doesn't have .keys directly; access via first sub-cache.
                 _cache = cache[0] if hasattr(cache, "caches") else cache  # type: ignore
-                _cache.keys = mx.depends(_cache.keys, output)  # type: ignore
+                if hasattr(_cache, "keys"):  # pyright: ignore[reportAny]
+                    _cache.keys = mx.depends(_cache.keys, output)  # type: ignore
             mx.eval(output)
-            if cache is not None:
+            if cache is not None and hasattr(_cache, "keys"):  # type: ignore
                 mx.eval(_cache.keys)  # type: ignore
 
         if not self.is_prefill:
@@ -248,6 +259,32 @@ def get_layers(inner_model_instance: nn.Module) -> list[_LayerCallable]:
     return layers
 
 
+def _patch_qwen35_cache(
+    model: Qwen3_5TextModel,
+    fa_idx: int,
+    has_full_attn: bool,
+    ssm_idx: int,
+    has_linear: bool,
+) -> None:
+    # Hacks to make make_mask happy.
+    original = model.make_cache
+
+    def patched() -> list[ArraysCache | KVCache]:
+        cache: list[ArraysCache | KVCache] = original()
+        if not has_full_attn:
+            entry = cache[fa_idx]
+            orig_make_mask = entry.make_mask
+            entry.make_mask = lambda n, **_kw: orig_make_mask(n)  # type: ignore
+        if not has_linear:
+            orig_ssm_make_mask = cache[ssm_idx].make_mask
+            cache[ssm_idx].make_mask = (  # type: ignore
+                lambda n, **kw: orig_ssm_make_mask(n, **kw) if kw else None  # type: ignore
+            )
+        return cache
+
+    model.make_cache = patched
+
+
 def pipeline_auto_parallel(
     model: nn.Module,
     group: mx.distributed.Group,
@@ -318,6 +355,24 @@ def pipeline_auto_parallel(
         inner_model_instance._swa_idx = 0 if not sliding_layers else sliding_layers[0]
         inner_model_instance._full_idx = 0 if not full_layers else full_layers[0]
 
+    if isinstance(inner_model_instance, Qwen3_5TextModelInner):
+        full_attn_layers = [
+            i for i, layer in enumerate(layers) if not getattr(layer, "is_linear", True)
+        ]
+        linear_layers = [
+            i for i, layer in enumerate(layers) if getattr(layer, "is_linear", False)
+        ]
+        inner_model_instance.fa_idx = full_attn_layers[0] if full_attn_layers else 0
+        inner_model_instance.ssm_idx = linear_layers[0] if linear_layers else 0
+        if not full_attn_layers or not linear_layers:
+            _patch_qwen35_cache(
+                cast(Qwen3_5TextModel, model),
+                fa_idx=inner_model_instance.fa_idx,
+                has_full_attn=bool(full_attn_layers),
+                ssm_idx=inner_model_instance.ssm_idx,
+                has_linear=bool(linear_layers),
+            )
+
     _set_layers(model, layers)
 
     assert isinstance(layers, list), (
@@ -347,7 +402,8 @@ def patch_pipeline_model[T](model: T, group: mx.distributed.Group) -> T:
         if cache is not None:
             last = cache[-1]  # type: ignore
             dep_cache = last[0] if hasattr(last, "caches") else last  # type: ignore
-            dep_cache.keys = mx.depends(dep_cache.keys, logits)  # type: ignore
+            if hasattr(dep_cache, "keys") and dep_cache.keys is not None:  # type: ignore
+                dep_cache.keys = mx.depends(dep_cache.keys, logits)  # type: ignore
 
         return logits
 
@@ -470,7 +526,9 @@ def tensor_auto_parallel(
             all_to_sharded_linear_in_place,
             sharded_to_all_linear_in_place,
         )
-    elif isinstance(model, (Qwen3MoeModel, Qwen3NextModel)):
+    elif isinstance(
+        model, (Qwen3MoeModel, Qwen3NextModel, Qwen3_5TextModel, Qwen3_5MoeModel)
+    ):
         tensor_parallel_sharding_strategy = QwenShardingStrategy(
             group,
             all_to_sharded_linear,
@@ -865,7 +923,9 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
         on_timeout: TimeoutCallback | None,
         on_layer_loaded: LayerLoadedCallback | None,
     ) -> nn.Module:
-        model = cast(Qwen3MoeModel | Qwen3NextModel, model)
+        model = cast(
+            Qwen3MoeModel | Qwen3NextModel | Qwen3_5TextModel | Qwen3_5MoeModel, model
+        )
         total = len(model.layers)
         for i, layer in enumerate(model.layers):
             eval_with_timeout(layer.parameters(), timeout_seconds / total, on_timeout)
@@ -886,16 +946,39 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
                 layer.self_attn.n_heads //= self.N
                 layer.self_attn.n_kv_heads //= self.N
             else:
-                assert isinstance(layer, Qwen3NextDecoderLayer)
+                assert isinstance(layer, (Qwen3NextDecoderLayer, Qwen3_5DecoderLayer))
                 if hasattr(layer, "linear_attn"):
                     linear_attn = layer.linear_attn
 
-                    linear_attn.in_proj_qkvz = self.all_to_sharded_linear(
-                        linear_attn.in_proj_qkvz
-                    )
-                    linear_attn.in_proj_ba = self.all_to_sharded_linear(
-                        linear_attn.in_proj_ba
-                    )
+                    if isinstance(linear_attn, Qwen3NextGatedDeltaNet):
+                        # Qwen3-Next: combined projections
+                        linear_attn.in_proj_qkvz = self.all_to_sharded_linear(
+                            linear_attn.in_proj_qkvz
+                        )
+                        linear_attn.in_proj_ba = self.all_to_sharded_linear(
+                            linear_attn.in_proj_ba
+                        )
+                    else:
+                        # Qwen3.5: separate projections
+                        # in_proj_qkv has sections [q(key_dim), k(key_dim), v(value_dim)]
+                        # that must be split section-aware, not as a contiguous block
+                        key_dim = linear_attn.key_dim
+                        value_dim = linear_attn.value_dim
+                        linear_attn.in_proj_qkv = shard_linear(
+                            linear_attn.in_proj_qkv,
+                            "all-to-sharded",
+                            segments=[key_dim, key_dim + key_dim],
+                            group=self.group,
+                        )
+                        linear_attn.in_proj_z = self.all_to_sharded_linear(
+                            linear_attn.in_proj_z
+                        )
+                        linear_attn.in_proj_b = self.all_to_sharded_linear(
+                            linear_attn.in_proj_b
+                        )
+                        linear_attn.in_proj_a = self.all_to_sharded_linear(
+                            linear_attn.in_proj_a
+                        )
                     linear_attn.out_proj = self.sharded_to_all_linear(
                         linear_attn.out_proj
                     )
@@ -957,11 +1040,20 @@ class QwenShardingStrategy(TensorParallelShardingStrategy):
                     layer.self_attn.num_key_value_heads //= self.N
 
             # Shard the MoE.
-            if isinstance(layer.mlp, (Qwen3MoeSparseMoeBlock, Qwen3NextSparseMoeBlock)):
+            if isinstance(
+                layer.mlp,
+                (
+                    Qwen3MoeSparseMoeBlock,
+                    Qwen3NextSparseMoeBlock,
+                    Qwen3_5SparseMoeBlock,
+                ),
+            ):
                 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)
-                if isinstance(layer.mlp, Qwen3NextSparseMoeBlock):
+                if isinstance(
+                    layer.mlp, (Qwen3NextSparseMoeBlock, Qwen3_5SparseMoeBlock)
+                ):
                     self.all_to_sharded_linear_in_place(
                         layer.mlp.shared_expert.gate_proj
                     )
diff --git a/src/exo/worker/engines/mlx/utils_mlx.py b/src/exo/worker/engines/mlx/utils_mlx.py
index f44a2f91..596db0b1 100644
--- a/src/exo/worker/engines/mlx/utils_mlx.py
+++ b/src/exo/worker/engines/mlx/utils_mlx.py
@@ -318,6 +318,9 @@ def get_eos_token_ids_for_model(model_id: ModelId) -> list[int] | None:
         return [151336, 151329, 151338]
     elif "gpt-oss" in model_id_lower:
         return [200002, 200012]
+    elif "qwen3.5" in model_id_lower or "qwen-3.5" in model_id_lower:
+        # For Qwen3.5: 248046 (<|im_end|>), 248044 (<|endoftext|>)
+        return [248046, 248044]
     return None
 
 
diff --git a/tmp/gen_card.py b/tmp/gen_card.py
new file mode 100644
index 00000000..9270dbb6
--- /dev/null
+++ b/tmp/gen_card.py
@@ -0,0 +1,33 @@
+"""
+Generates inference model cards for EXO.
+Usage:
+    uv run tmp/gen_card.py mlx-community/my_cool_model-8bit [repo-id/model-id-2] [...]
+
+Model Cards require cleanup for family & quantization data
+"""
+
+import sys
+
+import anyio
+
+from exo.shared.models.model_cards import ModelCard, ModelId
+
+
+async def main():
+    if len(sys.argv) == 1:
+        print(f"USAGE: {sys.argv[0]} repo-id/model-id-1 [repo-id/model-id-2] [...]")
+        quit(1)
+    print("Remember! Model Cards require cleanup for family & quantization data")
+    for arg in sys.argv[1:]:
+        mid = ModelId(arg)
+        mc = await ModelCard.fetch_from_hf(mid)
+        await mc.save(
+            anyio.Path(__file__).parent.parent
+            / "resources"
+            / "inference_model_cards"
+            / (mid.normalize() + ".toml")
+        )
+
+
+if __name__ == "__main__":
+    anyio.run(main)
diff --git a/uv.lock b/uv.lock
index 3657bf70..1f307ee7 100644
--- a/uv.lock
+++ b/uv.lock
@@ -418,7 +418,7 @@ requires-dist = [
     { name = "mflux", specifier = "==0.15.5" },
     { name = "mlx", marker = "sys_platform == 'darwin'", git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks" },
     { name = "mlx", extras = ["cpu"], marker = "sys_platform == 'linux'", specifier = "==0.30.6" },
-    { name = "mlx-lm", specifier = "==0.30.7" },
+    { name = "mlx-lm", git = "https://github.com/ml-explore/mlx-lm?rev=834fac934c4e04de9b3d723e2b9287a2c60cfd4a" },
     { name = "msgspec", specifier = ">=0.19.0" },
     { name = "openai-harmony", specifier = ">=0.0.8" },
     { name = "pillow", specifier = ">=11.0,<12.0" },
@@ -1104,8 +1104,8 @@ wheels = [
 
 [[package]]
 name = "mlx-lm"
-version = "0.30.7"
-source = { registry = "https://pypi.org/simple" }
+version = "0.30.8"
+source = { git = "https://github.com/ml-explore/mlx-lm?rev=834fac934c4e04de9b3d723e2b9287a2c60cfd4a#834fac934c4e04de9b3d723e2b9287a2c60cfd4a" }
 dependencies = [
     { name = "jinja2", marker = "sys_platform == 'darwin' or sys_platform == 'linux'" },
     { name = "mlx", version = "0.30.7.dev20260225+257d5692", source = { git = "https://github.com/rltakashige/mlx-jaccl-fix-small-recv.git?branch=address-rdma-gpu-locks#257d5692fc7af6bba3b8afaeb63c549b7d1e43d5" }, marker = "sys_platform == 'darwin'" },
@@ -1115,10 +1115,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/66/0d/56542e2ae13ec6f542d3977d7cff89a205d4f6c5122e0ce23f33265f61c9/mlx_lm-0.30.7.tar.gz", hash = "sha256:e5f31ac58d9f2381f28e1ba639ff903e64f7cff1bdc245c0bc97f72264be329c", size = 275764, upload-time = "2026-02-12T18:41:11.86Z" }
-wheels = [
-    { url = "https://files.pythonhosted.org/packages/1e/17/a41c798a3d9cbdc47f39c6db5bba4c2cd199203ead26bf911cb03b644070/mlx_lm-0.30.7-py3-none-any.whl", hash = "sha256:17442a4bf01c4c2d3bca1e647712fe44f19890c3f1eadc8589d389e57b44b9bf", size = 386591, upload-time = "2026-02-12T18:41:10.236Z" },
-]
 
 [[package]]
 name = "more-itertools"

← 0e1b9501 fix: mini topology sidebar navigates home on click (#1616)  ·  back to Exo  ·  Refactor runner for implementing batching (#1632) 37296c82 →