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)