Solar-Open2-250B-MLX-4bit / solar_open2.py
Vontra
Clean public model snapshot
c444e5d
Raw History Blame
13.6 kB
# 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