# Copyright 2026 # # Local MLX-LM compatibility loader for upstage/Solar-Open2-250B. # # Solar Open 2 uses a hybrid stack: GQA/full attention every fourth layer and # Kimi-style gated delta attention in the other layers, with a GLM/Solar MoE. from dataclasses import dataclass, field from typing import Any, Dict, List, Optional, Tuple import mlx.core as mx import mlx.nn as nn from mlx_lm.models.base import ( BaseModelArgs, create_attention_mask, create_ssm_mask, scaled_dot_product_attention, ) from mlx_lm.models.cache import ArraysCache, KVCache from mlx_lm.models.gated_delta import gated_delta_kernel, gated_delta_ops from mlx_lm.models.glm4_moe import MLP, MoE from mlx_lm.models.kimi_linear import KimiDeltaAttention from mlx_lm.models.pipeline import PipelineMixin @dataclass class ModelArgs(BaseModelArgs): model_type: str vocab_size: int hidden_size: int intermediate_size: int moe_intermediate_size: int num_hidden_layers: int num_attention_heads: int num_key_value_heads: int head_dim: int n_shared_experts: int n_routed_experts: int routed_scaling_factor: float num_experts_per_tok: int first_k_dense_replace: int norm_topk_prob: bool max_position_embeddings: int rms_norm_eps: float rope_theta: float = 10000.0 tie_word_embeddings: bool = False partial_rotary_factor: float = 1.0 linear_attn_config: Dict[str, Any] = field(default_factory=dict) gqa_layers: List[int] = field(default_factory=list) gqa_interval: int = 3 use_gqa_gate: bool = True use_gqa_gate_bias: bool = False use_rope: bool = False attention_bias: bool = False use_qk_norm: bool = False kda_use_full_proj: bool = False kda_gate_lower_bound: Optional[float] = -5.0 kda_allow_neg_eigval: bool = True n_group: int = 1 topk_group: int = 1 scoring_func: str = "sigmoid" topk_method: str = "noaux_tc" @mx.compile def _solar_kda_decay(A_log, a, dt_bias, lower_bound: Optional[float]): num_heads = A_log.size head_dim = dt_bias.size // num_heads A = mx.reshape(A_log.astype(mx.float32), (num_heads, 1)) dt = mx.reshape(dt_bias.astype(mx.float32), (num_heads, head_dim)) log_decay = -mx.exp(A) * nn.softplus(a.astype(mx.float32) + dt) if lower_bound is not None: log_decay = mx.maximum(log_decay, mx.array(lower_bound, dtype=log_decay.dtype)) return mx.exp(log_decay) def _solar_gated_delta_update( q: mx.array, k: mx.array, v: mx.array, a: mx.array, b: mx.array, A_log: mx.array, dt_bias: mx.array, state: Optional[mx.array] = None, mask: Optional[mx.array] = None, use_kernel: bool = True, lower_bound: Optional[float] = -5.0, allow_neg_eigval: bool = True, ) -> Tuple[mx.array, mx.array]: beta = mx.sigmoid(b) if allow_neg_eigval: beta = beta * 2.0 g = _solar_kda_decay(A_log, a, dt_bias, lower_bound) if state is None: B, _, Hk, Dk = q.shape Hv, Dv = v.shape[-2:] state = mx.zeros((B, Hv, Dv, Dk), dtype=mx.float32) if not use_kernel or mx.default_device() != mx.gpu or not mx.metal.is_available(): return gated_delta_ops(q, k, v, g, beta, state, mask) return gated_delta_kernel(q, k, v, g, beta, state, mask) class SolarOpen2Attention(nn.Module): def __init__(self, args: ModelArgs): super().__init__() dim = args.hidden_size self.n_heads = args.num_attention_heads self.n_kv_heads = args.num_key_value_heads self.head_dim = args.head_dim self.scale = self.head_dim**-0.5 self.use_gqa_gate = args.use_gqa_gate self.q_proj = nn.Linear(dim, self.n_heads * self.head_dim, bias=args.attention_bias) self.k_proj = nn.Linear(dim, self.n_kv_heads * self.head_dim, bias=args.attention_bias) self.v_proj = nn.Linear(dim, self.n_kv_heads * self.head_dim, bias=args.attention_bias) self.o_proj = nn.Linear(self.n_heads * self.head_dim, dim, bias=False) self.use_qk_norm = args.use_qk_norm if self.use_qk_norm: self.q_norm = nn.RMSNorm(self.head_dim, eps=args.rms_norm_eps) self.k_norm = nn.RMSNorm(self.head_dim, eps=args.rms_norm_eps) if self.use_gqa_gate: self.g_proj = nn.Linear(dim, self.n_heads * self.head_dim, bias=args.use_gqa_gate_bias) def __call__( self, x: mx.array, mask: Optional[mx.array] = None, cache: Optional[Any] = None, ) -> mx.array: B, L, _ = x.shape queries = self.q_proj(x).reshape(B, L, self.n_heads, self.head_dim).transpose(0, 2, 1, 3) keys = self.k_proj(x).reshape(B, L, self.n_kv_heads, self.head_dim).transpose(0, 2, 1, 3) values = self.v_proj(x).reshape(B, L, self.n_kv_heads, self.head_dim).transpose(0, 2, 1, 3) if self.use_qk_norm: queries = self.q_norm(queries) keys = self.k_norm(keys) if cache is not None: keys, values = cache.update_and_fetch(keys, values) output = scaled_dot_product_attention( queries, keys, values, cache=cache, scale=self.scale, mask=mask ) output = output.transpose(0, 2, 1, 3).reshape(B, L, -1) if self.use_gqa_gate: output = output * mx.sigmoid(self.g_proj(x)) return self.o_proj(output) class SolarOpen2LinearAttention(KimiDeltaAttention): def __init__(self, args: ModelArgs, layer_idx: int): if args.kda_use_full_proj: raise NotImplementedError("Solar Open2 full KDA projections are not supported by this MLX loader.") super().__init__(args, layer_idx) self.kda_gate_lower_bound = args.kda_gate_lower_bound self.kda_allow_neg_eigval = args.kda_allow_neg_eigval def __call__( self, x: mx.array, mask: Optional[mx.array] = None, cache: Optional[Any] = None, ) -> mx.array: B, T, _ = x.shape dtype = x.dtype if cache is not None: q_state, k_state, v_state, ssm_state = cache lengths = cache.lengths else: q_state = None k_state = None v_state = None ssm_state = None lengths = None if q_state is None: s = mx.zeros((B, self.conv_kernel - 1, self.projection_dim), dtype=dtype) q_state = s k_state = s v_state = s q_conv, q_state = self.q_conv(self.q_proj(x), q_state, mask, lengths) k_conv, k_state = self.k_conv(self.k_proj(x), k_state, mask, lengths) v_conv, v_state = self.v_conv(self.v_proj(x), v_state, mask, lengths) if cache is not None: cache[0] = q_state cache[1] = k_state cache[2] = v_state q = q_conv.reshape(B, T, self.num_heads, self.head_dim) k = k_conv.reshape(B, T, self.num_heads, self.head_dim) v = v_conv.reshape(B, T, self.num_heads, self.head_dim) inv_scale = self.scale q = (inv_scale**2) * mx.fast.rms_norm(q, None, 1e-6) k = inv_scale * mx.fast.rms_norm(k, None, 1e-6) a_logits = self.f_b_proj(self.f_a_proj(x)).reshape(B, T, self.num_heads, self.head_dim) b_logits = self.b_proj(x).reshape(B, T, self.num_heads) out, ssm_state = _solar_gated_delta_update( q, k, v, a_logits, b_logits, self.A_log.reshape(self.num_heads, 1), self.dt_bias.reshape(self.num_heads, self.head_dim), state=ssm_state, mask=mask, use_kernel=not self.training, lower_bound=self.kda_gate_lower_bound, allow_neg_eigval=self.kda_allow_neg_eigval, ) if cache is not None: cache[3] = ssm_state cache.advance(T) gate = self.g_b_proj(self.g_a_proj(x)).reshape(B, T, self.num_heads, self.head_dim) out = (self.o_norm(out.reshape(B, T, self.num_heads, self.head_dim)) * mx.sigmoid(gate)).reshape(B, T, -1) return self.o_proj(out) class SolarOpen2DecoderLayer(nn.Module): def __init__(self, args: ModelArgs, layer_idx: int): super().__init__() gqa_layers = set(args.gqa_layers or list(range(0, args.num_hidden_layers, args.gqa_interval + 1))) self.is_linear = layer_idx not in gqa_layers self.self_attn = ( SolarOpen2LinearAttention(args, layer_idx) if self.is_linear else SolarOpen2Attention(args) ) self.mlp = ( MoE(args) if args.n_routed_experts is not None and layer_idx >= args.first_k_dense_replace else MLP(args) ) self.input_layernorm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps) self.post_attention_layernorm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps) def __call__( self, x: mx.array, mask: Optional[mx.array] = None, cache: Optional[Any] = None, ) -> mx.array: h = x + self.self_attn(self.input_layernorm(x), mask=mask, cache=cache) return h + self.mlp(self.post_attention_layernorm(h)) class SolarOpen2Model(PipelineMixin, nn.Module): def __init__(self, args: ModelArgs): super().__init__() self.embed_tokens = nn.Embedding(args.vocab_size, args.hidden_size) self.layers = [SolarOpen2DecoderLayer(args, i) for i in range(args.num_hidden_layers)] self.norm = nn.RMSNorm(args.hidden_size, eps=args.rms_norm_eps) self.linear_idx = next((i for i, layer in enumerate(self.layers) if layer.is_linear), 0) self.attn_idx = next((i for i, layer in enumerate(self.layers) if not layer.is_linear), 0) def __call__( self, inputs: mx.array, cache: Optional[List[Any]] = None, ) -> mx.array: h = self.embed_tokens(inputs) if cache is None: cache = [None] * len(self.layers) ssm_mask = create_ssm_mask(h, cache[self.linear_idx]) attn_mask = create_attention_mask(h, cache[self.attn_idx], return_array=True) for layer, layer_cache in zip(self.layers, cache): mask = ssm_mask if layer.is_linear else attn_mask h = layer(h, mask=mask, cache=layer_cache) return self.norm(h) class Model(nn.Module): def __init__(self, args: ModelArgs): super().__init__() self.args = args self.model_type = args.model_type self.model = SolarOpen2Model(args) if args.tie_word_embeddings: self.lm_head = None else: self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False) def __call__( self, inputs: mx.array, cache: Optional[List[Any]] = None, ) -> mx.array: out = self.model(inputs, cache) if self.lm_head is None: return self.model.embed_tokens.as_linear(out) return self.lm_head(out) @property def layers(self): return self.model.layers def make_cache(self): caches: List[Any] = [] for layer in self.layers: caches.append(ArraysCache(size=4) if layer.is_linear else KVCache()) return caches def sanitize(self, weights: Dict[str, mx.array]) -> Dict[str, mx.array]: # Stack per-expert HF tensors into MLX SwitchGLU tensors. for layer_idx in range(self.args.num_hidden_layers): prefix = f"model.layers.{layer_idx}" for dst, src in (("gate_proj", "gate_proj"), ("down_proj", "down_proj"), ("up_proj", "up_proj")): for suffix in ("weight", "scales", "biases"): first = f"{prefix}.mlp.experts.0.{src}.{suffix}" if first in weights: weights[f"{prefix}.mlp.switch_mlp.{dst}.{suffix}"] = mx.stack( [ weights.pop(f"{prefix}.mlp.experts.{expert}.{src}.{suffix}") for expert in range(self.args.n_routed_experts) ] ) layer = self.layers[layer_idx] if layer.is_linear: attn_prefix = f"{prefix}.self_attn" for src_name, dst_name in ( ("q_conv1d", "q_conv"), ("k_conv1d", "k_conv"), ("v_conv1d", "v_conv"), ): src_key = f"{attn_prefix}.{src_name}.weight" if src_key in weights: w = weights.pop(src_key) if w.ndim == 3: w = w.moveaxis(2, 1) weights[f"{attn_prefix}.{dst_name}.conv.weight"] = w dt_key = f"{attn_prefix}.dt_bias" if dt_key in weights and weights[dt_key].ndim > 1: weights[dt_key] = mx.reshape(weights[dt_key], (-1,)) return weights @property def cast_predicate(self): def predicate(path: str): if "e_score_correction_bias" in path: return False if path.endswith("A_log") or path.endswith("dt_bias"): return False return True return predicate @property def quant_predicate(self): def predicate(path, _): if path.endswith("mlp.gate"): return {"group_size": 64, "bits": 8} return True return predicate