Instructions to use ukisai/Swift-1.5-5bit-MLX with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use ukisai/Swift-1.5-5bit-MLX with MLX:
# Make sure mlx-lm is installed # pip install --upgrade mlx-lm # Generate text with mlx-lm from mlx_lm import load, generate model, tokenizer = load("ukisai/Swift-1.5-5bit-MLX") prompt = "Write a story about Einstein" messages = [{"role": "user", "content": prompt}] prompt = tokenizer.apply_chat_template( messages, add_generation_prompt=True ) text = generate(model, tokenizer, prompt=prompt, verbose=True) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Pi
How to use ukisai/Swift-1.5-5bit-MLX with Pi:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "ukisai/Swift-1.5-5bit-MLX"
Configure the model in Pi
# Install Pi: npm install -g @earendil-works/pi-coding-agent # Add to ~/.pi/agent/models.json: { "providers": { "mlx-lm": { "baseUrl": "http://localhost:8080/v1", "api": "openai-completions", "apiKey": "none", "models": [ { "id": "ukisai/Swift-1.5-5bit-MLX" } ] } } }Run Pi
# Start Pi in your project directory: pi
- MLX LM
How to use ukisai/Swift-1.5-5bit-MLX with MLX LM:
Generate or start a chat session
# Install MLX LM uv tool install mlx-lm # Interactive chat REPL mlx_lm.chat --model "ukisai/Swift-1.5-5bit-MLX"
Run an OpenAI-compatible server
# Install MLX LM uv tool install mlx-lm # Start the server mlx_lm.server --model "ukisai/Swift-1.5-5bit-MLX" # Calling the OpenAI-compatible server with curl curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "ukisai/Swift-1.5-5bit-MLX", "messages": [ {"role": "user", "content": "Hello"} ] }' - Hermes Agent
How to use ukisai/Swift-1.5-5bit-MLX with Hermes Agent:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "ukisai/Swift-1.5-5bit-MLX"
Configure Hermes
# Install Hermes: curl -fsSL https://hermes-agent.nousresearch.com/install.sh | bash hermes setup # Point Hermes at the local server: hermes config set model.provider custom hermes config set model.base_url http://127.0.0.1:8080/v1 hermes config set model.default ukisai/Swift-1.5-5bit-MLX
Run Hermes
hermes
- Atomic Chat
- OpenClaw
How to use ukisai/Swift-1.5-5bit-MLX with OpenClaw:
Start the MLX server
# Install MLX LM: uv tool install mlx-lm # Start a local OpenAI-compatible server: mlx_lm.server --model "ukisai/Swift-1.5-5bit-MLX"
Configure OpenClaw
# Install OpenClaw: npm install -g openclaw@latest # Register the local server and set it as the default model: openclaw onboard --non-interactive --mode local \ --auth-choice custom-api-key \ --custom-base-url http://127.0.0.1:8080/v1 \ --custom-model-id "ukisai/Swift-1.5-5bit-MLX" \ --custom-provider-id mlx-lm \ --custom-compatibility openai \ --custom-text-input \ --accept-risk \ --skip-health
Run OpenClaw
openclaw agent --local --agent main --message "Hello from Hugging Face"
Download compatibility/swift15-mlx-lm.patch from ukisai/Swift-1.5-5bit-MLX: direct link, hf CLI and curl.
- Browser
- Download file 39 kB
-
https://huggingface.co/ukisai/Swift-1.5-5bit-MLX/resolve/main/compatibility/swift15-mlx-lm.patch
- Command line
-
hf download hf://ukisai/Swift-1.5-5bit-MLX/compatibility/swift15-mlx-lm.patch
-
curl -L -o swift15-mlx-lm.patch https://huggingface.co/ukisai/Swift-1.5-5bit-MLX/resolve/main/compatibility/swift15-mlx-lm.patch
39 kB
| diff --git a/mlx_lm/models/qwen3_5_full.py b/mlx_lm/models/qwen3_5_full.py | |
| new file mode 100644 | |
| index 0000000..eb3d09b | |
| --- /dev/null | |
| +++ b/mlx_lm/models/qwen3_5_full.py | |
| +# Copyright © 2026 Apple Inc. | |
| + | |
| +"""Complete Qwen3.5 parameter model, including the vision encoder and MTP. | |
| + | |
| +Text generation, the vision encoder, and explicit MTP steps are separate APIs. | |
| +Image/video token insertion, multimodal text positions, and speculative decoding | |
| +are not implemented here. Unsupported multimodal calls raise an error. | |
| +""" | |
| + | |
| +import copy | |
| +from dataclasses import dataclass | |
| +from typing import Optional | |
| + | |
| +import mlx.core as mx | |
| +import mlx.nn as nn | |
| +import numpy as np | |
| +from mlx.utils import tree_flatten, tree_unflatten | |
| + | |
| +from . import qwen3_5 | |
| +from .base import BaseModelArgs, create_attention_mask | |
| +from .cache import KVCache | |
| + | |
| + | |
| +class OffsetRMSNorm(nn.Module): | |
| + """Keep HF's zero-centered weights without rounding weight + 1 to BF16.""" | |
| + | |
| + def __init__(self, dims, eps=1e-6): | |
| + super().__init__() | |
| + self.weight = mx.zeros((dims,)) | |
| + self.eps = eps | |
| + | |
| + def __call__(self, x): | |
| + y = x.astype(mx.float32) | |
| + y = y * mx.rsqrt(mx.mean(y * y, axis=-1, keepdims=True) + self.eps) | |
| + return (y * (1 + self.weight.astype(mx.float32))).astype(x.dtype) | |
| + | |
| + | |
| +def _use_offset_norms(module): | |
| + replacements = [ | |
| + (name, OffsetRMSNorm(norm.weight.shape[0], norm.eps)) | |
| + for name, norm in module.named_modules() | |
| + if isinstance(norm, nn.RMSNorm) | |
| + ] | |
| + module.update_modules(tree_unflatten(replacements)) | |
| + | |
| + | |
| +@dataclass | |
| +class VisionArgs(BaseModelArgs): | |
| + depth: int | |
| + hidden_size: int | |
| + intermediate_size: int | |
| + num_heads: int | |
| + out_hidden_size: int | |
| + num_position_embeddings: int | |
| + patch_size: int = 16 | |
| + temporal_patch_size: int = 2 | |
| + spatial_merge_size: int = 2 | |
| + in_channels: int = 3 | |
| + hidden_act: str = "gelu_pytorch_tanh" | |
| + deepstack_visual_indexes: Optional[list] = None | |
| + | |
| + def __post_init__(self): | |
| + if self.deepstack_visual_indexes: | |
| + raise ValueError("DeepStack vision features are not supported") | |
| + if self.hidden_act != "gelu_pytorch_tanh": | |
| + raise ValueError(f"Unsupported vision activation: {self.hidden_act}") | |
| + if self.hidden_size % self.num_heads or self.hidden_size // self.num_heads % 4: | |
| + raise ValueError("Vision head dimension must be divisible by four") | |
| + if int(self.num_position_embeddings**0.5) ** 2 != self.num_position_embeddings: | |
| + raise ValueError("Vision position table must be square") | |
| + | |
| + | |
| +class VisionPatchEmbed(nn.Module): | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + self.args = args | |
| + kernel = (args.temporal_patch_size, args.patch_size, args.patch_size) | |
| + self.proj = nn.Conv3d( | |
| + args.in_channels, args.hidden_size, kernel, stride=kernel, bias=True | |
| + ) | |
| + | |
| + def __call__(self, pixels): | |
| + a = self.args | |
| + x = pixels.reshape( | |
| + -1, a.in_channels, a.temporal_patch_size, a.patch_size, a.patch_size | |
| + ) | |
| + x = x.transpose(0, 2, 3, 4, 1).astype(self.proj.weight.dtype) | |
| + return self.proj(x).reshape(-1, a.hidden_size) | |
| + | |
| + | |
| +class VisionAttention(nn.Module): | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + self.num_heads = args.num_heads | |
| + self.head_dim = args.hidden_size // args.num_heads | |
| + self.qkv = nn.Linear(args.hidden_size, 3 * args.hidden_size) | |
| + self.proj = nn.Linear(args.hidden_size, args.hidden_size) | |
| + | |
| + def __call__(self, x, cos, sin, boundaries): | |
| + qkv = self.qkv(x).reshape(x.shape[0], 3, self.num_heads, self.head_dim) | |
| + q, k, v = qkv.transpose(1, 0, 2, 3) | |
| + | |
| + def rotate(y): | |
| + z = y.astype(mx.float32) | |
| + half = self.head_dim // 2 | |
| + rotated = mx.concatenate([-z[..., half:], z[..., :half]], axis=-1) | |
| + return (z * cos[:, None] + rotated * sin[:, None]).astype(y.dtype) | |
| + | |
| + q, k = rotate(q), rotate(k) | |
| + outputs = [] | |
| + for start, end in zip(boundaries, boundaries[1:]): | |
| + out = mx.fast.scaled_dot_product_attention( | |
| + q[start:end].transpose(1, 0, 2)[None], | |
| + k[start:end].transpose(1, 0, 2)[None], | |
| + v[start:end].transpose(1, 0, 2)[None], | |
| + scale=self.head_dim**-0.5, | |
| + ) | |
| + outputs.append(out[0].transpose(1, 0, 2).reshape(end - start, -1)) | |
| + return self.proj(mx.concatenate(outputs, axis=0)) | |
| + | |
| + | |
| +class VisionMLP(nn.Module): | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + self.linear_fc1 = nn.Linear(args.hidden_size, args.intermediate_size) | |
| + self.linear_fc2 = nn.Linear(args.intermediate_size, args.hidden_size) | |
| + | |
| + def __call__(self, x): | |
| + return self.linear_fc2(nn.gelu_approx(self.linear_fc1(x))) | |
| + | |
| + | |
| +class VisionBlock(nn.Module): | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + self.norm1 = nn.LayerNorm(args.hidden_size, eps=1e-6) | |
| + self.norm2 = nn.LayerNorm(args.hidden_size, eps=1e-6) | |
| + self.attn = VisionAttention(args) | |
| + self.mlp = VisionMLP(args) | |
| + | |
| + def __call__(self, x, cos, sin, boundaries): | |
| + x = x + self.attn(self.norm1(x), cos, sin, boundaries) | |
| + return x + self.mlp(self.norm2(x)) | |
| + | |
| + | |
| +class VisionMerger(nn.Module): | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + self.hidden_size = args.hidden_size * args.spatial_merge_size**2 | |
| + self.norm = nn.LayerNorm(args.hidden_size, eps=1e-6) | |
| + self.linear_fc1 = nn.Linear(self.hidden_size, self.hidden_size) | |
| + self.linear_fc2 = nn.Linear(self.hidden_size, args.out_hidden_size) | |
| + | |
| + def __call__(self, x): | |
| + x = self.norm(x).reshape(-1, self.hidden_size) | |
| + return self.linear_fc2(nn.gelu(self.linear_fc1(x))) | |
| + | |
| + | |
| +class VisionModel(nn.Module): | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + self.args = args | |
| + self.patch_embed = VisionPatchEmbed(args) | |
| + self.pos_embed = nn.Embedding(args.num_position_embeddings, args.hidden_size) | |
| + self.blocks = [VisionBlock(args) for _ in range(args.depth)] | |
| + self.merger = VisionMerger(args) | |
| + | |
| + def _positions(self, grid_thw): | |
| + grid = np.asarray( | |
| + grid_thw.tolist() if hasattr(grid_thw, "tolist") else grid_thw | |
| + ) | |
| + if ( | |
| + grid.ndim != 2 | |
| + or grid.shape[1] != 3 | |
| + or not np.issubdtype(grid.dtype, np.integer) | |
| + ): | |
| + raise ValueError("grid_thw must be an integer array with shape (N, 3)") | |
| + merge = self.args.spatial_merge_size | |
| + side = int(self.args.num_position_embeddings**0.5) | |
| + positions, indices, weights, boundaries = [], [], [], [0] | |
| + for t, h, w in grid.tolist(): | |
| + if min(t, h, w) <= 0 or h % merge or w % merge: | |
| + raise ValueError("Invalid vision grid or spatial merge dimensions") | |
| + hp, wp = np.indices((h, w)) | |
| + block = (h // merge, merge, w // merge, merge) | |
| + hp = hp.reshape(block).transpose(0, 2, 1, 3).reshape(-1) | |
| + wp = wp.reshape(block).transpose(0, 2, 1, 3).reshape(-1) | |
| + positions.append(np.tile(np.stack([hp, wp], axis=-1), (t, 1))) | |
| + reorder = np.tile(hp * w + wp, t) | |
| + hs = np.linspace(0, side - 1, h, dtype=np.float32) | |
| + ws = np.linspace(0, side - 1, w, dtype=np.float32) | |
| + hf, wf = hs.astype(np.int32), ws.astype(np.int32) | |
| + hc, wc = np.minimum(hf + 1, side - 1), np.minimum(wf + 1, side - 1) | |
| + dh, dw = hs - hf, ws - wf | |
| + idx = [ | |
| + (a[:, None] * side + b[None]).reshape(-1) | |
| + for a, b in [(hf, wf), (hf, wc), (hc, wf), (hc, wc)] | |
| + ] | |
| + coeff = [ | |
| + (a[:, None] * b[None]).reshape(-1) | |
| + for a, b in [(1 - dh, 1 - dw), (1 - dh, dw), (dh, 1 - dw), (dh, dw)] | |
| + ] | |
| + indices.append(np.stack(idx)[:, reorder]) | |
| + weights.append(np.stack(coeff)[:, reorder]) | |
| + offset = boundaries[-1] | |
| + boundaries.extend(offset + (i + 1) * h * w for i in range(t)) | |
| + if not positions: | |
| + raise ValueError("At least one vision grid is required") | |
| + return ( | |
| + mx.array(np.concatenate(positions), dtype=mx.float32), | |
| + mx.array(np.concatenate(indices, axis=1), dtype=mx.int32), | |
| + mx.array(np.concatenate(weights, axis=1), dtype=mx.float32), | |
| + boundaries, | |
| + ) | |
| + | |
| + def __call__(self, pixels, grid_thw, return_hidden_states=False): | |
| + positions, indices, weights, boundaries = self._positions(grid_thw) | |
| + x = self.patch_embed(pixels) | |
| + if x.shape[0] != boundaries[-1]: | |
| + raise ValueError("Pixel patch count does not match grid_thw") | |
| + pos = (self.pos_embed(indices) * weights[..., None]).sum(axis=0) | |
| + x = x + pos.astype(x.dtype) | |
| + dim = self.args.hidden_size // self.args.num_heads // 2 | |
| + inv_freq = 1.0 / (10000 ** (mx.arange(0, dim, 2, dtype=mx.float32) / dim)) | |
| + angles = (positions[..., None] * inv_freq).reshape(x.shape[0], -1) | |
| + angles = mx.concatenate([angles, angles], axis=-1) | |
| + cos, sin = mx.cos(angles), mx.sin(angles) | |
| + for block in self.blocks: | |
| + x = block(x, cos, sin, boundaries) | |
| + merged = self.merger(x) | |
| + return (x, merged) if return_hidden_states else merged | |
| + | |
| + | |
| +class MultiTokenPredictor(nn.Module): | |
| + """An explicit MTP step; embeddings and the output head belong to the LM.""" | |
| + | |
| + def __init__(self, args, num_layers): | |
| + super().__init__() | |
| + self.fc = nn.Linear(2 * args.hidden_size, args.hidden_size, bias=False) | |
| + self.pre_fc_norm_embedding = OffsetRMSNorm(args.hidden_size, args.rms_norm_eps) | |
| + self.pre_fc_norm_hidden = OffsetRMSNorm(args.hidden_size, args.rms_norm_eps) | |
| + self.layers = [ | |
| + qwen3_5.DecoderLayer(args, args.full_attention_interval - 1) | |
| + for _ in range(num_layers) | |
| + ] | |
| + self.norm = OffsetRMSNorm(args.hidden_size, args.rms_norm_eps) | |
| + _use_offset_norms(self) | |
| + | |
| + def __call__(self, hidden_states, next_token_embeddings, cache=None, step=0): | |
| + if hidden_states.shape != next_token_embeddings.shape: | |
| + raise ValueError("MTP hidden states and next-token embeddings must align") | |
| + if step < 0: | |
| + raise ValueError("MTP step must be non-negative") | |
| + x = mx.concatenate( | |
| + [ | |
| + self.pre_fc_norm_embedding(next_token_embeddings), | |
| + self.pre_fc_norm_hidden(hidden_states), | |
| + ], | |
| + axis=-1, | |
| + ) | |
| + x = self.fc(x) | |
| + layer = step % len(self.layers) | |
| + if cache is not None and len(cache) != len(self.layers): | |
| + raise ValueError("MTP requires one KV cache per MTP layer") | |
| + c = cache[layer] if cache is not None else None | |
| + return self.norm(self.layers[layer](x, create_attention_mask(x, c), c)) | |
| + | |
| + def make_cache(self): | |
| + return [KVCache() for _ in self.layers] | |
| + | |
| + | |
| +@dataclass | |
| +class ModelArgs(BaseModelArgs): | |
| + model_type: str | |
| + text_config: dict | |
| + vision_config: dict | |
| + language_model_only: bool = False | |
| + image_token_id: Optional[int] = None | |
| + video_token_id: Optional[int] = None | |
| + vision_start_token_id: Optional[int] = None | |
| + tie_word_embeddings: bool = False | |
| + quantization: Optional[dict] = None | |
| + | |
| + @classmethod | |
| + def from_dict(cls, params): | |
| + return super().from_dict(copy.deepcopy(params)) | |
| + | |
| + | |
| +class Model(nn.Module): | |
| + extra_save_files = ( | |
| + "preprocessor_config.json", | |
| + "video_preprocessor_config.json", | |
| + "processor_config.json", | |
| + "tokenizer.json", | |
| + "tokenizer_config.json", | |
| + "special_tokens_map.json", | |
| + "vocab.json", | |
| + "merges.txt", | |
| + "chat_template.jinja", | |
| + ) | |
| + | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + self.args = args | |
| + self.model_type = args.model_type | |
| + text = args.text_config | |
| + if args.language_model_only: | |
| + raise ValueError("The complete model requires language_model_only=false") | |
| + if text.get("hidden_act", "silu") not in ("silu", "swish"): | |
| + raise ValueError("Unsupported text MLP activation") | |
| + if text.get("attn_output_gate", True) is not True: | |
| + raise ValueError("Ungated attention is not supported") | |
| + if text.get("output_gate_type", "swish") not in ("swish", "sigmoid"): | |
| + raise ValueError("Unknown attention gate declaration") | |
| + if text.get("mtp_use_dedicated_embeddings", False): | |
| + raise ValueError("Dedicated MTP embeddings are not supported") | |
| + if args.tie_word_embeddings != text.get("tie_word_embeddings", False): | |
| + raise ValueError("Conflicting text and top-level weight tying settings") | |
| + self._text_args = qwen3_5.TextModelArgs.from_dict(copy.deepcopy(text)) | |
| + expected_layers = [ | |
| + ( | |
| + "full_attention" | |
| + if (i + 1) % self._text_args.full_attention_interval == 0 | |
| + else "linear_attention" | |
| + ) | |
| + for i in range(self._text_args.num_hidden_layers) | |
| + ] | |
| + if text.get("layer_types", expected_layers) != expected_layers: | |
| + raise ValueError("Layer schedule differs from the supported architecture") | |
| + if self._text_args.num_experts: | |
| + raise ValueError("This complete model supports the dense architecture") | |
| + self.language_model = qwen3_5.TextModel(self._text_args) | |
| + _use_offset_norms(self.language_model) | |
| + self.visual = VisionModel(VisionArgs.from_dict(args.vision_config)) | |
| + if self.visual.args.out_hidden_size != self._text_args.hidden_size: | |
| + raise ValueError("Vision output size does not match text hidden size") | |
| + num_mtp = text.get("mtp_num_hidden_layers", 0) | |
| + if num_mtp < 0: | |
| + raise ValueError("Invalid MTP layer count") | |
| + if num_mtp: | |
| + self.mtp = MultiTokenPredictor(self._text_args, num_mtp) | |
| + self._parameter_shapes = { | |
| + k: tuple(v.shape) for k, v in tree_flatten(self.parameters()) | |
| + } | |
| + | |
| + @property | |
| + def model(self): | |
| + return self.language_model.model | |
| + | |
| + @property | |
| + def layers(self): | |
| + return self.language_model.layers | |
| + | |
| + def make_cache(self): | |
| + return self.language_model.make_cache() | |
| + | |
| + def __call__(self, inputs, cache=None, input_embeddings=None, **kwargs): | |
| + if kwargs: | |
| + raise NotImplementedError( | |
| + "Multimodal text integration is not implemented; use visual() for encoder features" | |
| + ) | |
| + for token in ( | |
| + self.args.image_token_id, | |
| + self.args.video_token_id, | |
| + self.args.vision_start_token_id, | |
| + ): | |
| + if ( | |
| + token is not None | |
| + and inputs is not None | |
| + and bool(mx.any(inputs == token)) | |
| + ): | |
| + raise NotImplementedError( | |
| + "Multimodal token positions require a multimodal text runtime" | |
| + ) | |
| + return self.language_model(inputs, cache, input_embeddings) | |
| + | |
| + def mtp_logits(self, next_token_ids, previous_hidden_states, cache=None, step=0): | |
| + if not hasattr(self, "mtp"): | |
| + raise ValueError("This checkpoint has no MTP layers") | |
| + embeddings = self.model.embed_tokens(next_token_ids) | |
| + hidden = self.mtp(previous_hidden_states, embeddings, cache, step) | |
| + if self._text_args.tie_word_embeddings: | |
| + return self.model.embed_tokens.as_linear(hidden) | |
| + return self.language_model.lm_head(hidden) | |
| + | |
| + def weight_mapping(self): | |
| + rows = [] | |
| + for name, shape in sorted(self._parameter_shapes.items()): | |
| + source, source_shape, transform = name, shape, "identity" | |
| + if name.startswith("language_model.model."): | |
| + source = "model.language_model." + name[len("language_model.model.") :] | |
| + elif name.startswith("language_model.lm_head."): | |
| + source = name[len("language_model.") :] | |
| + elif name.startswith("visual."): | |
| + source = "model." + name | |
| + if name.endswith(".conv1d.weight"): | |
| + source_shape = (shape[0], shape[2], shape[1]) | |
| + transform = "transpose(0,2,1)" | |
| + elif name == "visual.patch_embed.proj.weight": | |
| + source_shape = (shape[0], shape[4], shape[1], shape[2], shape[3]) | |
| + transform = "transpose(0,2,3,4,1)" | |
| + category = ( | |
| + "vision" | |
| + if name.startswith("visual.") | |
| + else "MTP" if name.startswith("mtp.") else "text" | |
| + ) | |
| + rows.append( | |
| + dict( | |
| + source=source, | |
| + source_shape=source_shape, | |
| + destination=name, | |
| + destination_shape=shape, | |
| + category=category, | |
| + transform=transform, | |
| + ) | |
| + ) | |
| + return rows | |
| + | |
| + def sanitize(self, weights): | |
| + rows = self.weight_mapping() | |
| + hf_names = {r["source"] for r in rows} | |
| + native_names = set(self._parameter_shapes) | |
| + is_hf = any(k.startswith("model.language_model.") for k in weights) | |
| + expected = hf_names if is_hf else native_names | |
| + if self.args.quantization: | |
| + if is_hf: | |
| + raise ValueError("Only original unquantized HF weights can be mapped") | |
| + return self._check_native_quantized(weights) | |
| + missing, unexpected = expected - set(weights), set(weights) - expected | |
| + if missing or unexpected: | |
| + raise ValueError( | |
| + f"Incomplete weight mapping: missing={sorted(missing)}, unexpected={sorted(unexpected)}" | |
| + ) | |
| + mapped = {} | |
| + for row in rows: | |
| + key = row["source"] if is_hf else row["destination"] | |
| + value = weights[key] | |
| + shape = row["source_shape"] if is_hf else row["destination_shape"] | |
| + if tuple(value.shape) != shape: | |
| + raise ValueError(f"Shape mismatch for {key}: {value.shape} != {shape}") | |
| + if is_hf and row["transform"] == "transpose(0,2,1)": | |
| + value = value.transpose(0, 2, 1) | |
| + elif is_hf and row["transform"] == "transpose(0,2,3,4,1)": | |
| + value = value.transpose(0, 2, 3, 4, 1) | |
| + mapped[row["destination"]] = value | |
| + return mapped | |
| + | |
| + def _check_native_quantized(self, weights): | |
| + shapes = dict(self._parameter_shapes) | |
| + q = self.args.quantization | |
| + for path, module in self.named_modules(): | |
| + if not hasattr(module, "to_quantized"): | |
| + continue | |
| + settings = q.get(path, q) | |
| + if settings is False: | |
| + continue | |
| + if not isinstance(settings, dict): | |
| + raise ValueError(f"Invalid native quantization metadata for {path}") | |
| + bits, group, mode = ( | |
| + settings.get("bits"), | |
| + settings.get("group_size"), | |
| + settings.get("mode", "affine"), | |
| + ) | |
| + if (bits, group, mode) != (4, 64, "affine"): | |
| + raise ValueError( | |
| + "This extension only prepares affine/4-bit/group-64 native checkpoints" | |
| + ) | |
| + original = shapes[f"{path}.weight"] | |
| + if original[-1] % group: | |
| + continue | |
| + shapes[f"{path}.weight"] = (*original[:-1], original[-1] * bits // 32) | |
| + shapes[f"{path}.scales"] = (*original[:-1], original[-1] // group) | |
| + shapes[f"{path}.biases"] = shapes[f"{path}.scales"] | |
| + missing, unexpected = set(shapes) - set(weights), set(weights) - set(shapes) | |
| + if missing or unexpected: | |
| + raise ValueError( | |
| + f"Incomplete native checkpoint: missing={sorted(missing)}, unexpected={sorted(unexpected)}" | |
| + ) | |
| + for name, shape in shapes.items(): | |
| + if tuple(weights[name].shape) != shape: | |
| + raise ValueError(f"Native checkpoint shape mismatch: {name}") | |
| + return weights | |
| + | |
| + @property | |
| + def cast_predicate(self): | |
| + return self.language_model.cast_predicate | |
| diff --git a/mlx_lm/utils.py b/mlx_lm/utils.py | |
| index a00fed8..6a751a2 100644 | |
| --- a/mlx_lm/utils.py | |
| +++ b/mlx_lm/utils.py | |
| def _transform_awq_weights( | |
| return new_weights, mlx_quantization | |
| +def _is_complete_qwen3_5(config: dict) -> bool: | |
| + return ( | |
| + config.get("model_type") == "qwen3_5" | |
| + and config.get("language_model_only") is False | |
| + ) | |
| + | |
| + | |
| def _get_classes(config: dict): | |
| """ | |
| Retrieve the model and model args classes based on the configuration. | |
| def _get_classes(config: dict): | |
| break | |
| else: | |
| model_type = MODEL_REMAPPING.get(model_type, model_type) | |
| + if _is_complete_qwen3_5(config): | |
| + model_type = "qwen3_5_full" | |
| try: | |
| arch = importlib.import_module(f"mlx_lm.models.{model_type}") | |
| except ImportError as e: | |
| def load_model( | |
| weight_files = glob.glob(str(model_path / "model*.safetensors")) | |
| + complete_qwen = _is_complete_qwen3_5(config) | |
| + if complete_qwen: | |
| + with open(model_path / "model.safetensors.index.json") as stream: | |
| + index = json.load(stream) | |
| + expected_files = set(index["weight_map"].values()) | |
| + actual_files = {Path(file).name for file in weight_files} | |
| + if actual_files != expected_files: | |
| + raise ValueError( | |
| + "Incomplete full-model checkpoint: " | |
| + f"missing shards={sorted(expected_files - actual_files)}, " | |
| + f"unexpected shards={sorted(actual_files - expected_files)}" | |
| + ) | |
| + | |
| if not weight_files and strict: | |
| raise FileNotFoundError(f"No safetensors found in {model_path}") | |
| weights = {} | |
| for wf in weight_files: | |
| - weights.update(mx.load(wf)) | |
| + shard = mx.load(wf) | |
| + if complete_qwen: | |
| + duplicates = weights.keys() & shard.keys() | |
| + if duplicates: | |
| + raise ValueError(f"Duplicate checkpoint tensors: {sorted(duplicates)}") | |
| + for name in shard: | |
| + if index["weight_map"].get(name) != Path(wf).name: | |
| + raise ValueError(f"Checkpoint index mismatch: {name}") | |
| + weights.update(shard) | |
| + if complete_qwen and weights.keys() != index["weight_map"].keys(): | |
| + raise ValueError("Checkpoint tensor set does not match its index") | |
| if (model_file := config.get("model_file")) is not None: | |
| if not trust_remote_code: | |
| def save_config( | |
| config (dict): The model configuration. | |
| config_path (Union[str, Path]): Model configuration file path. | |
| """ | |
| - # Clean unused keys | |
| + config = copy.deepcopy(config) | |
| + # Complete multimodal checkpoints need the vision architecture on reload. | |
| config.pop("_name_or_path", None) | |
| - config.pop("vision_config", None) | |
| + if not _is_complete_qwen3_5(config): | |
| + config.pop("vision_config", None) | |
| if "quantization" in config: | |
| config["quantization_config"] = config["quantization"] | |
| def save( | |
| save_config(config, config_path=dst_path / "config.json") | |
| tokenizer.save_pretrained(dst_path) | |
| - for p in ["*.py", "generation_config.json"]: | |
| + extra_files = getattr(model, "extra_save_files", ()) | |
| + for p in ["*.py", "generation_config.json", *extra_files]: | |
| for file in glob.glob(str(src_path / p)): | |
| shutil.copy(file, dst_path) | |
| diff --git a/tests/test_qwen3_5_full.py b/tests/test_qwen3_5_full.py | |
| new file mode 100644 | |
| index 0000000..2b3e5ae | |
| --- /dev/null | |
| +++ b/tests/test_qwen3_5_full.py | |
| +# Copyright © 2026 Apple Inc. | |
| + | |
| +"""Small synthetic fixtures test architecture code, not Swift model quality.""" | |
| + | |
| +import copy | |
| +import json | |
| +from pathlib import Path | |
| +from unittest.mock import patch | |
| + | |
| +import mlx.core as mx | |
| +import numpy as np | |
| +import pytest | |
| +import torch | |
| +from mlx.utils import tree_flatten | |
| +from mlx_lm.convert import convert | |
| +from mlx_lm.models.qwen3_5_full import Model, ModelArgs, OffsetRMSNorm | |
| +from mlx_lm.utils import _get_classes, load_model, save_config | |
| +from transformers import Qwen3_5Config, Qwen3_5ForConditionalGeneration | |
| +from transformers.models.qwen3_5.modeling_qwen3_5 import ( | |
| + Qwen3_5DecoderLayer, | |
| + Qwen3_5RMSNorm, | |
| + Qwen3_5TextRotaryEmbedding, | |
| +) | |
| + | |
| + | |
| +def small_config(): | |
| + return { | |
| + "model_type": "qwen3_5", | |
| + "architectures": ["Qwen3_5ForConditionalGeneration"], | |
| + "language_model_only": False, | |
| + "tie_word_embeddings": False, | |
| + "image_token_id": 125, | |
| + "video_token_id": 126, | |
| + "vision_start_token_id": 127, | |
| + "text_config": { | |
| + "model_type": "qwen3_5_text", | |
| + "hidden_size": 128, | |
| + "intermediate_size": 256, | |
| + "num_hidden_layers": 4, | |
| + "num_attention_heads": 4, | |
| + "num_key_value_heads": 2, | |
| + "head_dim": 32, | |
| + "vocab_size": 128, | |
| + "full_attention_interval": 4, | |
| + "layer_types": ["linear_attention"] * 3 + ["full_attention"], | |
| + "linear_num_key_heads": 2, | |
| + "linear_num_value_heads": 4, | |
| + "linear_key_head_dim": 32, | |
| + "linear_value_head_dim": 32, | |
| + "linear_conv_kernel_dim": 4, | |
| + "hidden_act": "silu", | |
| + "attn_output_gate": True, | |
| + "output_gate_type": "swish", | |
| + "mamba_ssm_dtype": "float32", | |
| + "rms_norm_eps": 1e-6, | |
| + "max_position_embeddings": 256, | |
| + "tie_word_embeddings": False, | |
| + "attention_bias": False, | |
| + "attention_dropout": 0.0, | |
| + "mtp_num_hidden_layers": 1, | |
| + "mtp_use_dedicated_embeddings": False, | |
| + "rope_parameters": { | |
| + "rope_type": "default", | |
| + "rope_theta": 10000000, | |
| + "partial_rotary_factor": 0.5, | |
| + "mrope_interleaved": True, | |
| + "mrope_section": [3, 3, 2], | |
| + }, | |
| + }, | |
| + "vision_config": { | |
| + "model_type": "qwen3_5", | |
| + "depth": 2, | |
| + "hidden_size": 32, | |
| + "intermediate_size": 48, | |
| + "num_heads": 4, | |
| + "out_hidden_size": 128, | |
| + "num_position_embeddings": 16, | |
| + "patch_size": 2, | |
| + "temporal_patch_size": 2, | |
| + "spatial_merge_size": 2, | |
| + "in_channels": 3, | |
| + "hidden_act": "gelu_pytorch_tanh", | |
| + "deepstack_visual_indexes": [], | |
| + }, | |
| + } | |
| + | |
| + | |
| +class ReferenceMTP(torch.nn.Module): | |
| + """Single-step composition used by the source vLLM MTP implementation.""" | |
| + | |
| + def __init__(self, args): | |
| + super().__init__() | |
| + h = args.hidden_size | |
| + self.fc = torch.nn.Linear(2 * h, h, bias=False) | |
| + self.pre_fc_norm_embedding = Qwen3_5RMSNorm(h, args.rms_norm_eps) | |
| + self.pre_fc_norm_hidden = Qwen3_5RMSNorm(h, args.rms_norm_eps) | |
| + self.layers = torch.nn.ModuleList([Qwen3_5DecoderLayer(args, 3)]) | |
| + self.norm = Qwen3_5RMSNorm(h, args.rms_norm_eps) | |
| + self.rotary = Qwen3_5TextRotaryEmbedding(args) | |
| + | |
| + def forward(self, hidden, embeds): | |
| + x = self.fc( | |
| + torch.cat( | |
| + [self.pre_fc_norm_embedding(embeds), self.pre_fc_norm_hidden(hidden)], | |
| + dim=-1, | |
| + ) | |
| + ) | |
| + positions = torch.arange(x.shape[1])[None, None].expand(3, x.shape[0], -1) | |
| + rotary = self.rotary(x, positions) | |
| + mask = torch.triu( | |
| + torch.full((x.shape[0], 1, x.shape[1], x.shape[1]), float("-inf")), | |
| + diagonal=1, | |
| + ) | |
| + return self.norm( | |
| + self.layers[0](x, position_embeddings=rotary, attention_mask=mask) | |
| + ) | |
| + | |
| + | |
| +@pytest.fixture(scope="module") | |
| +def reference(): | |
| + torch.set_num_threads(2) | |
| + torch.manual_seed(71) | |
| + config = small_config() | |
| + hf_config = Qwen3_5Config(**copy.deepcopy(config)) | |
| + hf_config._attn_implementation = "eager" | |
| + hf_config.text_config._attn_implementation = "eager" | |
| + hf_config.vision_config._attn_implementation = "eager" | |
| + hf = Qwen3_5ForConditionalGeneration(hf_config).eval() | |
| + mtp = ReferenceMTP(hf_config.text_config).eval() | |
| + with torch.no_grad(): | |
| + for module in (hf, mtp): | |
| + for name, value in module.named_parameters(): | |
| + if "norm" in name and value.ndim == 1 and "model.visual" not in name: | |
| + value.uniform_(-0.15, 0.15) | |
| + state = {k: mx.array(v.detach().numpy()) for k, v in hf.state_dict().items()} | |
| + state.update( | |
| + {"mtp." + k: mx.array(v.detach().numpy()) for k, v in mtp.state_dict().items()} | |
| + ) | |
| + model = Model(ModelArgs.from_dict(config)) | |
| + model.load_weights(list(model.sanitize(state).items()), strict=True) | |
| + model.eval() | |
| + return config, hf, mtp, state, model | |
| + | |
| + | |
| +def assert_close(mlx_value, torch_value, atol=3e-5, rtol=3e-5): | |
| + np.testing.assert_allclose( | |
| + np.array(mlx_value), torch_value.detach().numpy(), atol=atol, rtol=rtol | |
| + ) | |
| + | |
| + | |
| +def test_dispatch_and_mapping_preserve_every_component(reference): | |
| + config, _, _, state, model = reference | |
| + before = copy.deepcopy(config) | |
| + cls, args = _get_classes(config) | |
| + fresh = cls(args.from_dict(config)) | |
| + assert cls is Model | |
| + assert config == before | |
| + rows = fresh.weight_mapping() | |
| + assert {r["source"] for r in rows} == set(state) | |
| + assert len(rows) == len(dict(tree_flatten(model.parameters()))) | |
| + assert sum(r["category"] == "MTP" for r in rows) == 15 | |
| + assert sum(r["category"] == "vision" for r in rows) == 33 | |
| + assert model.language_model.lm_head is not model.model.embed_tokens | |
| + assert not any("mtp.embed" in r["destination"] for r in rows) | |
| + | |
| + | |
| +@pytest.mark.parametrize("prefix", ["model.language_model.", "model.visual.", "mtp."]) | |
| +def test_missing_critical_tensor_is_rejected(reference, prefix): | |
| + _, _, _, state, model = reference | |
| + missing = dict(state) | |
| + missing.pop(next(k for k in state if k.startswith(prefix))) | |
| + with pytest.raises(ValueError, match="missing="): | |
| + model.sanitize(missing) | |
| + | |
| + | |
| +def test_unexpected_and_misshaped_tensors_are_rejected(reference): | |
| + _, _, _, state, model = reference | |
| + with pytest.raises(ValueError, match="unexpected="): | |
| + model.sanitize(dict(state, **{"mtp.unknown.weight": mx.zeros((1,))})) | |
| + wrong = dict(state) | |
| + wrong["mtp.fc.weight"] = mx.zeros((1,)) | |
| + with pytest.raises(ValueError, match="Shape mismatch"): | |
| + model.sanitize(wrong) | |
| + | |
| + | |
| +def test_native_mapping_is_idempotent_and_norm_values_are_exact(reference): | |
| + _, _, _, state, model = reference | |
| + once = model.sanitize(state) | |
| + twice = model.sanitize(once) | |
| + for key in once: | |
| + np.testing.assert_array_equal(np.array(once[key]), np.array(twice[key])) | |
| + for row in model.weight_mapping(): | |
| + if "norm" in row["source"]: | |
| + np.testing.assert_array_equal( | |
| + np.array(state[row["source"]]), np.array(once[row["destination"]]) | |
| + ) | |
| + | |
| + | |
| +def test_offset_norm_preserves_small_bf16_deltas(): | |
| + layer = OffsetRMSNorm(4) | |
| + raw = mx.array([0.0001, -0.0002, 0.0003, -0.0004], dtype=mx.bfloat16) | |
| + layer.weight = raw | |
| + x = mx.array([[0.8, -1.2, 0.4, 2.0]], dtype=mx.bfloat16) | |
| + torch_layer = Qwen3_5RMSNorm(4) | |
| + torch_layer.weight.data.copy_(torch.tensor(np.array(raw.astype(mx.float32)))) | |
| + expected = torch_layer( | |
| + torch.tensor(np.array(x.astype(mx.float32))).to(torch.bfloat16) | |
| + ).float() | |
| + assert_close(layer(x).astype(mx.float32), expected, atol=0, rtol=0) | |
| + np.testing.assert_array_equal( | |
| + np.array(layer.weight.astype(mx.float32)), np.array(raw.astype(mx.float32)) | |
| + ) | |
| + | |
| + | |
| +def test_text_forward_matches_transformers(reference): | |
| + _, hf, _, _, model = reference | |
| + ids = torch.tensor([[3, 19, 8, 21, 11]]) | |
| + with torch.no_grad(): | |
| + expected = hf(input_ids=ids, use_cache=False).logits | |
| + actual = model(mx.array(ids.numpy())) | |
| + assert_close(actual, expected) | |
| + | |
| + | |
| +def test_text_decode_cache_matches_full_forward(reference): | |
| + _, _, _, _, model = reference | |
| + ids = mx.array([[3, 19, 8, 21, 11]]) | |
| + full = model(ids) | |
| + cache = model.make_cache() | |
| + parts = [model(ids[:, :3], cache)] | |
| + parts.extend(model(ids[:, i : i + 1], cache) for i in range(3, 5)) | |
| + np.testing.assert_allclose( | |
| + np.array(mx.concatenate(parts, axis=1)), np.array(full), atol=3e-5, rtol=3e-5 | |
| + ) | |
| + | |
| + | |
| +@pytest.mark.parametrize("grid", [[[1, 4, 4]], [[2, 2, 4], [1, 4, 2]], [[1, 6, 4]]]) | |
| +def test_vision_encoder_matches_transformers(reference, grid): | |
| + _, hf, _, _, model = reference | |
| + torch.manual_seed(11) | |
| + count = sum(t * h * w for t, h, w in grid) | |
| + pixels = torch.randn(count, 3 * 2 * 2 * 2) | |
| + with torch.no_grad(): | |
| + expected = hf.model.visual(pixels, torch.tensor(grid)) | |
| + hidden, pooled = model.visual( | |
| + mx.array(pixels.numpy()), grid, return_hidden_states=True | |
| + ) | |
| + assert_close(hidden, expected.last_hidden_state) | |
| + assert_close(pooled, expected.pooler_output) | |
| + | |
| + | |
| +def test_mtp_forward_and_shared_head_match_reference(reference): | |
| + _, hf, mtp, _, model = reference | |
| + torch.manual_seed(9) | |
| + hidden = torch.randn(1, 4, 128) | |
| + ids = torch.tensor([[9, 8, 7, 6]]) | |
| + with torch.no_grad(): | |
| + embeds = hf.model.language_model.embed_tokens(ids) | |
| + expected = mtp(hidden, embeds) | |
| + expected_logits = hf.lm_head(expected) | |
| + assert_close( | |
| + model.mtp(mx.array(hidden.numpy()), mx.array(embeds.numpy())), expected | |
| + ) | |
| + assert_close( | |
| + model.mtp_logits(mx.array(ids.numpy()), mx.array(hidden.numpy())), | |
| + expected_logits, | |
| + ) | |
| + cache = model.mtp.make_cache() | |
| + cached = mx.concatenate( | |
| + [ | |
| + model.mtp( | |
| + mx.array(hidden[:, i : i + 1].numpy()), | |
| + mx.array(embeds[:, i : i + 1].numpy()), | |
| + cache, | |
| + ) | |
| + for i in range(4) | |
| + ], | |
| + axis=1, | |
| + ) | |
| + assert_close(cached, expected) | |
| + | |
| + | |
| +def test_unsupported_multimodal_generation_fails_explicitly(reference): | |
| + _, _, _, _, model = reference | |
| + with pytest.raises(NotImplementedError, match="Multimodal"): | |
| + model(mx.array([[1, 2]]), pixel_values=mx.zeros((1, 24))) | |
| + with pytest.raises(NotImplementedError, match="Multimodal"): | |
| + model(mx.array([[1, 125]])) | |
| + | |
| + | |
| +def test_save_config_keeps_vision_and_does_not_mutate_input(tmp_path): | |
| + config = small_config() | |
| + before = copy.deepcopy(config) | |
| + save_config(config, tmp_path / "config.json") | |
| + assert config == before | |
| + assert json.loads((tmp_path / "config.json").read_text()) == before | |
| + | |
| + | |
| +def test_official_nonquantized_convert_and_reload_preserve_all_tensors( | |
| + reference, tmp_path | |
| +): | |
| + config, _, _, state, model = reference | |
| + source, output = tmp_path / "hf", tmp_path / "mlx" | |
| + source.mkdir() | |
| + (source / "config.json").write_text(json.dumps(config)) | |
| + mx.save_safetensors(str(source / "model.safetensors"), state) | |
| + (source / "model.safetensors.index.json").write_text( | |
| + json.dumps({"weight_map": {k: "model.safetensors" for k in state}}) | |
| + ) | |
| + assets = ["generation_config.json", *model.extra_save_files] | |
| + for name in assets: | |
| + (source / name).write_text("original synthetic asset: " + name) | |
| + (source / "generation_config.json").write_text('{"eos_token_id": 2}') | |
| + | |
| + class FixtureTokenizer: | |
| + def save_pretrained(self, target): | |
| + Path(target, "tokenizer_config.json").write_text("rewritten") | |
| + | |
| + def fixture_load(path, **kwargs): | |
| + loaded, loaded_config = load_model(Path(path), lazy=True, strict=True) | |
| + return loaded, FixtureTokenizer(), loaded_config | |
| + | |
| + with patch("mlx_lm.convert.load", side_effect=fixture_load): | |
| + convert(str(source), str(output), quantize=False) | |
| + loaded, saved_config = load_model(output, lazy=False, strict=True) | |
| + expected = model.sanitize(state) | |
| + actual = dict(tree_flatten(loaded.parameters())) | |
| + assert set(actual) == set(expected) | |
| + for name in actual: | |
| + np.testing.assert_array_equal(np.array(actual[name]), np.array(expected[name])) | |
| + assert saved_config["vision_config"] == config["vision_config"] | |
| + assert saved_config["text_config"] == config["text_config"] | |
| + assert "quantization" not in saved_config | |
| + for name in assets: | |
| + assert (source / name).read_bytes() == (output / name).read_bytes() | |
| + ids = mx.array([[3, 19, 8]]) | |
| + np.testing.assert_array_equal(np.array(model(ids)), np.array(loaded(ids))) | |
| + (output / "model.safetensors.index.json").write_text( | |
| + json.dumps( | |
| + { | |
| + "weight_map": { | |
| + **{k: "model.safetensors" for k in actual}, | |
| + "mtp.missing.weight": "missing.safetensors", | |
| + } | |
| + } | |
| + ) | |
| + ) | |
| + with pytest.raises(ValueError, match="missing shards"): | |
| + load_model(output, lazy=True) | |
| + | |
| + | |
| +def test_fixed_affine_quantized_roundtrip_preserves_component_tree(reference, tmp_path): | |
| + from mlx_lm.utils import quantize_model, save_model | |
| + | |
| + config, _, _, state, _ = reference | |
| + model = Model(ModelArgs.from_dict(config)) | |
| + model.load_weights(list(model.sanitize(state).items()), strict=True) | |
| + original_names = set(dict(tree_flatten(model.parameters()))) | |
| + model, quantized_config = quantize_model(model, config, 64, 4, mode="affine") | |
| + save_model(tmp_path, model) | |
| + save_config(quantized_config, tmp_path / "config.json") | |
| + loaded, saved_config = load_model(tmp_path, lazy=False, strict=True) | |
| + actual = dict(tree_flatten(loaded.parameters())) | |
| + expected = dict(tree_flatten(model.parameters())) | |
| + assert set(actual) == set(expected) | |
| + assert original_names <= set(actual) | |
| + assert saved_config["quantization"] == { | |
| + "bits": 4, | |
| + "group_size": 64, | |
| + "mode": "affine", | |
| + } | |
| + for name in actual: | |
| + np.testing.assert_array_equal(np.array(actual[name]), np.array(expected[name])) | |
| + assert bool(mx.all(mx.isfinite(loaded(mx.array([[3, 19, 8]]))))) | |
| + missing = dict(actual) | |
| + missing.pop("mtp.fc.scales") | |
| + with pytest.raises(ValueError, match="Incomplete native checkpoint"): | |
| + loaded.sanitize(missing) | |