"""MLX implementation for Laneformer causal language models. This file is copied into converted Laneformer artifacts as ``laneformer.py`` and loaded by ``mlx_lm`` via ``config.json``'s ``model_file`` field. """ from __future__ import annotations import inspect import math from enum import Enum from typing import Any, Optional import mlx.core as mx import mlx.nn as nn class ModelArgs: def __init__( self, model_type: str = "laneformer", hidden_size: int = 4096, num_hidden_layers: int = 32, num_attention_heads: int = 32, num_key_value_heads: Optional[int] = None, intermediate_size: int = 16384, max_position_embeddings: int = 4096, vocab_size: int = 32000, rope_theta: float = 10000.0, norm_eps: float = 1e-5, num_lanes: int = 1, broadcast_delay: int = 1, use_attention_comm: bool = True, use_mlp_comm: bool = True, use_early_comm: bool = True, lm_head_type: str = "replicate", pre_norm_lane_agg: bool = False, replicated_rmsn_scale: bool = True, tie_word_embeddings: bool = False, sliding_window: Optional[int] = None, swa_layers: Optional[list[int]] = None, layer_types: Optional[list[str]] = None, rope_scaling: Optional[dict[str, Any]] = None, ) -> None: self.model_type = model_type self.hidden_size = hidden_size self.num_hidden_layers = num_hidden_layers self.num_attention_heads = num_attention_heads self.num_key_value_heads = num_key_value_heads self.intermediate_size = intermediate_size self.max_position_embeddings = max_position_embeddings self.vocab_size = vocab_size self.rope_theta = rope_theta self.norm_eps = norm_eps self.num_lanes = num_lanes self.broadcast_delay = broadcast_delay self.use_attention_comm = use_attention_comm self.use_mlp_comm = use_mlp_comm self.use_early_comm = use_early_comm self.lm_head_type = lm_head_type self.pre_norm_lane_agg = pre_norm_lane_agg self.replicated_rmsn_scale = replicated_rmsn_scale self.tie_word_embeddings = tie_word_embeddings self.sliding_window = sliding_window self.swa_layers = swa_layers self.layer_types = layer_types self.rope_scaling = rope_scaling self.__post_init__() @classmethod def from_dict(cls, params: dict[str, Any]) -> "ModelArgs": return cls( **{ key: value for key, value in params.items() if key in inspect.signature(cls).parameters } ) def __post_init__(self) -> None: if self.num_key_value_heads is None: self.num_key_value_heads = self.num_attention_heads if self.swa_layers is None: self.swa_layers = [] if self.layer_types is None: self.layer_types = [ "sliding_attention" if i in set(self.swa_layers) else "full_attention" for i in range(self.num_hidden_layers) ] class ReduceMode(Enum): NO_REDUCE = "no_reduce" PRESENT = "present" PAST = "past" class HistoryCache: """Token-history cache for slow but correct MLX-LM generation. Laneformer's inter-layer lane communication makes a normal per-layer KV cache more involved. For artifact smoke tests and short generations we keep the token history and recompute the full context each step. """ def __init__(self) -> None: self.tokens: Optional[mx.array] = None self.offset = 0 def update_and_fetch(self, inputs: mx.array) -> mx.array: if self.tokens is None: self.tokens = inputs else: self.tokens = mx.concatenate([self.tokens, inputs], axis=1) self.offset = int(self.tokens.shape[1]) return self.tokens @property def state(self): return [] if self.tokens is None else self.tokens @state.setter def state(self, value) -> None: if value is None or value == []: self.tokens = None self.offset = 0 return self.tokens = value self.offset = int(value.shape[1]) @property def meta_state(self) -> str: return "" @meta_state.setter def meta_state(self, value) -> None: if value: raise ValueError("HistoryCache does not store metadata.") @property def nbytes(self) -> int: return 0 if self.tokens is None else self.tokens.nbytes def empty(self) -> bool: return self.tokens is None def is_trimmable(self) -> bool: return False def size(self) -> int: return self.offset class LaneModule: def __init__(self, no_reduce_scale: float, reduce_mode: ReduceMode): super().__init__() self.no_reduce_scale = no_reduce_scale self.reduce_mode = reduce_mode def reduce_lanes(self, x: mx.array, past: Optional[mx.array]) -> mx.array: if self.reduce_mode is ReduceMode.NO_REDUCE: return x * self.no_reduce_scale if self.reduce_mode is ReduceMode.PRESENT: return mx.sum(x, axis=2, keepdims=True) if self.reduce_mode is ReduceMode.PAST: if past is None: raise ValueError("ReduceMode.PAST requires a past lane tensor.") sum_past = mx.sum(past, axis=-2, keepdims=True) - past return x + sum_past raise ValueError(f"Unknown reduce mode: {self.reduce_mode}") class LaneRowLinear(nn.Module): """Per-lane row-parallel projection with weight shape ``[L, O/L, I]``.""" def __init__(self, in_features: int, out_features: int, num_lanes: int): super().__init__() self.in_features = in_features self.out_features = out_features self.num_lanes = num_lanes self.weight = mx.zeros((num_lanes, out_features // num_lanes, in_features)) def __call__(self, x: mx.array) -> mx.array: return mx.einsum("bsli,loi->bslo", x, self.weight) def to_quantized( self, group_size: Optional[int] = None, bits: Optional[int] = None, mode: str = "affine", ): return QuantizedLaneRowLinear.from_lane_linear(self, group_size, bits, mode) class LaneColumnLinear(nn.Module): """Per-lane column-parallel projection with weight shape ``[L, O, I/L]``.""" def __init__(self, in_features: int, out_features: int, num_lanes: int): super().__init__() self.in_features = in_features self.out_features = out_features self.num_lanes = num_lanes self.weight = mx.zeros((num_lanes, out_features, in_features // num_lanes)) def __call__(self, x: mx.array) -> mx.array: return mx.einsum("bsli,loi->bslo", x, self.weight) def to_quantized( self, group_size: Optional[int] = None, bits: Optional[int] = None, mode: str = "affine", ): return QuantizedLaneColumnLinear.from_lane_linear(self, group_size, bits, mode) class _QuantizedLaneLinear(nn.Module): def _quantized_lane_matmul(self, x: mx.array) -> mx.array: outputs = [] biases = self.get("biases") for lane in range(self.num_lanes): lane_biases = None if biases is None else biases[lane] outputs.append( mx.quantized_matmul( x[:, :, lane, :], self.weight[lane], scales=self.scales[lane], biases=lane_biases, transpose=True, group_size=self.group_size, bits=self.bits, mode=self.mode, ) ) return mx.stack(outputs, axis=2) def __call__(self, x: mx.array) -> mx.array: return self._quantized_lane_matmul(x) class QuantizedLaneRowLinear(_QuantizedLaneLinear): @classmethod def from_lane_linear( cls, linear: LaneRowLinear, group_size: Optional[int], bits: Optional[int], mode: str, ) -> "QuantizedLaneRowLinear": quantized = cls() quantized.in_features = linear.in_features quantized.out_features = linear.out_features quantized.num_lanes = linear.num_lanes quantized.group_size = group_size or 64 quantized.bits = bits or 4 quantized.mode = mode quantized.weight, quantized.scales, *biases = mx.quantize( linear.weight, group_size=quantized.group_size, bits=quantized.bits, mode=mode, ) quantized.biases = biases[0] if biases else None return quantized class QuantizedLaneColumnLinear(_QuantizedLaneLinear): @classmethod def from_lane_linear( cls, linear: LaneColumnLinear, group_size: Optional[int], bits: Optional[int], mode: str, ) -> "QuantizedLaneColumnLinear": quantized = cls() quantized.in_features = linear.in_features quantized.out_features = linear.out_features quantized.num_lanes = linear.num_lanes quantized.group_size = group_size or 64 quantized.bits = bits or 4 quantized.mode = mode quantized.weight, quantized.scales, *biases = mx.quantize( linear.weight, group_size=quantized.group_size, bits=quantized.bits, mode=mode, ) quantized.biases = biases[0] if biases else None return quantized class LaneRMSNorm(nn.Module): def __init__(self, dim: int, num_lanes: int, eps: float = 1e-5): super().__init__() self.eps = eps self.scale = mx.ones((num_lanes, dim)) def __call__(self, x: mx.array) -> mx.array: variance = mx.mean(mx.square(x), axis=-1, keepdims=True) return x * mx.rsqrt(variance + self.eps) * self.scale class LaneLMHead(nn.Module): def __init__(self, in_features: int, out_features: int, num_lanes: int): super().__init__() self.linear = nn.Linear(num_lanes * in_features, out_features, bias=False) def __call__(self, x: mx.array) -> mx.array: batch, seq_len, lanes, dim = x.shape return self.linear(x.reshape(batch, seq_len, lanes * dim)) def _rope_inv_freq(dim: int, theta: float, rope_scaling: Optional[dict[str, Any]]) -> mx.array: freqs = 1.0 / (theta ** (mx.arange(0, dim, 2, dtype=mx.float32) / dim)) if rope_scaling is None: return freqs scaling_factor = float(rope_scaling.get("scaling_factor", 1.0)) low_freq_factor = float(rope_scaling.get("low_freq_factor", 1.0)) high_freq_factor = float(rope_scaling.get("high_freq_factor", 1.0)) original_max_position_embeddings = float( rope_scaling.get("original_max_position_embeddings", 4096) ) wavelen = 2 * math.pi / freqs high_freq_wavelen = original_max_position_embeddings / high_freq_factor low_freq_wavelen = original_max_position_embeddings / low_freq_factor scaled = mx.where(wavelen > low_freq_wavelen, freqs / scaling_factor, freqs) smooth_factor = ( original_max_position_embeddings / wavelen - low_freq_factor ) / (high_freq_factor - low_freq_factor) smoothed = (1 - smooth_factor) * freqs / scaling_factor + smooth_factor * freqs medium = mx.logical_and(wavelen >= high_freq_wavelen, wavelen <= low_freq_wavelen) return mx.where(medium, smoothed, scaled) def _apply_rotary( queries: mx.array, keys: mx.array, position_ids: mx.array, theta: float, rope_scaling: Optional[dict[str, Any]], ) -> tuple[mx.array, mx.array]: head_dim = queries.shape[-1] freqs = _rope_inv_freq(head_dim, theta, rope_scaling) angles = position_ids.astype(mx.float32)[..., None] * freqs cos = mx.cos(angles)[:, :, None, :] sin = mx.sin(angles)[:, :, None, :] def rotate(x: mx.array) -> mx.array: even = x[..., 0::2] odd = x[..., 1::2] rotated = mx.stack((even * cos - odd * sin, even * sin + odd * cos), axis=-1) return rotated.reshape(x.shape).astype(x.dtype) return rotate(queries), rotate(keys) def _attention_mask(seq_len: int, window_size: Optional[int] = None): if seq_len == 1: return None if window_size is None or seq_len <= window_size: return "causal" positions = mx.arange(seq_len) return (positions[:, None] >= positions[None, :]) & ( positions[:, None] < positions[None, :] + window_size ) class LaneformerAttention(LaneModule, nn.Module): def __init__( self, args: ModelArgs, layer_idx: int, reduce_mode: ReduceMode = ReduceMode.PRESENT, ): super().__init__( no_reduce_scale=math.sqrt(args.num_lanes), reduce_mode=reduce_mode, ) self.layer_idx = layer_idx self.n_heads = args.num_attention_heads self.num_lanes = args.num_lanes self.n_kv_heads = args.num_key_value_heads or args.num_attention_heads self.head_dim = args.hidden_size // args.num_attention_heads self.scale = self.head_dim**-0.5 self.rope_theta = args.rope_theta self.rope_scaling = args.rope_scaling self.wq = LaneRowLinear(args.hidden_size, self.n_heads * self.head_dim, args.num_lanes) self.wk = LaneRowLinear(args.hidden_size, self.n_kv_heads * self.head_dim, args.num_lanes) self.wv = LaneRowLinear(args.hidden_size, self.n_kv_heads * self.head_dim, args.num_lanes) self.wo = LaneColumnLinear(self.n_heads * self.head_dim, args.hidden_size, args.num_lanes) def __call__( self, x: mx.array, position_ids: mx.array, mask, ) -> mx.array: batch, seq_len, _, _ = x.shape queries = self.wq(x).reshape(batch, seq_len, self.n_heads, self.head_dim) keys = self.wk(x).reshape(batch, seq_len, self.n_kv_heads, self.head_dim) values = self.wv(x).reshape(batch, seq_len, self.n_kv_heads, self.head_dim) queries, keys = _apply_rotary( queries, keys, position_ids, self.rope_theta, self.rope_scaling, ) queries = queries.transpose(0, 2, 1, 3) keys = keys.transpose(0, 2, 1, 3) values = values.transpose(0, 2, 1, 3) output = mx.fast.scaled_dot_product_attention( queries, keys, values, scale=self.scale, mask=mask, ) output = output.transpose(0, 2, 1, 3).reshape( batch, seq_len, self.num_lanes, self.n_heads * self.head_dim // self.num_lanes, ) return self.wo(output) class LaneformerMLP(LaneModule, nn.Module): def __init__(self, args: ModelArgs, reduce_mode: ReduceMode = ReduceMode.PRESENT): super().__init__( no_reduce_scale=math.sqrt(args.num_lanes), reduce_mode=reduce_mode, ) self.w1 = LaneRowLinear(args.hidden_size, args.intermediate_size, args.num_lanes) self.w2 = LaneColumnLinear(args.intermediate_size, args.hidden_size, args.num_lanes) self.w3 = LaneRowLinear(args.hidden_size, args.intermediate_size, args.num_lanes) def __call__(self, x: mx.array) -> mx.array: return self.w2(nn.silu(self.w1(x)) * self.w3(x)) class LaneformerDecoderLayer(nn.Module): def __init__( self, args: ModelArgs, layer_idx: int, attention_reduce_mode: ReduceMode, mlp_reduce_mode: ReduceMode, broadcast_attention_to_future: bool, broadcast_mlp_to_future: bool, ): super().__init__() self.attention = LaneformerAttention(args, layer_idx, attention_reduce_mode) self.feed_forward = LaneformerMLP(args, mlp_reduce_mode) norm_class = ( nn.RMSNorm if args.replicated_rmsn_scale else lambda dim, eps: LaneRMSNorm(dim, args.num_lanes, eps) ) self.attention_norm = norm_class(args.hidden_size, eps=args.norm_eps) self.ffn_norm = norm_class(args.hidden_size, eps=args.norm_eps) self.broadcast_attention_to_future = broadcast_attention_to_future self.broadcast_mlp_to_future = broadcast_mlp_to_future def __call__( self, hidden_states: mx.array, position_ids: mx.array, mask, past_attention: Optional[mx.array], past_mlp: Optional[mx.array], ) -> tuple[mx.array, Optional[mx.array], Optional[mx.array]]: residual = hidden_states hidden_states = self.attention_norm(hidden_states) hidden_states = self.attention(hidden_states, position_ids, mask) future_attention = hidden_states if self.broadcast_attention_to_future else None hidden_states = self.attention.reduce_lanes(hidden_states, past_attention) hidden_states = residual + hidden_states residual = hidden_states hidden_states = self.ffn_norm(hidden_states) hidden_states = self.feed_forward(hidden_states) future_mlp = hidden_states if self.broadcast_mlp_to_future else None hidden_states = self.feed_forward.reduce_lanes(hidden_states, past_mlp) hidden_states = residual + hidden_states return hidden_states, future_attention, future_mlp class LaneformerModel(nn.Module): def __init__(self, args: ModelArgs): super().__init__() self.args = args self.num_lanes = args.num_lanes self.broadcast_delay = args.broadcast_delay self.use_attention_comm = args.use_attention_comm self.use_mlp_comm = args.use_mlp_comm self.use_early_comm = args.use_early_comm self.pre_norm_lane_agg = args.pre_norm_lane_agg self.lm_head_type = args.lm_head_type self.sliding_window = args.sliding_window self.layer_types = args.layer_types or ["full_attention"] * args.num_hidden_layers self.tok_embeddings = nn.Embedding(args.vocab_size, args.hidden_size) self.layers = [ self._make_layer(args, layer_id) for layer_id in range(args.num_hidden_layers) ] norm_class = ( nn.RMSNorm if args.replicated_rmsn_scale or args.pre_norm_lane_agg else lambda dim, eps: LaneRMSNorm(dim, args.num_lanes, eps) ) self.norm = norm_class(args.hidden_size, eps=args.norm_eps) def _make_layer(self, args: ModelArgs, layer_id: int) -> LaneformerDecoderLayer: if self.broadcast_delay == 0 or layer_id < self.broadcast_delay: attention_reduce_mode = ( ReduceMode.PRESENT if self.use_early_comm else ReduceMode.NO_REDUCE ) mlp_reduce_mode = ( ReduceMode.PRESENT if self.use_early_comm else ReduceMode.NO_REDUCE ) else: attention_reduce_mode = ( ReduceMode.NO_REDUCE if not self.use_attention_comm else ReduceMode.PAST ) mlp_reduce_mode = ( ReduceMode.NO_REDUCE if not self.use_mlp_comm else ReduceMode.PAST ) broadcast_attention_to_future = ( args.num_hidden_layers - self.broadcast_delay > layer_id ) and self.use_attention_comm broadcast_mlp_to_future = ( args.num_hidden_layers - self.broadcast_delay > layer_id ) and self.use_mlp_comm return LaneformerDecoderLayer( args=args, layer_idx=layer_id, attention_reduce_mode=attention_reduce_mode, mlp_reduce_mode=mlp_reduce_mode, broadcast_attention_to_future=broadcast_attention_to_future, broadcast_mlp_to_future=broadcast_mlp_to_future, ) def __call__( self, inputs: mx.array, input_embeddings: Optional[mx.array] = None, ) -> mx.array: if input_embeddings is None: hidden_states = self.tok_embeddings(inputs) else: hidden_states = input_embeddings batch, seq_len, _ = hidden_states.shape position_ids = mx.broadcast_to(mx.arange(seq_len)[None, :], (batch, seq_len)) hidden_states = mx.broadcast_to( hidden_states[:, :, None, :], (batch, seq_len, self.num_lanes, self.args.hidden_size), ) masks = { "full_attention": _attention_mask(seq_len), } if "sliding_attention" in self.layer_types: masks["sliding_attention"] = _attention_mask(seq_len, self.sliding_window) past_attentions = [] past_mlps = [] for i, layer in enumerate(self.layers): if self.broadcast_delay == 0 or i < self.broadcast_delay: past_attention = None past_mlp = None else: past_attention = past_attentions[i - self.broadcast_delay] past_mlp = past_mlps[i - self.broadcast_delay] hidden_states, future_attention, future_mlp = layer( hidden_states, position_ids, masks[self.layer_types[i]], past_attention, past_mlp, ) past_attentions.append(future_attention) past_mlps.append(future_mlp) if self.pre_norm_lane_agg: hidden_states = mx.sum(hidden_states, axis=-2) return self.norm(hidden_states) if self.lm_head_type == "replicate": hidden_states = self.norm(hidden_states) return mx.mean(hidden_states, axis=-2) return self.norm(hidden_states) class Model(nn.Module): def __init__(self, args: ModelArgs): super().__init__() if args.tie_word_embeddings: raise ValueError("Laneformer does not support tied input/output embeddings.") self.args = args self.model_type = args.model_type self.model = LaneformerModel(args) self.lm_head_type = args.lm_head_type if self.lm_head_type == "replicate": self.lm_head = nn.Linear(args.hidden_size, args.vocab_size, bias=False) elif self.lm_head_type == "lane": self.lm_head = LaneLMHead(args.hidden_size, args.vocab_size, args.num_lanes) elif self.lm_head_type == "vocab_parallel": self.lm_head = LaneRowLinear(args.hidden_size, args.vocab_size, args.num_lanes) else: raise ValueError(f"Unsupported lm_head_type: {self.lm_head_type!r}") def __call__( self, inputs: mx.array, cache: Optional[list[Any]] = None, input_embeddings: Optional[mx.array] = None, ) -> mx.array: if cache is not None and cache and isinstance(cache[0], HistoryCache): if input_embeddings is not None: raise ValueError("HistoryCache does not support input_embeddings.") inputs = cache[0].update_and_fetch(inputs) hidden_states = self.model(inputs, input_embeddings=input_embeddings) if self.lm_head_type == "replicate": return self.lm_head(hidden_states) logits = self.lm_head(hidden_states) if self.lm_head_type == "vocab_parallel": batch, seq_len, lanes, dim = logits.shape return logits.reshape(batch, seq_len, lanes * dim) return logits def make_cache(self) -> list[HistoryCache]: return [HistoryCache()] def sanitize(self, weights: dict[str, mx.array]) -> dict[str, mx.array]: return { key: value for key, value in weights.items() if "rotary_emb.inv_freq" not in key }