Swift-1.5-5bit-MLX / compatibility /swift15-mlx-lm.patch
ukisai's picture
initial release
e476358
Raw History Blame Contribute Delete
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
@@ -0,0 +1,485 @@
+# 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
@@ -195,6 +195,13 @@ 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.
@@ -216,6 +223,8 @@ 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:
@@ -446,12 +455,35 @@ 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:
@@ -1064,9 +1096,11 @@ 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"]
@@ -1099,7 +1133,8 @@ 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
@@ -0,0 +1,377 @@
+# 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)