← back to Exo
Add support for Nemotron sharding (#1693)
82c54dd6d64dcb160b830e5cad8ff7999df2e11d · 2026-03-10 15:51:07 +0000 · rltakashige
### Automated Testing
tested logits match
Files touched
M .mlx_typings/mlx/nn/layers/convolution.pyiM .mlx_typings/mlx/nn/layers/linear.pyiM .mlx_typings/mlx/nn/layers/normalization.pyiA .mlx_typings/mlx_lm/models/nemotron_h.pyiM .mlx_typings/mlx_lm/models/qwen3_next.pyiM .mlx_typings/mlx_lm/models/switch_layers.pyiA resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-4bit.tomlA resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-8bit.tomlA resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-bf16.tomlA resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-4bit.tomlA resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-8bit.tomlA resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-bf16.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-4Bit.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-5Bit.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-6Bit.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-8Bit.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-BF16.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-MXFP4.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-4bits.tomlA resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-6bit.tomlM src/exo/shared/models/model_cards.pyM src/exo/worker/engines/image/models/flux/adapter.pyM src/exo/worker/engines/image/models/flux/kontext_adapter.pyM src/exo/worker/engines/mlx/auto_parallel.py
Diff
commit 82c54dd6d64dcb160b830e5cad8ff7999df2e11d
Author: rltakashige <rl.takashige@gmail.com>
Date: Tue Mar 10 15:51:07 2026 +0000
Add support for Nemotron sharding (#1693)
### Automated Testing
tested logits match
---
.mlx_typings/mlx/nn/layers/convolution.pyi | 1 +
.mlx_typings/mlx/nn/layers/linear.pyi | 4 +
.mlx_typings/mlx/nn/layers/normalization.pyi | 3 +
.mlx_typings/mlx_lm/models/nemotron_h.pyi | 154 +++++++++++++++
.mlx_typings/mlx_lm/models/qwen3_next.pyi | 4 +
.mlx_typings/mlx_lm/models/switch_layers.pyi | 3 +
...y--Llama-3.1-Nemotron-70B-Instruct-HF-4bit.toml | 12 ++
...y--Llama-3.1-Nemotron-70B-Instruct-HF-8bit.toml | 12 ++
...y--Llama-3.1-Nemotron-70B-Instruct-HF-bf16.toml | 12 ++
...nity--Llama-3.1-Nemotron-Nano-4B-v1.1-4bit.toml | 12 ++
...nity--Llama-3.1-Nemotron-Nano-4B-v1.1-8bit.toml | 12 ++
...nity--Llama-3.1-Nemotron-Nano-4B-v1.1-bf16.toml | 12 ++
...y--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-4Bit.toml | 12 ++
...y--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-5Bit.toml | 12 ++
...y--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-6Bit.toml | 12 ++
...y--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-8Bit.toml | 12 ++
...y--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-BF16.toml | 12 ++
...--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-MXFP4.toml | 12 ++
...nity--NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.toml | 12 ++
...ommunity--NVIDIA-Nemotron-Nano-9B-v2-4bits.toml | 12 ++
...community--NVIDIA-Nemotron-Nano-9B-v2-6bit.toml | 12 ++
src/exo/shared/models/model_cards.py | 1 +
.../worker/engines/image/models/flux/adapter.py | 2 +-
.../engines/image/models/flux/kontext_adapter.py | 2 +-
src/exo/worker/engines/mlx/auto_parallel.py | 209 +++++++++++++++++++--
25 files changed, 549 insertions(+), 14 deletions(-)
diff --git a/.mlx_typings/mlx/nn/layers/convolution.pyi b/.mlx_typings/mlx/nn/layers/convolution.pyi
index 3c3ddcb0..28b4ffd3 100644
--- a/.mlx_typings/mlx/nn/layers/convolution.pyi
+++ b/.mlx_typings/mlx/nn/layers/convolution.pyi
@@ -32,6 +32,7 @@ class Conv1d(Module):
"""
weight: mx.array
+ bias: mx.array | None
groups: int
def __init__(
self,
diff --git a/.mlx_typings/mlx/nn/layers/linear.pyi b/.mlx_typings/mlx/nn/layers/linear.pyi
index f9c91874..07e93a43 100644
--- a/.mlx_typings/mlx/nn/layers/linear.pyi
+++ b/.mlx_typings/mlx/nn/layers/linear.pyi
@@ -40,6 +40,10 @@ class Linear(Module):
bias (bool, optional): If set to ``False`` then the layer will
not use a bias. Default is ``True``.
"""
+
+ weight: mx.array
+ bias: mx.array | None
+
def __init__(self, input_dims: int, output_dims: int, bias: bool = ...) -> None: ...
def __call__(self, x: mx.array) -> mx.array: ...
def to_quantized(
diff --git a/.mlx_typings/mlx/nn/layers/normalization.pyi b/.mlx_typings/mlx/nn/layers/normalization.pyi
index 4116f860..216ccfff 100644
--- a/.mlx_typings/mlx/nn/layers/normalization.pyi
+++ b/.mlx_typings/mlx/nn/layers/normalization.pyi
@@ -88,6 +88,9 @@ class RMSNorm(Module):
dims (int): The feature dimension of the input to normalize over
eps (float): A small additive constant for numerical stability
"""
+
+ weight: mx.array
+
def __init__(self, dims: int, eps: float = ...) -> None: ...
def __call__(self, x) -> mx.array: ...
diff --git a/.mlx_typings/mlx_lm/models/nemotron_h.pyi b/.mlx_typings/mlx_lm/models/nemotron_h.pyi
new file mode 100644
index 00000000..6f38f3ad
--- /dev/null
+++ b/.mlx_typings/mlx_lm/models/nemotron_h.pyi
@@ -0,0 +1,154 @@
+from dataclasses import dataclass
+from typing import Any, List, Optional, Tuple
+
+import mlx.core as mx
+import mlx.nn as nn
+
+from .cache import ArraysCache, KVCache
+from .switch_layers import SwitchMLP
+
+@dataclass
+class ModelArgs:
+ model_type: str
+ vocab_size: int
+ hidden_size: int
+ intermediate_size: int
+ num_hidden_layers: int
+ max_position_embeddings: int
+ num_attention_heads: int
+ num_key_value_heads: int
+ attention_bias: bool
+ mamba_num_heads: int
+ mamba_head_dim: int
+ mamba_proj_bias: bool
+ ssm_state_size: int
+ conv_kernel: int
+ n_groups: int
+ mlp_bias: bool
+ layer_norm_epsilon: float
+ use_bias: bool
+ use_conv_bias: bool
+ hybrid_override_pattern: List[str]
+ head_dim: Optional[int]
+ moe_intermediate_size: Optional[int]
+ moe_shared_expert_intermediate_size: Optional[int]
+ n_group: Optional[int]
+ n_routed_experts: Optional[int]
+ n_shared_experts: Optional[int]
+ topk_group: Optional[int]
+ num_experts_per_tok: Optional[int]
+ norm_topk_prob: Optional[bool]
+ routed_scaling_factor: Optional[float]
+ time_step_limit: Optional[Tuple[float, float]]
+ time_step_min: Optional[float]
+ time_step_max: Optional[float]
+
+ @classmethod
+ def from_dict(cls, params: dict[str, Any]) -> ModelArgs: ...
+ def __post_init__(self) -> None: ...
+
+class NemotronHMamba2Mixer(nn.Module):
+ num_heads: int
+ hidden_size: int
+ ssm_state_size: int
+ conv_kernel_size: int
+ intermediate_size: int
+ n_groups: int
+ head_dim: int
+ conv_dim: int
+ conv1d: nn.Conv1d
+ in_proj: nn.Linear
+ dt_bias: mx.array
+ A_log: mx.array
+ D: mx.array
+ norm: nn.RMSNorm
+ heads_per_group: int
+ out_proj: nn.Linear
+
+ def __init__(self, args: ModelArgs) -> None: ...
+ def __call__(
+ self,
+ hidden_states: mx.array,
+ mask: Optional[mx.array],
+ cache: Optional[ArraysCache] = None,
+ ) -> mx.array: ...
+
+class NemotronHAttention(nn.Module):
+ hidden_size: int
+ num_heads: int
+ head_dim: int
+ num_key_value_heads: int
+ scale: float
+ q_proj: nn.Linear
+ k_proj: nn.Linear
+ v_proj: nn.Linear
+ o_proj: nn.Linear
+
+ def __init__(self, args: ModelArgs) -> None: ...
+ def __call__(
+ self,
+ x: mx.array,
+ mask: Optional[mx.array] = None,
+ cache: Optional[KVCache] = None,
+ ) -> mx.array: ...
+
+class NemotronHMLP(nn.Module):
+ up_proj: nn.Linear
+ down_proj: nn.Linear
+
+ def __init__(
+ self, args: ModelArgs, intermediate_size: Optional[int] = None
+ ) -> None: ...
+ def __call__(self, x: mx.array) -> mx.array: ...
+
+class NemotronHMoE(nn.Module):
+ num_experts_per_tok: int
+ switch_mlp: SwitchMLP
+ shared_experts: NemotronHMLP
+
+ def __init__(self, config: ModelArgs) -> None: ...
+ def __call__(self, x: mx.array) -> mx.array: ...
+
+class NemotronHBlock(nn.Module):
+ block_type: str
+ norm: nn.RMSNorm
+ mixer: NemotronHMamba2Mixer | NemotronHAttention | NemotronHMLP | NemotronHMoE
+
+ def __init__(self, args: ModelArgs, block_type: str) -> None: ...
+ def __call__(
+ self,
+ x: mx.array,
+ mask: Optional[mx.array] = None,
+ cache: Optional[Any] = None,
+ ) -> mx.array: ...
+
+class NemotronHModel(nn.Module):
+ embeddings: nn.Embedding
+ layers: list[NemotronHBlock]
+ norm_f: nn.RMSNorm
+ fa_idx: int
+ ssm_idx: int
+
+ def __init__(self, args: ModelArgs) -> None: ...
+ def __call__(
+ self,
+ inputs: mx.array,
+ cache: Optional[Any] = None,
+ ) -> mx.array: ...
+
+class Model(nn.Module):
+ args: ModelArgs
+ backbone: NemotronHModel
+ lm_head: nn.Linear
+ model_type: str
+
+ def __init__(self, args: ModelArgs) -> None: ...
+ def __call__(
+ self,
+ inputs: mx.array,
+ cache: Optional[Any] = None,
+ ) -> mx.array: ...
+ @property
+ def layers(self) -> list[NemotronHBlock]: ...
+ def make_cache(self) -> list[ArraysCache | KVCache]: ...
+ 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 10a508c0..649669cf 100644
--- a/.mlx_typings/mlx_lm/models/qwen3_next.pyi
+++ b/.mlx_typings/mlx_lm/models/qwen3_next.pyi
@@ -5,6 +5,7 @@ from typing import Any, Optional
import mlx.core as mx
import mlx.nn as nn
+from .cache import ArraysCache, KVCache
from .switch_layers import SwitchGLU
class Qwen3NextRMSNormGated(nn.Module):
@@ -99,6 +100,8 @@ class Qwen3NextModel(nn.Module):
embed_tokens: nn.Embedding
layers: list[Qwen3NextDecoderLayer]
norm: nn.RMSNorm
+ ssm_idx: int
+ fa_idx: int
def __init__(self, args: Any) -> None: ...
def __call__(
@@ -121,3 +124,4 @@ class Model(nn.Module):
def sanitize(self, weights: dict[str, Any]) -> dict[str, Any]: ...
@property
def layers(self) -> list[Qwen3NextDecoderLayer]: ...
+ def make_cache(self) -> list[ArraysCache | KVCache]: ...
diff --git a/.mlx_typings/mlx_lm/models/switch_layers.pyi b/.mlx_typings/mlx_lm/models/switch_layers.pyi
index d617c5b0..235d4168 100644
--- a/.mlx_typings/mlx_lm/models/switch_layers.pyi
+++ b/.mlx_typings/mlx_lm/models/switch_layers.pyi
@@ -73,6 +73,9 @@ class SwitchGLU(nn.Module):
def __call__(self, x, indices) -> mx.array: ...
class SwitchMLP(nn.Module):
+ fc1: SwitchLinear
+ fc2: SwitchLinear
+
def __init__(
self,
input_dims: int,
diff --git a/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-4bit.toml
new file mode 100644
index 00000000..2aacb04e
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-4bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Llama-3.1-Nemotron-70B-Instruct-HF-4bit"
+n_layers = 80
+hidden_size = 8192
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "llama"
+quantization = "4bit"
+base_model = "NVIDIA Llama-3.1-Nemotron-70B-Instruct"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 39688355840
diff --git a/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-8bit.toml
new file mode 100644
index 00000000..f7c34f8e
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Llama-3.1-Nemotron-70B-Instruct-HF-8bit"
+n_layers = 80
+hidden_size = 8192
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "llama"
+quantization = "8bit"
+base_model = "NVIDIA Llama-3.1-Nemotron-70B-Instruct"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 74964549632
diff --git a/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-bf16.toml b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-bf16.toml
new file mode 100644
index 00000000..f513934b
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-70B-Instruct-HF-bf16.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Llama-3.1-Nemotron-70B-Instruct-HF-bf16"
+n_layers = 80
+hidden_size = 8192
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "llama"
+quantization = "bf16"
+base_model = "NVIDIA Llama-3.1-Nemotron-70B-Instruct"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 141107412992
diff --git a/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-4bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-4bit.toml
new file mode 100644
index 00000000..12707721
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-4bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Llama-3.1-Nemotron-Nano-4B-v1.1-4bit"
+n_layers = 32
+hidden_size = 3072
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "llama"
+quantization = "4bit"
+base_model = "NVIDIA Llama-3.1-Nemotron-Nano-4B-v1.1"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 2538706944
diff --git a/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-8bit.toml b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-8bit.toml
new file mode 100644
index 00000000..b4baebc6
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-8bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Llama-3.1-Nemotron-Nano-4B-v1.1-8bit"
+n_layers = 32
+hidden_size = 3072
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "llama"
+quantization = "8bit"
+base_model = "NVIDIA Llama-3.1-Nemotron-Nano-4B-v1.1"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 4794980352
diff --git a/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-bf16.toml b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-bf16.toml
new file mode 100644
index 00000000..2258b856
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--Llama-3.1-Nemotron-Nano-4B-v1.1-bf16.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/Llama-3.1-Nemotron-Nano-4B-v1.1-bf16"
+n_layers = 32
+hidden_size = 3072
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "llama"
+quantization = "bf16"
+base_model = "NVIDIA Llama-3.1-Nemotron-Nano-4B-v1.1"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 9025492992
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-4Bit.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-4Bit.toml
new file mode 100644
index 00000000..f40798ea
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-4Bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-4Bit"
+n_layers = 52
+hidden_size = 2688
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "4bit"
+base_model = "NVIDIA Nemotron-3-Nano-30B-A3B"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 17775342336
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-5Bit.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-5Bit.toml
new file mode 100644
index 00000000..15761184
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-5Bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-5Bit"
+n_layers = 52
+hidden_size = 2688
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "5bit"
+base_model = "NVIDIA Nemotron-3-Nano-30B-A3B"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 21721476864
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-6Bit.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-6Bit.toml
new file mode 100644
index 00000000..238d4fe3
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-6Bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-6Bit"
+n_layers = 52
+hidden_size = 2688
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "6bit"
+base_model = "NVIDIA Nemotron-3-Nano-30B-A3B"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 25667611392
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-8Bit.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-8Bit.toml
new file mode 100644
index 00000000..7317de54
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-8Bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-8Bit"
+n_layers = 52
+hidden_size = 2688
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "8bit"
+base_model = "NVIDIA Nemotron-3-Nano-30B-A3B"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 33559880448
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-BF16.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-BF16.toml
new file mode 100644
index 00000000..fd8886af
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-BF16.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-BF16"
+n_layers = 52
+hidden_size = 2688
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "bf16"
+base_model = "NVIDIA Nemotron-3-Nano-30B-A3B"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 63155889408
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-MXFP4.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-MXFP4.toml
new file mode 100644
index 00000000..ec402a22
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-MXFP4.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-3-Nano-30B-A3B-MLX-MXFP4"
+n_layers = 52
+hidden_size = 2688
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "4bit"
+base_model = "NVIDIA Nemotron-3-Nano-30B-A3B"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 16788808704
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.toml
new file mode 100644
index 00000000..d9b86095
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-3-Nano-30B-A3B-NVFP4"
+n_layers = 52
+hidden_size = 2688
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "4bit"
+base_model = "NVIDIA Nemotron-3-Nano-30B-A3B"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 19323906944
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-4bits.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-4bits.toml
new file mode 100644
index 00000000..8084787e
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-4bits.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-Nano-9B-v2-4bits"
+n_layers = 56
+hidden_size = 4480
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "4bit"
+base_model = "NVIDIA Nemotron-Nano-9B-v2"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 5002791936
diff --git a/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-6bit.toml b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-6bit.toml
new file mode 100644
index 00000000..4f09bde3
--- /dev/null
+++ b/resources/inference_model_cards/mlx-community--NVIDIA-Nemotron-Nano-9B-v2-6bit.toml
@@ -0,0 +1,12 @@
+model_id = "mlx-community/NVIDIA-Nemotron-Nano-9B-v2-6bit"
+n_layers = 56
+hidden_size = 4480
+supports_tensor = true
+tasks = ["TextGeneration"]
+family = "nemotron"
+quantization = "6bit"
+base_model = "NVIDIA Nemotron-Nano-9B-v2"
+capabilities = ["text"]
+
+[storage_size]
+in_bytes = 7224298496
diff --git a/src/exo/shared/models/model_cards.py b/src/exo/shared/models/model_cards.py
index b6d94e74..f62363a0 100644
--- a/src/exo/shared/models/model_cards.py
+++ b/src/exo/shared/models/model_cards.py
@@ -196,6 +196,7 @@ class ConfigData(BaseModel):
["LlamaForCausalLM"],
["GptOssForCausalLM"],
["Step3p5ForCausalLM"],
+ ["NemotronHForCausalLM"],
]
@model_validator(mode="before")
diff --git a/src/exo/worker/engines/image/models/flux/adapter.py b/src/exo/worker/engines/image/models/flux/adapter.py
index 90b7bacd..e1d5d612 100644
--- a/src/exo/worker/engines/image/models/flux/adapter.py
+++ b/src/exo/worker/engines/image/models/flux/adapter.py
@@ -93,7 +93,7 @@ class FluxModelAdapter(ModelAdapter[Flux1, Transformer]):
@property
def hidden_dim(self) -> int:
- return self._transformer.x_embedder.weight.shape[0] # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
+ return self._transformer.x_embedder.weight.shape[0]
@property
def needs_cfg(self) -> bool:
diff --git a/src/exo/worker/engines/image/models/flux/kontext_adapter.py b/src/exo/worker/engines/image/models/flux/kontext_adapter.py
index 19d0be56..98dafad7 100644
--- a/src/exo/worker/engines/image/models/flux/kontext_adapter.py
+++ b/src/exo/worker/engines/image/models/flux/kontext_adapter.py
@@ -133,7 +133,7 @@ class FluxKontextModelAdapter(ModelAdapter[Flux1Kontext, Transformer]):
@property
def hidden_dim(self) -> int:
- return self._transformer.x_embedder.weight.shape[0] # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
+ return self._transformer.x_embedder.weight.shape[0]
@property
def needs_cfg(self) -> bool:
diff --git a/src/exo/worker/engines/mlx/auto_parallel.py b/src/exo/worker/engines/mlx/auto_parallel.py
index daf010ff..031fa4a1 100644
--- a/src/exo/worker/engines/mlx/auto_parallel.py
+++ b/src/exo/worker/engines/mlx/auto_parallel.py
@@ -4,7 +4,7 @@ from abc import ABC, abstractmethod
from collections.abc import Callable
from functools import partial
from inspect import signature
-from typing import TYPE_CHECKING, Any, Protocol, cast
+from typing import TYPE_CHECKING, Any, Literal, Protocol, cast
import mlx.core as mx
import mlx.nn as nn
@@ -32,6 +32,13 @@ 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.nemotron_h import Model as NemotronHModel
+from mlx_lm.models.nemotron_h import (
+ NemotronHAttention,
+ NemotronHMamba2Mixer,
+ NemotronHMoE,
+)
+from mlx_lm.models.nemotron_h import NemotronHModel as NemotronHInnerModel
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
@@ -45,6 +52,7 @@ from mlx_lm.models.qwen3_next import (
Qwen3NextGatedDeltaNet,
Qwen3NextSparseMoeBlock,
)
+from mlx_lm.models.qwen3_next import Qwen3NextModel as Qwen3NextInnerModel
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
@@ -243,7 +251,13 @@ def get_inner_model(model: nn.Module) -> nn.Module:
if isinstance(inner_inner, nn.Module):
return inner_inner
- raise ValueError("Model must either have a 'model' or 'transformer' attribute")
+ inner = getattr(model, "backbone", None)
+ if isinstance(inner, nn.Module):
+ return inner
+
+ raise ValueError(
+ "Model must either have a 'model', 'transformer', or 'backbone' attribute"
+ )
def get_layers(inner_model_instance: nn.Module) -> list[_LayerCallable]:
@@ -259,8 +273,8 @@ def get_layers(inner_model_instance: nn.Module) -> list[_LayerCallable]:
return layers
-def _patch_qwen35_cache(
- model: Qwen3_5TextModel,
+def _patch_hybrid_cache(
+ model: Qwen3_5TextModel | Qwen3NextModel | NemotronHModel,
fa_idx: int,
has_full_attn: bool,
ssm_idx: int,
@@ -270,16 +284,20 @@ def _patch_qwen35_cache(
original = model.make_cache
def patched() -> list[ArraysCache | KVCache]:
- cache: list[ArraysCache | KVCache] = original()
+ cache = 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
- )
+
+ def _ssm_mask(
+ n: int, **kw: bool | int | None
+ ) -> mx.array | Literal["causal"] | None:
+ return orig_ssm_make_mask(n, **kw) if kw else None
+
+ cache[ssm_idx].make_mask = _ssm_mask # type: ignore
return cache
model.make_cache = patched
@@ -355,7 +373,7 @@ 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):
+ if isinstance(inner_model_instance, (Qwen3_5TextModelInner, Qwen3NextInnerModel)):
full_attn_layers = [
i for i, layer in enumerate(layers) if not getattr(layer, "is_linear", True)
]
@@ -365,14 +383,44 @@ def pipeline_auto_parallel(
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),
+ _patch_hybrid_cache(
+ cast(Qwen3_5TextModel | Qwen3NextModel, 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),
)
+ if isinstance(inner_model_instance, NemotronHInnerModel):
+ # NemotronH uses block_type: "M" (Mamba/SSM), "*" (Attention), "E" (MoE), "-" (MLP)
+ # Only "M" and "*" blocks have cache entries.
+ # Recompute fa_idx and ssm_idx as cache-array indices for the shard's layers.
+ cache_idx = 0
+ fa_idx: int | None = None
+ ssm_idx: int | None = None
+ for layer in layers:
+ block_type = getattr(layer, "block_type", None)
+ if block_type == "*":
+ if fa_idx is None:
+ fa_idx = cache_idx
+ cache_idx += 1
+ elif block_type == "M":
+ if ssm_idx is None:
+ ssm_idx = cache_idx
+ cache_idx += 1
+ has_attn = fa_idx is not None
+ has_mamba = ssm_idx is not None
+ inner_model_instance.fa_idx = fa_idx if fa_idx is not None else 0
+ inner_model_instance.ssm_idx = ssm_idx if ssm_idx is not None else 0
+ if not has_attn or not has_mamba:
+ _patch_hybrid_cache(
+ cast(NemotronHModel, model),
+ fa_idx=inner_model_instance.fa_idx,
+ has_full_attn=has_attn,
+ ssm_idx=inner_model_instance.ssm_idx,
+ has_linear=has_mamba,
+ )
+
_set_layers(model, layers)
assert isinstance(layers, list), (
@@ -431,7 +479,8 @@ def patch_tensor_model[T](model: T) -> T:
if cache is not None and len(cache) > 0: # pyright: ignore[reportAny]
last = cache[-1] # pyright: ignore[reportAny]
dep_cache = last[0] if hasattr(last, "caches") else last # pyright: ignore[reportAny]
- dep_cache.keys = mx.depends(dep_cache.keys, logits) # pyright: ignore[reportAny,reportUnknownMemberType]
+ if hasattr(dep_cache, "keys"): # type: ignore
+ dep_cache.keys = mx.depends(dep_cache.keys, logits) # pyright: ignore[reportAny,reportUnknownMemberType]
return logits
@@ -552,6 +601,14 @@ def tensor_auto_parallel(
all_to_sharded_linear_in_place,
sharded_to_all_linear_in_place,
)
+ elif isinstance(model, NemotronHModel):
+ tensor_parallel_sharding_strategy = NemotronHShardingStrategy(
+ group,
+ all_to_sharded_linear,
+ sharded_to_all_linear,
+ all_to_sharded_linear_in_place,
+ sharded_to_all_linear_in_place,
+ )
else:
raise ValueError(f"Unsupported model type: {type(model)}")
@@ -1210,3 +1267,131 @@ class Step35ShardingStrategy(TensorParallelShardingStrategy):
if on_layer_loaded is not None:
on_layer_loaded(i, total)
return model
+
+
+class NemotronHShardingStrategy(TensorParallelShardingStrategy):
+ def shard_model(
+ self,
+ model: nn.Module,
+ timeout_seconds: float,
+ on_timeout: TimeoutCallback | None,
+ on_layer_loaded: LayerLoadedCallback | None,
+ ) -> nn.Module:
+ model = cast(NemotronHModel, model)
+ rank = self.group.rank()
+ total = len(model.layers)
+ for i, layer in enumerate(model.layers):
+ eval_with_timeout(layer.parameters(), timeout_seconds / total, on_timeout)
+
+ mixer = layer.mixer
+
+ if isinstance(mixer, NemotronHAttention):
+ mixer.q_proj = self.all_to_sharded_linear(mixer.q_proj)
+ mixer.k_proj = self.all_to_sharded_linear(mixer.k_proj)
+ mixer.v_proj = self.all_to_sharded_linear(mixer.v_proj)
+ mixer.o_proj = self.sharded_to_all_linear(mixer.o_proj)
+ mixer.num_heads //= self.N
+ mixer.num_key_value_heads //= self.N
+
+ elif isinstance(mixer, NemotronHMamba2Mixer):
+ self._shard_mamba2_mixer(mixer, rank)
+
+ elif isinstance(mixer, NemotronHMoE):
+ # Shard routed experts (SwitchMLP uses fc1/fc2)
+ self.all_to_sharded_linear_in_place(mixer.switch_mlp.fc1)
+ self.sharded_to_all_linear_in_place(mixer.switch_mlp.fc2)
+ # Shard shared expert in-place (no all-reduce — ShardedMoE handles that)
+ if hasattr(mixer, "shared_experts"):
+ self.all_to_sharded_linear_in_place(mixer.shared_experts.up_proj)
+ self.sharded_to_all_linear_in_place(mixer.shared_experts.down_proj)
+ mixer = ShardedMoE(mixer) # pyright: ignore[reportArgumentType]
+ mixer.sharding_group = self.group
+ layer.mixer = mixer # pyright: ignore[reportAttributeAccessIssue]
+
+ mx.eval(layer)
+ if on_layer_loaded is not None:
+ on_layer_loaded(i, total)
+ return model
+
+ def _shard_mamba2_mixer(self, mixer: NemotronHMamba2Mixer, rank: int) -> None:
+ """Shard the Mamba2 mixer along the head dimension."""
+ world_size = self.N
+ num_heads = mixer.num_heads
+ head_dim = mixer.head_dim
+ n_groups = mixer.n_groups
+ ssm_state_size = mixer.ssm_state_size
+ intermediate_size = mixer.intermediate_size # = num_heads * head_dim
+
+ # Per-rank sizes
+ heads_per_rank = num_heads // world_size
+ groups_per_rank = n_groups // world_size
+ is_per_rank = heads_per_rank * head_dim
+ bc_per_rank = groups_per_rank * ssm_state_size
+
+ # === in_proj: output layout is [gate:IS | conv_ssm:IS | B:NG*SS | C:NG*SS | dt:NH] ===
+ gate_start = 0
+ conv_ssm_start = intermediate_size
+ b_start = 2 * intermediate_size
+ c_start = b_start + n_groups * ssm_state_size
+ dt_start = c_start + n_groups * ssm_state_size
+
+ # Build index tensor for this rank's slice of each section
+ gate_idx = mx.arange(
+ gate_start + rank * is_per_rank, gate_start + (rank + 1) * is_per_rank
+ )
+ conv_ssm_idx = mx.arange(
+ conv_ssm_start + rank * is_per_rank,
+ conv_ssm_start + (rank + 1) * is_per_rank,
+ )
+ b_idx = mx.arange(
+ b_start + rank * bc_per_rank, b_start + (rank + 1) * bc_per_rank
+ )
+ c_idx = mx.arange(
+ c_start + rank * bc_per_rank, c_start + (rank + 1) * bc_per_rank
+ )
+ dt_idx = mx.arange(
+ dt_start + rank * heads_per_rank, dt_start + (rank + 1) * heads_per_rank
+ )
+
+ indices = mx.concatenate([gate_idx, conv_ssm_idx, b_idx, c_idx, dt_idx])
+ mixer.in_proj.weight = mixer.in_proj.weight[indices]
+
+ # === out_proj: input is intermediate_size (sharded) → hidden_size (reduce) ===
+ mixer.out_proj = self.sharded_to_all_linear(mixer.out_proj)
+
+ # === conv1d: depthwise conv on conv_dim channels ===
+ # conv_dim layout: [ssm_hidden:IS | B:NG*SS | C:NG*SS]
+ conv_ssm_idx_local = mx.arange(rank * is_per_rank, (rank + 1) * is_per_rank)
+ conv_b_idx = mx.arange(
+ intermediate_size + rank * bc_per_rank,
+ intermediate_size + (rank + 1) * bc_per_rank,
+ )
+ conv_c_idx = mx.arange(
+ intermediate_size + n_groups * ssm_state_size + rank * bc_per_rank,
+ intermediate_size + n_groups * ssm_state_size + (rank + 1) * bc_per_rank,
+ )
+ conv_indices = mx.concatenate([conv_ssm_idx_local, conv_b_idx, conv_c_idx])
+ mixer.conv1d.weight = mixer.conv1d.weight[conv_indices]
+ new_conv_dim = is_per_rank + 2 * bc_per_rank
+ mixer.conv1d.groups = new_conv_dim
+ if mixer.conv1d.bias is not None:
+ mixer.conv1d.bias = mixer.conv1d.bias[conv_indices]
+
+ # === Per-head parameters ===
+ h_start = rank * heads_per_rank
+ h_end = h_start + heads_per_rank
+ mixer.dt_bias = mixer.dt_bias[h_start:h_end]
+ mixer.A_log = mixer.A_log[h_start:h_end]
+ mixer.D = mixer.D[h_start:h_end]
+
+ # === Norm: weight is intermediate_size ===
+ mixer.norm.weight = mixer.norm.weight[
+ rank * is_per_rank : (rank + 1) * is_per_rank
+ ]
+
+ # === Update dimensions ===
+ mixer.num_heads = heads_per_rank
+ mixer.n_groups = groups_per_rank
+ mixer.intermediate_size = is_per_rank
+ mixer.conv_dim = new_conv_dim
+ mixer.heads_per_group = heads_per_rank // groups_per_rank
← 3536161f Fix stale state.runners (#1684)
·
back to Exo
·
Include power usage in bench responses (#1692) f36fd56c →