diff --git a/mlx_lm/benchmark.py b/mlx_lm/benchmark.py index 560d5e5..9cd5057 100644 --- a/mlx_lm/benchmark.py +++ b/mlx_lm/benchmark.py @@ -75,11 +75,11 @@ def main(): if group.size() > 1: model, tokenizer, config = sharded_load( - args.model, pipeline_group, tensor_group, return_config=True + model_path, pipeline_group, tensor_group, return_config=True ) else: model, tokenizer, config = load( - args.model, return_config=True, tokenizer_config={"trust_remote_code": True} + model_path, return_config=True, tokenizer_config={"trust_remote_code": True} ) # Empty to avoid early stopping diff --git a/mlx_lm/models/Klear.py b/mlx_lm/models/Klear.py index 2bf5323..fb7a560 100644 --- a/mlx_lm/models/Klear.py +++ b/mlx_lm/models/Klear.py @@ -6,6 +6,7 @@ from typing import Any, List, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .switch_layers import SwitchGLU @@ -114,7 +115,7 @@ class KlearMLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class KlearSparseMoeBlock(nn.Module): diff --git a/mlx_lm/models/activations.py b/mlx_lm/models/activations.py new file mode 100644 index 0000000..7507899 --- /dev/null +++ b/mlx_lm/models/activations.py @@ -0,0 +1,11 @@ +# Copyright © 2023-2026 Apple Inc. + +from functools import partial + +import mlx.core as mx +import mlx.nn as nn + + +@partial(mx.compile, shapeless=True) +def swiglu(gate, x): + return nn.silu(gate) * x diff --git a/mlx_lm/models/afm7.py b/mlx_lm/models/afm7.py index d8b0dfe..8a87b45 100644 --- a/mlx_lm/models/afm7.py +++ b/mlx_lm/models/afm7.py @@ -9,6 +9,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import ConcatenateKVCache, KVCache from .rope_utils import initialize_rope @@ -262,11 +263,6 @@ class KVReuseAttention(nn.Module): return self.out_proj(output) -@partial(mx.compile, shapeless=True) -def _swiglu(g, x): - return nn.silu(g) * x - - class MLP(nn.Module): def __init__(self, args: ModelArgs): super().__init__() @@ -281,7 +277,7 @@ class MLP(nn.Module): def __call__(self, x) -> mx.array: g = self.gate_proj(x) x = self.up_proj(x) - return self.down_proj(_swiglu(g, x)) + return self.down_proj(swiglu(g, x)) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/afmoe.py b/mlx_lm/models/afmoe.py index a2898b0..e23dc23 100644 --- a/mlx_lm/models/afmoe.py +++ b/mlx_lm/models/afmoe.py @@ -7,6 +7,7 @@ from typing import Any, Dict, List, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .rope_utils import initialize_rope @@ -149,7 +150,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class MoERouter(nn.Module): diff --git a/mlx_lm/models/baichuan_m1.py b/mlx_lm/models/baichuan_m1.py index 6c82816..3221c02 100644 --- a/mlx_lm/models/baichuan_m1.py +++ b/mlx_lm/models/baichuan_m1.py @@ -6,6 +6,7 @@ from typing import Any, List, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import CacheList, KVCache, MambaCache, RotatingKVCache @@ -140,7 +141,7 @@ class MLP(nn.Module): ) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class DecoderLayer(nn.Module): diff --git a/mlx_lm/models/bailing_moe.py b/mlx_lm/models/bailing_moe.py index 18195a6..ff7d908 100644 --- a/mlx_lm/models/bailing_moe.py +++ b/mlx_lm/models/bailing_moe.py @@ -7,6 +7,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope from .switch_layers import SwitchGLU @@ -49,11 +50,6 @@ class ModelArgs(BaseModelArgs): moe_router_enable_shared_expert: bool = True -@partial(mx.compile, shapeless=True) -def swiglu(gate, up): - return nn.silu(gate) * up - - @partial(mx.compile, shapeless=True) def aggregate_expert_outputs(expert_outputs, scores): return ( diff --git a/mlx_lm/models/bailing_moe_linear.py b/mlx_lm/models/bailing_moe_linear.py index 8b2919f..360f99e 100644 --- a/mlx_lm/models/bailing_moe_linear.py +++ b/mlx_lm/models/bailing_moe_linear.py @@ -7,6 +7,7 @@ from typing import Any, Dict, Optional, Tuple, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -130,7 +131,7 @@ class MLP(nn.Module): ) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Attention(nn.Module): diff --git a/mlx_lm/models/cohere.py b/mlx_lm/models/cohere.py index c4cdd8e..c9b8888 100644 --- a/mlx_lm/models/cohere.py +++ b/mlx_lm/models/cohere.py @@ -6,6 +6,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -109,7 +110,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(hidden_dim, dim, bias=False) def __call__(self, x): - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/cohere2.py b/mlx_lm/models/cohere2.py index e288309..c93e238 100644 --- a/mlx_lm/models/cohere2.py +++ b/mlx_lm/models/cohere2.py @@ -6,6 +6,7 @@ from typing import Optional, Tuple import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache @@ -106,7 +107,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(hidden_dim, dim, bias=False) def __call__(self, x): - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/dbrx.py b/mlx_lm/models/dbrx.py index aff066f..cdd03ef 100644 --- a/mlx_lm/models/dbrx.py +++ b/mlx_lm/models/dbrx.py @@ -7,6 +7,7 @@ import mlx.core as mx import mlx.nn as nn import numpy as np +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -107,7 +108,7 @@ class MLP(nn.Module): self.w2 = nn.Linear(ffn_dim, d_model, bias=False) def __call__(self, x: mx.array) -> mx.array: - current_hidden_states = nn.silu(self.w1(x)) * self.v1(x) + current_hidden_states = swiglu(self.w1(x), self.v1(x)) current_hidden_states = self.w2(current_hidden_states) return current_hidden_states diff --git a/mlx_lm/models/deepseek.py b/mlx_lm/models/deepseek.py index c6fb159..2876f7b 100644 --- a/mlx_lm/models/deepseek.py +++ b/mlx_lm/models/deepseek.py @@ -4,6 +4,7 @@ from typing import Any, Dict, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .switch_layers import SwitchGLU @@ -120,7 +121,7 @@ class DeepseekMLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class MoEGate(nn.Module): diff --git a/mlx_lm/models/deepseek_v2.py b/mlx_lm/models/deepseek_v2.py index c56a748..a263d85 100644 --- a/mlx_lm/models/deepseek_v2.py +++ b/mlx_lm/models/deepseek_v2.py @@ -8,6 +8,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_inplace, shard_linear, sum_gradients +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .pipeline import PipelineMixin from .switch_layers import SwitchGLU @@ -260,7 +261,7 @@ class DeepseekV2MLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x): - down_proj = self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + down_proj = self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) return down_proj diff --git a/mlx_lm/models/deepseek_v3.py b/mlx_lm/models/deepseek_v3.py index 9f8360a..e9a95dd 100644 --- a/mlx_lm/models/deepseek_v3.py +++ b/mlx_lm/models/deepseek_v3.py @@ -9,6 +9,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_inplace, shard_linear, sum_gradients +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .pipeline import PipelineMixin from .rope_utils import initialize_rope @@ -174,7 +175,7 @@ class DeepseekV3MLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x): - down_proj = self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + down_proj = self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) return down_proj diff --git a/mlx_lm/models/deepseek_v32.py b/mlx_lm/models/deepseek_v32.py index a3069a3..fa3be06 100644 --- a/mlx_lm/models/deepseek_v32.py +++ b/mlx_lm/models/deepseek_v32.py @@ -8,6 +8,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_inplace, shard_linear, sum_gradients +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import CacheList, KVCache from .rope_utils import initialize_rope @@ -251,7 +252,7 @@ class DeepseekV32MLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x): - down_proj = self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + down_proj = self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) return down_proj diff --git a/mlx_lm/models/dots1.py b/mlx_lm/models/dots1.py index e5382ac..c568f39 100644 --- a/mlx_lm/models/dots1.py +++ b/mlx_lm/models/dots1.py @@ -7,6 +7,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope from .switch_layers import SwitchGLU @@ -180,7 +181,7 @@ class Dots1MLP(nn.Module): ) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Dots1MoE(nn.Module): diff --git a/mlx_lm/models/ernie4_5.py b/mlx_lm/models/ernie4_5.py index b667918..d6551f8 100644 --- a/mlx_lm/models/ernie4_5.py +++ b/mlx_lm/models/ernie4_5.py @@ -6,6 +6,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -87,7 +88,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=use_bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class DecoderLayer(nn.Module): diff --git a/mlx_lm/models/ernie4_5_moe.py b/mlx_lm/models/ernie4_5_moe.py index 0e28e9c..597e9c4 100644 --- a/mlx_lm/models/ernie4_5_moe.py +++ b/mlx_lm/models/ernie4_5_moe.py @@ -6,6 +6,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope from .switch_layers import SwitchGLU @@ -98,7 +99,7 @@ class Ernie4_5_MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=use_bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Ernie4_5_MoeMLP(nn.Module): diff --git a/mlx_lm/models/exaone.py b/mlx_lm/models/exaone.py index 93953f1..3cd223d 100644 --- a/mlx_lm/models/exaone.py +++ b/mlx_lm/models/exaone.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -91,7 +92,7 @@ class MLP(nn.Module): self.c_proj = nn.Linear(hidden_dim, dim, bias=args.mlp_bias) def __call__(self, x: mx.array) -> mx.array: - return self.c_proj(nn.silu(self.c_fc_0(x)) * self.c_fc_1(x)) + return self.c_proj(swiglu(self.c_fc_0(x), self.c_fc_1(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/exaone4.py b/mlx_lm/models/exaone4.py index 441e44a..c0e15f5 100644 --- a/mlx_lm/models/exaone4.py +++ b/mlx_lm/models/exaone4.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .rope_utils import initialize_rope @@ -102,7 +103,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/exaone_moe.py b/mlx_lm/models/exaone_moe.py index af36e45..a6da99b 100644 --- a/mlx_lm/models/exaone_moe.py +++ b/mlx_lm/models/exaone_moe.py @@ -7,6 +7,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_inplace, shard_linear, sum_gradients +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .rope_utils import initialize_rope @@ -119,7 +120,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) def __call__(self, x): - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class MoE(nn.Module): diff --git a/mlx_lm/models/falcon_h1.py b/mlx_lm/models/falcon_h1.py index 8e1b1b1..35a0b81 100644 --- a/mlx_lm/models/falcon_h1.py +++ b/mlx_lm/models/falcon_h1.py @@ -6,6 +6,7 @@ from typing import List, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -81,14 +82,14 @@ class FalconH1RMSNormGated(nn.Module): def __call__(self, hidden_states, gate=None): if not self.norm_before_gate and gate is not None: - hidden_states = hidden_states * nn.silu(gate) + hidden_states = swiglu(gate, hidden_states) hidden_states = mx.fast.rms_norm( hidden_states, self.weight, self.variance_epsilon ) if self.norm_before_gate and gate is not None: - hidden_states = hidden_states * nn.silu(gate) + hidden_states = swiglu(gate, hidden_states) return hidden_states @@ -329,7 +330,7 @@ class FalconH1Mixer(nn.Module): if self.mamba_rms_norm: y = self.norm(y, gate) else: - y = y * nn.silu(gate) + y = swiglu(gate, y) return self.out_proj(y) @@ -347,7 +348,7 @@ class FalconH1MLP(nn.Module): self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=args.mlp_bias) def __call__(self, x): - y = self.up_proj(x) * nn.silu(self.gate_proj(x)) + y = swiglu(self.gate_proj(x), self.up_proj(x)) y = self.down_proj(y) return y diff --git a/mlx_lm/models/glm.py b/mlx_lm/models/glm.py index d1eef71..ab3853d 100644 --- a/mlx_lm/models/glm.py +++ b/mlx_lm/models/glm.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -102,7 +103,7 @@ class GLMMLP(nn.Module): def __call__(self, x) -> mx.array: x = self.gate_up_proj(x) gate, x = mx.split(x, 2, axis=-1) - return self.down_proj(nn.silu(gate) * x) + return self.down_proj(swiglu(gate, x)) class GLMBlock(nn.Module): diff --git a/mlx_lm/models/glm4.py b/mlx_lm/models/glm4.py index c0cdc78..611175d 100644 --- a/mlx_lm/models/glm4.py +++ b/mlx_lm/models/glm4.py @@ -6,6 +6,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -38,7 +39,7 @@ class Glm4MLP(nn.Module): def __call__(self, x) -> mx.array: x = self.gate_up_proj(x) gate, up_states = mx.split(x, 2, axis=-1) - return self.down_proj(nn.silu(gate) * up_states) + return self.down_proj(swiglu(gate, up_states)) class Glm4Attention(nn.Module): diff --git a/mlx_lm/models/glm4_moe.py b/mlx_lm/models/glm4_moe.py index 3e0bb29..bcf7cc9 100644 --- a/mlx_lm/models/glm4_moe.py +++ b/mlx_lm/models/glm4_moe.py @@ -9,6 +9,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_inplace, shard_linear, sum_gradients +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .pipeline import PipelineMixin from .switch_layers import SwitchGLU @@ -123,7 +124,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x): - down_proj = self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + down_proj = self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) return down_proj diff --git a/mlx_lm/models/granite.py b/mlx_lm/models/granite.py index d17019c..6892445 100644 --- a/mlx_lm/models/granite.py +++ b/mlx_lm/models/granite.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -104,7 +105,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=mlp_bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/granitemoehybrid.py b/mlx_lm/models/granitemoehybrid.py index b2a8d2a..40dae99 100644 --- a/mlx_lm/models/granitemoehybrid.py +++ b/mlx_lm/models/granitemoehybrid.py @@ -6,6 +6,7 @@ from typing import Any, List, Optional, Tuple import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -75,7 +76,7 @@ class GraniteMoeHybridRMSNormGated(nn.Module): def __call__(self, hidden_states: mx.array, gate: mx.array = None) -> mx.array: if gate is not None: - hidden_states = hidden_states * nn.silu(gate) + hidden_states = swiglu(gate, hidden_states) return mx.fast.rms_norm(hidden_states, self.weight, self.eps) @@ -337,7 +338,7 @@ class GraniteMoeHybridSharedMLP(nn.Module): def __call__(self, x: mx.array) -> mx.array: gate, up = mx.split(self.input_linear(x), 2, axis=-1) - return self.output_linear(nn.silu(gate) * up) + return self.output_linear(swiglu(gate, up)) class GraniteMoeHybridMLP(nn.Module): @@ -352,7 +353,7 @@ class GraniteMoeHybridMLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=mlp_bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class GraniteMoeHybridLayer(nn.Module): diff --git a/mlx_lm/models/helium.py b/mlx_lm/models/helium.py index f7bc3bf..238f205 100644 --- a/mlx_lm/models/helium.py +++ b/mlx_lm/models/helium.py @@ -6,6 +6,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -92,7 +93,7 @@ class HeliumMLP(nn.Module): ) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class HeliumDecoderLayer(nn.Module): diff --git a/mlx_lm/models/hunyuan.py b/mlx_lm/models/hunyuan.py index 0eba6e5..0ee3a3a 100644 --- a/mlx_lm/models/hunyuan.py +++ b/mlx_lm/models/hunyuan.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Tuple, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .switch_layers import SwitchGLU @@ -148,7 +149,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Gate(nn.Module): diff --git a/mlx_lm/models/hunyuan_v1_dense.py b/mlx_lm/models/hunyuan_v1_dense.py index e2a86eb..d94edea 100644 --- a/mlx_lm/models/hunyuan_v1_dense.py +++ b/mlx_lm/models/hunyuan_v1_dense.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -144,7 +145,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/internlm2.py b/mlx_lm/models/internlm2.py index 9d3fdc9..8fd2301 100644 --- a/mlx_lm/models/internlm2.py +++ b/mlx_lm/models/internlm2.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -156,7 +157,7 @@ class MLP(nn.Module): self.w3 = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.w2(nn.silu(self.w1(x)) * self.w3(x)) + return self.w2(swiglu(self.w1(x), self.w3(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/internlm3.py b/mlx_lm/models/internlm3.py index f537bd7..766fe97 100644 --- a/mlx_lm/models/internlm3.py +++ b/mlx_lm/models/internlm3.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -154,7 +155,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/jamba.py b/mlx_lm/models/jamba.py index 1dd6630..f7515c0 100644 --- a/mlx_lm/models/jamba.py +++ b/mlx_lm/models/jamba.py @@ -7,6 +7,7 @@ from typing import Any, List, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -65,7 +66,7 @@ class JambaMLP(nn.Module): self.down_proj = nn.Linear(args.intermediate_size, args.hidden_size, bias=False) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class JambaAttention(nn.Module): @@ -205,7 +206,7 @@ class JambaMambaMixer(nn.Module): x = nn.silu(conv_out) A = -mx.exp(self.A_log) y, ssm_state = self.ssm_step(x, A, ssm_state) - z = self.out_proj(nn.silu(z) * y) + z = self.out_proj(swiglu(z, y)) return z, (conv_state, ssm_state) def __call__(self, x, cache): diff --git a/mlx_lm/models/kimi_linear.py b/mlx_lm/models/kimi_linear.py index fa73d12..bf61f5a 100644 --- a/mlx_lm/models/kimi_linear.py +++ b/mlx_lm/models/kimi_linear.py @@ -6,6 +6,7 @@ from typing import Any, Dict, List, Optional, Tuple import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -68,7 +69,7 @@ class KimiMLP(nn.Module): self.down_proj = nn.Linear(hidden, dim, bias=False) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) @mx.compile diff --git a/mlx_lm/models/lfm2.py b/mlx_lm/models/lfm2.py index a46ad38..649c7a8 100644 --- a/mlx_lm/models/lfm2.py +++ b/mlx_lm/models/lfm2.py @@ -5,6 +5,7 @@ from typing import Any, List, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -187,7 +188,7 @@ class MLP(nn.Module): self.w2 = nn.Linear(ff_dim, dim, bias=False) def __call__(self, x) -> mx.array: - return self.w2(nn.silu(self.w1(x)) * self.w3(x)) + return self.w2(swiglu(self.w1(x), self.w3(x))) class Lfm2DecoderLayer(nn.Module): diff --git a/mlx_lm/models/lfm2_moe.py b/mlx_lm/models/lfm2_moe.py index 64b42eb..3de939e 100644 --- a/mlx_lm/models/lfm2_moe.py +++ b/mlx_lm/models/lfm2_moe.py @@ -5,6 +5,7 @@ from typing import Any, List, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -179,7 +180,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Lfm2MoeSparseMoeBlock(nn.Module): diff --git a/mlx_lm/models/lille-130m.py b/mlx_lm/models/lille-130m.py index 839e0ed..b790572 100644 --- a/mlx_lm/models/lille-130m.py +++ b/mlx_lm/models/lille-130m.py @@ -6,6 +6,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -87,7 +88,7 @@ class Lille130mMLP(nn.Module): def __call__(self, x: mx.array) -> mx.array: h = self.norm(x) - return self.down_proj(nn.silu(self.gate_proj(h)) * self.up_proj(h)) + return self.down_proj(swiglu(self.gate_proj(h), self.up_proj(h))) class Lille130Block(nn.Module): diff --git a/mlx_lm/models/llama.py b/mlx_lm/models/llama.py index dcb8949..826bc9f 100644 --- a/mlx_lm/models/llama.py +++ b/mlx_lm/models/llama.py @@ -7,6 +7,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_linear +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .rope_utils import initialize_rope @@ -117,7 +118,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=mlp_bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/llama4.py b/mlx_lm/models/llama4.py index 36edfe4..e4e284d 100644 --- a/mlx_lm/models/llama4.py +++ b/mlx_lm/models/llama4.py @@ -6,6 +6,7 @@ from typing import Any, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import ChunkedKVCache, KVCache from .rope_utils import initialize_rope @@ -145,7 +146,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class MoE(nn.Module): diff --git a/mlx_lm/models/llama4_text.py b/mlx_lm/models/llama4_text.py index ba6b40c..6f45a85 100644 --- a/mlx_lm/models/llama4_text.py +++ b/mlx_lm/models/llama4_text.py @@ -6,6 +6,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -95,7 +96,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(intermediate_size, dim, bias=False) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/longcat_flash.py b/mlx_lm/models/longcat_flash.py index 85a26aa..3a649ad 100644 --- a/mlx_lm/models/longcat_flash.py +++ b/mlx_lm/models/longcat_flash.py @@ -5,6 +5,7 @@ from typing import Any, Dict, Optional, Tuple import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import CacheList, KVCache from .switch_layers import SwitchGLU @@ -168,7 +169,7 @@ class LongcatFlashMLP(nn.Module): self.down_proj = nn.Linear(hidden_size, args.hidden_size, bias=False) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class LongcatFlashTopkRouter(nn.Module): diff --git a/mlx_lm/models/mamba.py b/mlx_lm/models/mamba.py index 3fc4389..319a950 100644 --- a/mlx_lm/models/mamba.py +++ b/mlx_lm/models/mamba.py @@ -6,6 +6,7 @@ from dataclasses import dataclass import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs from .cache import MambaCache @@ -139,7 +140,7 @@ class MambaBlock(nn.Module): y_t, current_state = self.ssm_step(x[:, t], A, current_state) y.append(y_t) y = mx.stack(y, axis=1) - z = self.out_proj(nn.silu(z) * y) + z = self.out_proj(swiglu(z, y)) return z, (new_conv_cache, current_state) def __call__(self, x, cache): diff --git a/mlx_lm/models/mamba2.py b/mlx_lm/models/mamba2.py index a026f6e..87db6a6 100644 --- a/mlx_lm/models/mamba2.py +++ b/mlx_lm/models/mamba2.py @@ -7,6 +7,7 @@ from typing import Optional, Tuple, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_ssm_mask from .cache import MambaCache from .ssm import ssm_update @@ -48,7 +49,7 @@ class MambaRMSNormGated(nn.Module): def __call__(self, hidden_states: mx.array, gate: mx.array = None) -> mx.array: if gate is not None: - hidden_states = hidden_states * nn.silu(gate) + hidden_states = swiglu(gate, hidden_states) return mx.fast.rms_norm(hidden_states, self.weight, self.eps) diff --git a/mlx_lm/models/mimo.py b/mlx_lm/models/mimo.py index 78bb780..d87a2d9 100644 --- a/mlx_lm/models/mimo.py +++ b/mlx_lm/models/mimo.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -90,7 +91,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/mimo_v2_flash.py b/mlx_lm/models/mimo_v2_flash.py index b02e0e1..bec873e 100644 --- a/mlx_lm/models/mimo_v2_flash.py +++ b/mlx_lm/models/mimo_v2_flash.py @@ -8,6 +8,7 @@ from typing import Any, Dict, List, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .switch_layers import SwitchGLU @@ -139,7 +140,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x): - down_proj = self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + down_proj = self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) return down_proj diff --git a/mlx_lm/models/minicpm.py b/mlx_lm/models/minicpm.py index 9671e4e..660db90 100644 --- a/mlx_lm/models/minicpm.py +++ b/mlx_lm/models/minicpm.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -38,7 +39,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(args.intermediate_size, args.hidden_size, bias=False) def __call__(self, x): - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Attention(nn.Module): diff --git a/mlx_lm/models/minicpm3.py b/mlx_lm/models/minicpm3.py index 6165b04..977f00e 100644 --- a/mlx_lm/models/minicpm3.py +++ b/mlx_lm/models/minicpm3.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import SuScaledRoPE @@ -156,7 +157,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(args.intermediate_size, args.hidden_size, bias=False) def __call__(self, x): - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class DecoderLayer(nn.Module): diff --git a/mlx_lm/models/ministral3.py b/mlx_lm/models/ministral3.py index d61f63e..4c12658 100644 --- a/mlx_lm/models/ministral3.py +++ b/mlx_lm/models/ministral3.py @@ -7,6 +7,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_linear +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .pipeline import PipelineMixin @@ -121,7 +122,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/nemotron_h.py b/mlx_lm/models/nemotron_h.py index 5bc2419..c4a9f59 100644 --- a/mlx_lm/models/nemotron_h.py +++ b/mlx_lm/models/nemotron_h.py @@ -7,6 +7,7 @@ from typing import Any, List, Optional, Tuple import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -62,7 +63,7 @@ class MambaRMSNormGated(nn.Module): def __call__(self, x: mx.array, gate: mx.array = None) -> mx.array: if gate is not None: - x = x * nn.silu(gate) + x = swiglu(gate, x) x = mx.unflatten(x, axis=-1, shape=(-1, self.group_size)) x = mx.fast.rms_norm(x, weight=None, eps=self.eps) return self.weight * x.flatten(-2) diff --git a/mlx_lm/models/olmo.py b/mlx_lm/models/olmo.py index 1ab7818..bb7592d 100644 --- a/mlx_lm/models/olmo.py +++ b/mlx_lm/models/olmo.py @@ -7,6 +7,7 @@ from typing import Any, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask try: @@ -105,7 +106,7 @@ class TransformerBlock(nn.Module): x1, x2 = mx.split(self.ff_proj(self.ff_norm(h)), 2, axis=-1) - out = h + self.ff_out(nn.silu(x2) * x1) + out = h + self.ff_out(swiglu(x2, x1)) return out diff --git a/mlx_lm/models/olmo2.py b/mlx_lm/models/olmo2.py index f18f782..4591e00 100644 --- a/mlx_lm/models/olmo2.py +++ b/mlx_lm/models/olmo2.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -115,7 +116,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=mlp_bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/olmo3.py b/mlx_lm/models/olmo3.py index f0b829d..f87a818 100644 --- a/mlx_lm/models/olmo3.py +++ b/mlx_lm/models/olmo3.py @@ -6,6 +6,7 @@ from typing import Any, Dict, List, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .cache import KVCache, RotatingKVCache from .rope_utils import initialize_rope @@ -131,7 +132,7 @@ class Olmo3MLP(nn.Module): self.up_proj = nn.Linear(args.hidden_size, args.intermediate_size, bias=False) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Olmo3DecoderLayer(nn.Module): diff --git a/mlx_lm/models/openelm.py b/mlx_lm/models/openelm.py index 5b98528..1e5b895 100644 --- a/mlx_lm/models/openelm.py +++ b/mlx_lm/models/openelm.py @@ -6,6 +6,7 @@ from typing import Any, Dict, List, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -136,7 +137,7 @@ class MLP(nn.Module): def __call__(self, x) -> mx.array: x = self.proj_1(x) gate, x = mx.split(x, 2, axis=-1) - return self.proj_2(nn.silu(gate) * x) + return self.proj_2(swiglu(gate, x)) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/phi3.py b/mlx_lm/models/phi3.py index d82efa0..b3579ea 100644 --- a/mlx_lm/models/phi3.py +++ b/mlx_lm/models/phi3.py @@ -6,6 +6,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import SuScaledRoPE @@ -126,7 +127,7 @@ class MLP(nn.Module): def __call__(self, x) -> mx.array: x = self.gate_up_proj(x) gate, x = mx.split(x, 2, axis=-1) - return self.down_proj(nn.silu(gate) * x) + return self.down_proj(swiglu(gate, x)) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/plamo.py b/mlx_lm/models/plamo.py index df3fb4f..762a3c5 100644 --- a/mlx_lm/models/plamo.py +++ b/mlx_lm/models/plamo.py @@ -7,6 +7,7 @@ import mlx.core as mx import mlx.nn as nn import numpy as np +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -115,7 +116,7 @@ class MLP(nn.Module): self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) def __call__(self, x: mx.array) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) # type: ignore + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) # type: ignore class PlamoDecoderLayer(nn.Module): diff --git a/mlx_lm/models/plamo2.py b/mlx_lm/models/plamo2.py index ec16136..9e32aa1 100644 --- a/mlx_lm/models/plamo2.py +++ b/mlx_lm/models/plamo2.py @@ -9,6 +9,7 @@ import mlx.nn as nn from mlx_lm.models.base import BaseModelArgs, create_attention_mask, create_ssm_mask +from .activations import swiglu from .cache import KVCache, MambaCache from .ssm import ssm_update @@ -215,7 +216,7 @@ class Mamba(nn.Module): if cache: cache.advance(out.shape[1]) - out = out * nn.silu(z.flatten(-2)) + out = swiglu(z.flatten(-2), out) return self.out_proj(out) @@ -305,7 +306,7 @@ class MLP(nn.Module): def __call__(self, x: mx.array) -> mx.array: h = self.gate_up_proj(x) hs = mx.split(h, 2, axis=-1) - return self.down_proj(nn.silu(hs[0]) * hs[1]) + return self.down_proj(swiglu(hs[0], hs[1])) class PlamoDecoderLayer(nn.Module): diff --git a/mlx_lm/models/qwen.py b/mlx_lm/models/qwen.py index f3eef38..96e1a97 100644 --- a/mlx_lm/models/qwen.py +++ b/mlx_lm/models/qwen.py @@ -5,6 +5,7 @@ from dataclasses import dataclass import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -89,7 +90,7 @@ class MLP(nn.Module): def __call__(self, x): a1 = self.w1(x) a2 = self.w2(x) - return self.c_proj(a1 * nn.silu(a2)) + return self.c_proj(swiglu(a2, a1)) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/qwen2.py b/mlx_lm/models/qwen2.py index 7cfa966..45df85c 100644 --- a/mlx_lm/models/qwen2.py +++ b/mlx_lm/models/qwen2.py @@ -7,6 +7,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_linear +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -91,7 +92,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/qwen2_moe.py b/mlx_lm/models/qwen2_moe.py index d42275b..8d6022d 100644 --- a/mlx_lm/models/qwen2_moe.py +++ b/mlx_lm/models/qwen2_moe.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .switch_layers import SwitchGLU @@ -103,7 +104,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Qwen2MoeSparseMoeBlock(nn.Module): diff --git a/mlx_lm/models/qwen3.py b/mlx_lm/models/qwen3.py index 20586be..a59343b 100644 --- a/mlx_lm/models/qwen3.py +++ b/mlx_lm/models/qwen3.py @@ -7,6 +7,7 @@ import mlx.core as mx import mlx.nn as nn from mlx.nn.layers.distributed import shard_linear +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -96,7 +97,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/qwen3_moe.py b/mlx_lm/models/qwen3_moe.py index e07fed2..8557feb 100644 --- a/mlx_lm/models/qwen3_moe.py +++ b/mlx_lm/models/qwen3_moe.py @@ -6,6 +6,7 @@ from typing import Any, Dict, List, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .switch_layers import SwitchGLU @@ -103,7 +104,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Qwen3MoeSparseMoeBlock(nn.Module): diff --git a/mlx_lm/models/qwen3_next.py b/mlx_lm/models/qwen3_next.py index d3ca775..28b6c97 100644 --- a/mlx_lm/models/qwen3_next.py +++ b/mlx_lm/models/qwen3_next.py @@ -8,6 +8,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import ( BaseModelArgs, create_attention_mask, @@ -63,7 +64,7 @@ class Qwen3NextRMSNormGated(nn.Module): ) -> mx.array: x = mx.fast.rms_norm(hidden_states, self.weight, self.eps) if gate is not None: - x = x * nn.silu(gate) + x = swiglu(gate, x) return x @@ -155,7 +156,7 @@ class Qwen3NextMLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class Qwen3NextGatedDeltaNet(nn.Module): diff --git a/mlx_lm/models/seed_oss.py b/mlx_lm/models/seed_oss.py index e77ea1a..d23c9f3 100644 --- a/mlx_lm/models/seed_oss.py +++ b/mlx_lm/models/seed_oss.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional, Union import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -96,7 +97,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=bias) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class TransformerBlock(nn.Module): diff --git a/mlx_lm/models/stablelm.py b/mlx_lm/models/stablelm.py index 729f848..0f90f72 100644 --- a/mlx_lm/models/stablelm.py +++ b/mlx_lm/models/stablelm.py @@ -6,6 +6,7 @@ from dataclasses import dataclass import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention @@ -135,7 +136,7 @@ class MLP(nn.Module): self.up_proj = nn.Linear(dim, hidden_dim, bias=False) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class DecoderLayer(nn.Module): diff --git a/mlx_lm/models/switch_layers.py b/mlx_lm/models/switch_layers.py index 78e00ed..69f5623 100644 --- a/mlx_lm/models/switch_layers.py +++ b/mlx_lm/models/switch_layers.py @@ -6,6 +6,8 @@ from functools import partial import mlx.core as mx import mlx.nn as nn +from .activations import swiglu + def _gather_sort(x, indices): *_, M = indices.shape @@ -147,11 +149,6 @@ class SwitchLinear(nn.Module): return ql -@partial(mx.compile, shapeless=True) -def swiglu(x, gate): - return nn.silu(gate) * x - - class SwiGLU(nn.Module): def __init__(self): super().__init__() diff --git a/mlx_lm/models/youtu_llm.py b/mlx_lm/models/youtu_llm.py index a83d8a2..9d8c653 100644 --- a/mlx_lm/models/youtu_llm.py +++ b/mlx_lm/models/youtu_llm.py @@ -6,6 +6,7 @@ from typing import Any, Dict, Optional import mlx.core as mx import mlx.nn as nn +from .activations import swiglu from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_attention from .rope_utils import initialize_rope @@ -151,7 +152,7 @@ class YoutuLLMMLP(nn.Module): ) def __call__(self, x) -> mx.array: - return self.down_proj(nn.silu(self.gate_proj(x)) * self.up_proj(x)) + return self.down_proj(swiglu(self.gate_proj(x), self.up_proj(x))) class YoutuLLMDecoderLayer(nn.Module):