Instructions to use Blankyy/needle-mlx with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use Blankyy/needle-mlx with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir needle-mlx Blankyy/needle-mlx
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
| import json | |
| import math | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| import mlx.core as mx | |
| import mlx.nn as nn | |
| from mlx.utils import tree_flatten | |
| def default_init(): | |
| return nn.init.normal(std=0.02) | |
| def residual_init(num_layers): | |
| return nn.init.normal(std=0.02 / math.sqrt(2 * num_layers)) | |
| DTYPE_MAP = { | |
| "float32": mx.float32, | |
| "bfloat16": mx.bfloat16, | |
| "float16": mx.float16, | |
| } | |
| def _map_hf_weight(name): | |
| if name == "model.embed_tokens.weight": | |
| return "embedding.weight" | |
| if name == "model.encoder.final_norm.weight": | |
| return "encoder.final_norm.scale" | |
| if name == "model.decoder.norm.weight": | |
| return "decoder.final_norm.scale" | |
| if name == "lm_head.weight": | |
| return None | |
| if name.startswith("model.encoder.layers."): | |
| name = name.removeprefix("model.") | |
| name = name.replace(".input_layernorm.weight", ".norm.scale") | |
| elif name.startswith("model.decoder.layers."): | |
| name = name.removeprefix("model.") | |
| name = name.replace(".input_layernorm.weight", ".self_norm.scale") | |
| name = name.replace( | |
| ".encoder_attn_layer_norm.weight", ".cross_norm.scale" | |
| ) | |
| name = name.replace(".encoder_attn.", ".cross_attn.") | |
| else: | |
| raise ValueError(f"Unsupported checkpoint tensor: {name}") | |
| return name.replace("_norm.weight", "_norm.scale") | |
| def _linear(in_dims, out_dims, dtype, *, bias=False, init=None): | |
| layer = nn.Linear(in_dims, out_dims, bias=bias) | |
| layer.weight = (init or default_init())(layer.weight) | |
| if bias: | |
| layer.bias = mx.zeros_like(layer.bias) | |
| layer.set_dtype(dtype) | |
| return layer | |
| def _dropout(x, rate, deterministic): | |
| if deterministic or rate == 0: | |
| return x | |
| keep = 1 - rate | |
| return mx.random.bernoulli(keep, x.shape) * x / keep | |
| class ZCRMSNorm(nn.Module): | |
| def __init__(self, dims, dtype=mx.bfloat16, epsilon=1e-6): | |
| super().__init__() | |
| self.scale = mx.zeros((dims,)) | |
| self.dtype = dtype | |
| self.epsilon = epsilon | |
| def __call__(self, x): | |
| rms = mx.sqrt( | |
| mx.mean(x.astype(mx.float32) ** 2, axis=-1, keepdims=True) | |
| + self.epsilon | |
| ) | |
| return ((1 + self.scale) * x / rms).astype(self.dtype) | |
| class TransformerConfig: | |
| vocab_size: int = 8192 | |
| d_model: int = 128 | |
| num_heads: int = 4 | |
| num_kv_heads: int = 2 | |
| num_encoder_layers: int = 2 | |
| num_decoder_layers: int = 2 | |
| d_ff: int = 512 | |
| max_seq_len: int = 128 | |
| pad_token_id: int = 0 | |
| rope_theta: float = 10_000.0 | |
| dtype: str = "bfloat16" | |
| activation: str = "drelu" | |
| num_memory_slots: int = 64 | |
| dropout_rate: float = 0.1 | |
| contrastive_dim: int = 128 | |
| no_feedforward: bool = True | |
| def __init__(self, **kwargs): | |
| valid = self.__dataclass_fields__ | |
| for key, value in kwargs.items(): | |
| if key in valid: | |
| setattr(self, key, value) | |
| def mlx_dtype(self): | |
| return DTYPE_MAP[self.dtype] | |
| def total_layers(self): | |
| return self.num_encoder_layers + self.num_decoder_layers | |
| def precompute_rope_freqs(head_dim, seq_len, theta=10_000.0): | |
| freqs = 1 / theta ** (mx.arange(0, head_dim, 2, dtype=mx.float32) / head_dim) | |
| angles = mx.outer(mx.arange(seq_len, dtype=mx.float32), freqs) | |
| return mx.cos(angles), mx.sin(angles) | |
| def apply_rope(x, cos, sin): | |
| length = x.shape[2] | |
| half = x.shape[-1] // 2 | |
| cos = cos[:length][None, None, :, :] | |
| sin = sin[:length][None, None, :, :] | |
| x1, x2 = x[..., :half], x[..., half:] | |
| return mx.concatenate((x1 * cos - x2 * sin, x2 * cos + x1 * sin), axis=-1) | |
| class MultiHeadAttention(nn.Module): | |
| def __init__( | |
| self, | |
| num_heads, | |
| num_kv_heads, | |
| d_model, | |
| num_layers, | |
| dtype=mx.bfloat16, | |
| rope_keys_only=False, | |
| ): | |
| super().__init__() | |
| if d_model % num_heads or num_heads % num_kv_heads: | |
| raise ValueError( | |
| "d_model must be divisible by num_heads, and num_heads by num_kv_heads" | |
| ) | |
| self.num_heads = num_heads | |
| self.num_kv_heads = num_kv_heads | |
| self.d_model = d_model | |
| self.head_dim = d_model // num_heads | |
| self.repeats = num_heads // num_kv_heads | |
| self.dtype = dtype | |
| self.rope_keys_only = rope_keys_only | |
| kv_dim = num_kv_heads * self.head_dim | |
| self.q_proj = _linear(d_model, d_model, dtype) | |
| self.k_proj = _linear(d_model, kv_dim, dtype) | |
| self.v_proj = _linear(d_model, kv_dim, dtype) | |
| self.out_proj = _linear( | |
| d_model, d_model, dtype, init=residual_init(num_layers) | |
| ) | |
| self.q_norm = ZCRMSNorm(self.head_dim, dtype) | |
| self.k_norm = ZCRMSNorm(self.head_dim, dtype) | |
| def __call__(self, q_input, kv_input, mask=None, rope=None): | |
| batch, q_len, _ = q_input.shape | |
| kv_len = kv_input.shape[1] | |
| q = self.q_proj(q_input.astype(self.dtype)) | |
| k = self.k_proj(kv_input.astype(self.dtype)) | |
| v = self.v_proj(kv_input.astype(self.dtype)) | |
| q = q.reshape(batch, q_len, self.num_heads, self.head_dim).transpose( | |
| 0, 2, 1, 3 | |
| ) | |
| k = k.reshape(batch, kv_len, self.num_kv_heads, self.head_dim).transpose( | |
| 0, 2, 1, 3 | |
| ) | |
| v = v.reshape(batch, kv_len, self.num_kv_heads, self.head_dim).transpose( | |
| 0, 2, 1, 3 | |
| ) | |
| q = self.q_norm(q) | |
| k = self.k_norm(k) | |
| if self.repeats > 1: | |
| k = mx.repeat(k, self.repeats, axis=1) | |
| v = mx.repeat(v, self.repeats, axis=1) | |
| if rope is not None: | |
| cos, sin = rope | |
| if not self.rope_keys_only: | |
| q = apply_rope(q, cos, sin) | |
| k = apply_rope(k, cos, sin) | |
| scale = mx.sqrt(mx.array(self.head_dim, dtype=mx.float32)) | |
| weights = mx.matmul(q, k.transpose(0, 1, 3, 2)) / scale | |
| if mask is not None: | |
| weights = mx.where(mask, weights, mx.finfo(weights.dtype).min) | |
| weights = nn.softmax(weights, axis=-1) | |
| out = mx.matmul(weights, v) | |
| out = out.transpose(0, 2, 1, 3).reshape(batch, q_len, self.d_model) | |
| return self.out_proj(out.astype(self.dtype)) | |
| class FeedForward(nn.Module): | |
| def __init__( | |
| self, d_model, d_ff, num_layers, dtype=mx.bfloat16, activation="drelu" | |
| ): | |
| super().__init__() | |
| self.dtype = dtype | |
| self.activation = activation | |
| self.gate_proj = _linear(d_model, d_ff, dtype) | |
| self.up_proj = _linear(d_model, d_ff, dtype) | |
| self.down_proj = _linear( | |
| d_ff, d_model, dtype, init=residual_init(num_layers) | |
| ) | |
| def __call__(self, x, ffn_mask=None): | |
| gate = self.gate_proj(x.astype(self.dtype)) | |
| up = self.up_proj(x.astype(self.dtype)) | |
| if self.activation == "swiglu": | |
| h = nn.silu(gate) * up | |
| elif self.activation == "geglu": | |
| h = nn.gelu_approx(gate) * up | |
| else: | |
| h = nn.relu(gate) * nn.relu(up) | |
| if ffn_mask is not None: | |
| h = h * ffn_mask[:, None, :] | |
| return self.down_proj(h.astype(self.dtype)) | |
| class EncoderBlock(nn.Module): | |
| def __init__( | |
| self, | |
| num_heads, | |
| num_kv_heads, | |
| d_model, | |
| d_ff, | |
| num_layers, | |
| dtype=mx.bfloat16, | |
| activation="drelu", | |
| dropout_rate=0.0, | |
| no_feedforward=True, | |
| ): | |
| super().__init__() | |
| self.dtype = dtype | |
| self.dropout_rate = dropout_rate | |
| self.no_feedforward = no_feedforward | |
| self.attn_gate = mx.zeros(()) | |
| self.norm = ZCRMSNorm(d_model, dtype) | |
| self.self_attn = MultiHeadAttention( | |
| num_heads, num_kv_heads, d_model, num_layers, dtype | |
| ) | |
| if not no_feedforward: | |
| self.ffn_gate = mx.zeros(()) | |
| self.ffn_norm = ZCRMSNorm(d_model, dtype) | |
| self.ffn = FeedForward( | |
| d_model, d_ff, num_layers, dtype, activation | |
| ) | |
| def __call__( | |
| self, x, mask=None, rope=None, ffn_mask=None, deterministic=True | |
| ): | |
| residual = x | |
| normed = self.norm(x) | |
| x = self.self_attn(normed, normed, mask=mask, rope=rope) | |
| gate = nn.sigmoid(self.attn_gate).astype(self.dtype) | |
| x = residual + gate * _dropout(x, self.dropout_rate, deterministic) | |
| if not self.no_feedforward: | |
| residual = x | |
| x = self.ffn(self.ffn_norm(x), ffn_mask=ffn_mask) | |
| gate = nn.sigmoid(self.ffn_gate).astype(self.dtype) | |
| x = residual + gate * _dropout(x, self.dropout_rate, deterministic) | |
| return x | |
| class Encoder(nn.Module): | |
| def __init__(self, config): | |
| super().__init__() | |
| self.config = config | |
| self.layers = [ | |
| EncoderBlock( | |
| config.num_heads, | |
| config.num_kv_heads, | |
| config.d_model, | |
| config.d_ff, | |
| config.total_layers, | |
| config.mlx_dtype, | |
| config.activation, | |
| config.dropout_rate, | |
| config.no_feedforward, | |
| ) | |
| for _ in range(config.num_encoder_layers) | |
| ] | |
| self.final_norm = ZCRMSNorm(config.d_model, config.mlx_dtype) | |
| def __call__( | |
| self, x, mask=None, rope=None, ffn_mask=None, deterministic=True | |
| ): | |
| x = x.astype(self.config.mlx_dtype) | |
| for layer in self.layers: | |
| x = layer(x, mask, rope, ffn_mask, deterministic) | |
| return self.final_norm(x), mask | |
| class DecoderBlock(nn.Module): | |
| def __init__( | |
| self, | |
| num_heads, | |
| num_kv_heads, | |
| d_model, | |
| d_ff, | |
| num_layers, | |
| dtype=mx.bfloat16, | |
| activation="drelu", | |
| dropout_rate=0.0, | |
| no_feedforward=True, | |
| ): | |
| super().__init__() | |
| self.dtype = dtype | |
| self.dropout_rate = dropout_rate | |
| self.no_feedforward = no_feedforward | |
| self.self_attn_gate = mx.zeros(()) | |
| self.self_norm = ZCRMSNorm(d_model, dtype) | |
| self.self_attn = MultiHeadAttention( | |
| num_heads, num_kv_heads, d_model, num_layers, dtype | |
| ) | |
| self.cross_attn_gate = mx.zeros(()) | |
| self.cross_norm = ZCRMSNorm(d_model, dtype) | |
| self.cross_attn = MultiHeadAttention( | |
| num_heads, num_kv_heads, d_model, num_layers, dtype | |
| ) | |
| if not no_feedforward: | |
| self.ffn_gate = mx.zeros(()) | |
| self.ffn_norm = ZCRMSNorm(d_model, dtype) | |
| self.ffn = FeedForward( | |
| d_model, d_ff, num_layers, dtype, activation | |
| ) | |
| def __call__( | |
| self, | |
| x, | |
| encoder_out, | |
| self_mask=None, | |
| cross_mask=None, | |
| rope=None, | |
| ffn_mask=None, | |
| deterministic=True, | |
| ): | |
| residual = x | |
| normed = self.self_norm(x) | |
| x = self.self_attn(normed, normed, mask=self_mask, rope=rope) | |
| gate = nn.sigmoid(self.self_attn_gate).astype(self.dtype) | |
| x = residual + gate * _dropout(x, self.dropout_rate, deterministic) | |
| residual = x | |
| x = self.cross_attn( | |
| self.cross_norm(x), encoder_out, mask=cross_mask | |
| ) | |
| gate = nn.sigmoid(self.cross_attn_gate).astype(self.dtype) | |
| x = residual + gate * _dropout(x, self.dropout_rate, deterministic) | |
| if not self.no_feedforward: | |
| residual = x | |
| x = self.ffn(self.ffn_norm(x), ffn_mask=ffn_mask) | |
| gate = nn.sigmoid(self.ffn_gate).astype(self.dtype) | |
| x = residual + gate * _dropout(x, self.dropout_rate, deterministic) | |
| return x | |
| class Decoder(nn.Module): | |
| def __init__(self, config): | |
| super().__init__() | |
| self.config = config | |
| self.layers = [ | |
| DecoderBlock( | |
| config.num_heads, | |
| config.num_kv_heads, | |
| config.d_model, | |
| config.d_ff, | |
| config.total_layers, | |
| config.mlx_dtype, | |
| config.activation, | |
| config.dropout_rate, | |
| config.no_feedforward, | |
| ) | |
| for _ in range(config.num_decoder_layers) | |
| ] | |
| self.final_norm = ZCRMSNorm(config.d_model, config.mlx_dtype) | |
| def __call__( | |
| self, | |
| x, | |
| encoder_out, | |
| self_mask=None, | |
| cross_mask=None, | |
| rope=None, | |
| ffn_mask=None, | |
| deterministic=True, | |
| ): | |
| x = x.astype(self.config.mlx_dtype) | |
| for layer in self.layers: | |
| x = layer( | |
| x, | |
| encoder_out, | |
| self_mask, | |
| cross_mask, | |
| rope, | |
| ffn_mask, | |
| deterministic, | |
| ) | |
| return self.final_norm(x) | |
| class SimpleAttentionNetwork(nn.Module): | |
| def __init__(self, config): | |
| super().__init__() | |
| self.config = config | |
| self.embedding = nn.Embedding(config.vocab_size, config.d_model) | |
| self.embedding.weight = default_init()(self.embedding.weight) | |
| self.embed_scale = math.sqrt(config.d_model) | |
| self.encoder = Encoder(config) | |
| self.decoder = Decoder(config) | |
| self.contrastive_hidden = _linear( | |
| config.d_model, | |
| config.d_model // 4, | |
| config.mlx_dtype, | |
| bias=True, | |
| ) | |
| self.contrastive_proj = _linear( | |
| config.d_model // 4, config.contrastive_dim, config.mlx_dtype | |
| ) | |
| self.log_temp = mx.zeros(()) | |
| def from_pretrained(cls, model_dir="weights"): | |
| model_dir = Path(model_dir) | |
| with (model_dir / "config.json").open() as file: | |
| hf_config = json.load(file) | |
| config = TransformerConfig( | |
| vocab_size=hf_config["vocab_size"], | |
| d_model=hf_config["d_model"], | |
| num_heads=hf_config["num_heads"], | |
| num_kv_heads=hf_config["num_kv_heads"], | |
| num_encoder_layers=hf_config["num_encoder_layers"], | |
| num_decoder_layers=hf_config["num_decoder_layers"], | |
| pad_token_id=hf_config["pad_token_id"], | |
| rope_theta=hf_config["rope_theta"], | |
| dtype=hf_config.get("torch_dtype", "bfloat16"), | |
| no_feedforward=True, | |
| ) | |
| model = cls(config) | |
| checkpoint = mx.load(str(model_dir / "model.safetensors")) | |
| if "model.embed_tokens.weight" in checkpoint: | |
| if not mx.array_equal( | |
| checkpoint["model.embed_tokens.weight"], | |
| checkpoint["lm_head.weight"], | |
| ).item(): | |
| raise ValueError("Checkpoint input and output embeddings are not tied") | |
| mapped = [] | |
| for name, value in checkpoint.items(): | |
| target = _map_hf_weight(name) | |
| if target is not None: | |
| if target.endswith("_gate"): | |
| value = value.reshape(()) | |
| mapped.append((target, value)) | |
| else: | |
| mapped = list(checkpoint.items()) | |
| model.load_weights(mapped, strict=False) | |
| mx.eval(model.parameters()) | |
| return model | |
| def generate(self, src, max_new_tokens=512): | |
| """Greedy generation. Decoder starts with EOS, as in the JAX runtime.""" | |
| if src.shape[0] != 1: | |
| raise ValueError("generate currently supports batch size 1") | |
| encoder_out, enc_mask = self.encode_text( | |
| src, src_mask=make_padding_mask(src, self.config.pad_token_id) | |
| ) | |
| tokens = mx.full( | |
| (1, max_new_tokens + 1), | |
| self.config.pad_token_id, | |
| dtype=mx.int32, | |
| ) | |
| tokens[0, 0] = 1 | |
| causal = make_causal_mask(tokens.shape[1]) | |
| generated = [] | |
| # ponytail: full-buffer decode matches upstream; add a KV cache if latency matters. | |
| for position in range(max_new_tokens): | |
| logits = self.decode( | |
| tokens, | |
| encoder_out, | |
| self_mask=causal, | |
| cross_mask=enc_mask, | |
| ) | |
| token = int(mx.argmax(logits[0, position]).item()) | |
| if token == 1: | |
| break | |
| generated.append(token) | |
| tokens[0, position + 1] = token | |
| return generated | |
| def _rope(self, seq_len): | |
| head_dim = self.config.d_model // self.config.num_heads | |
| return precompute_rope_freqs( | |
| head_dim, seq_len, self.config.rope_theta | |
| ) | |
| def encode_text( | |
| self, src, src_mask=None, ffn_mask=None, deterministic=True | |
| ): | |
| x = self.embedding(src) * self.embed_scale | |
| return self.encoder( | |
| x, | |
| mask=src_mask, | |
| rope=self._rope(src.shape[1]), | |
| ffn_mask=ffn_mask, | |
| deterministic=deterministic, | |
| ) | |
| def encode(self, src, src_mask=None): | |
| return self.encode_text(src, src_mask=src_mask) | |
| def decode( | |
| self, | |
| tgt, | |
| encoder_out, | |
| self_mask=None, | |
| cross_mask=None, | |
| deterministic=True, | |
| ): | |
| x = self.embedding(tgt) * self.embed_scale | |
| x = self.decoder( | |
| x, | |
| encoder_out, | |
| self_mask=self_mask, | |
| cross_mask=cross_mask, | |
| rope=self._rope(tgt.shape[1]), | |
| deterministic=deterministic, | |
| ) | |
| return self.embedding.as_linear(x.astype(mx.float32)) | |
| def _mean_pool(self, encoder_out, enc_mask): | |
| if enc_mask is not None: | |
| mask_2d = enc_mask[:, 0, 0, :] | |
| else: | |
| mask_2d = mx.ones(encoder_out.shape[:2], dtype=encoder_out.dtype) | |
| mask_3d = mask_2d[:, :, None].astype(encoder_out.dtype) | |
| summed = mx.sum(encoder_out * mask_3d, axis=1) | |
| counts = mx.maximum(mx.sum(mask_2d, axis=1, keepdims=True), 1.0) | |
| return summed / counts | |
| def encode_contrastive(self, tokens, deterministic=True): | |
| src_mask = make_padding_mask(tokens, self.config.pad_token_id) | |
| encoder_out, enc_mask = self.encode_text( | |
| tokens, src_mask=src_mask, deterministic=deterministic | |
| ) | |
| pooled = self._mean_pool(encoder_out, enc_mask) | |
| projected = self.contrastive_proj( | |
| nn.relu(self.contrastive_hidden(pooled)) | |
| ) | |
| denom = mx.sqrt( | |
| mx.sum( | |
| projected.astype(mx.float32) ** 2, | |
| axis=-1, | |
| keepdims=True, | |
| ) | |
| + 1e-12 | |
| ) | |
| return projected / denom.astype(projected.dtype) | |
| def forward_contrastive( | |
| self, query_tokens, tool_tokens, deterministic=True | |
| ): | |
| q_emb = self.encode_contrastive(query_tokens, deterministic) | |
| t_emb = self.encode_contrastive(tool_tokens, deterministic) | |
| return q_emb, t_emb, self.log_temp | |
| def __call__( | |
| self, src, tgt, src_mask=None, tgt_mask=None, cross_mask=None | |
| ): | |
| encoder_out, enc_mask = self.encode_text(src, src_mask=src_mask) | |
| return self.decode( | |
| tgt, | |
| encoder_out, | |
| self_mask=tgt_mask, | |
| cross_mask=cross_mask if cross_mask is not None else enc_mask, | |
| ) | |
| def _run_decoder( | |
| self, | |
| encoder_out, | |
| tgt, | |
| tgt_mask=None, | |
| cross_mask=None, | |
| ffn_mask=None, | |
| deterministic=True, | |
| ): | |
| x = self.embedding(tgt) * self.embed_scale | |
| x = self.decoder( | |
| x, | |
| encoder_out, | |
| self_mask=tgt_mask, | |
| cross_mask=cross_mask, | |
| rope=self._rope(tgt.shape[1]), | |
| ffn_mask=ffn_mask, | |
| deterministic=deterministic, | |
| ) | |
| return x.astype(mx.float32) | |
| def forward_masked( | |
| self, | |
| src, | |
| tgt, | |
| src_mask=None, | |
| tgt_mask=None, | |
| cross_mask=None, | |
| ffn_mask=None, | |
| deterministic=True, | |
| ): | |
| encoder_out, enc_mask = self.encode_text( | |
| src, | |
| src_mask=src_mask, | |
| ffn_mask=ffn_mask, | |
| deterministic=deterministic, | |
| ) | |
| x = self._run_decoder( | |
| encoder_out, | |
| tgt, | |
| tgt_mask=tgt_mask, | |
| cross_mask=cross_mask if cross_mask is not None else enc_mask, | |
| ffn_mask=ffn_mask, | |
| deterministic=deterministic, | |
| ) | |
| return self.embedding.as_linear(x), 0.0 | |
| def _make_eval_ffn_mask(self, ff_width, batch, dtype): | |
| mask = (mx.arange(self.config.d_ff) < ff_width).astype(dtype) | |
| return mx.broadcast_to(mask[None, :], (batch, self.config.d_ff)) | |
| def forward_with_aux( | |
| self, | |
| src, | |
| tgt, | |
| src_mask=None, | |
| tgt_mask=None, | |
| cross_mask=None, | |
| mat_ff_widths=None, | |
| ): | |
| batch = src.shape[0] | |
| encoder_out, enc_mask = self.encode_text(src, src_mask=src_mask) | |
| cm = cross_mask if cross_mask is not None else enc_mask | |
| x = self._run_decoder( | |
| encoder_out, tgt, tgt_mask=tgt_mask, cross_mask=cm | |
| ) | |
| logits = self.embedding.as_linear(x) | |
| mat_logits = [] | |
| if mat_ff_widths is not None: | |
| for width in mat_ff_widths: | |
| mask = self._make_eval_ffn_mask(width, batch, x.dtype) | |
| enc_m, _ = self.encode_text( | |
| src, src_mask=src_mask, ffn_mask=mask | |
| ) | |
| x_m = self._run_decoder( | |
| enc_m, | |
| tgt, | |
| tgt_mask=tgt_mask, | |
| cross_mask=cm, | |
| ffn_mask=mask, | |
| ) | |
| mat_logits.append(self.embedding.as_linear(x_m)) | |
| return logits, 0.0, mat_logits | |
| def init_all(self, src, tgt): | |
| src_mask = make_padding_mask(src, self.config.pad_token_id) | |
| tgt_mask = make_causal_mask(tgt.shape[1]) & make_padding_mask( | |
| tgt, self.config.pad_token_id | |
| ) | |
| encoder_out, enc_mask = self.encode_text(src, src_mask=src_mask) | |
| self._run_decoder( | |
| encoder_out, tgt, tgt_mask=tgt_mask, cross_mask=enc_mask | |
| ) | |
| self.encode_contrastive(src) | |
| return mx.zeros(()) | |
| def make_causal_mask(seq_len): | |
| return mx.tril(mx.ones((seq_len, seq_len), dtype=mx.bool_))[None, None] | |
| def make_padding_mask(tokens, pad_token_id): | |
| return (tokens != pad_token_id)[:, None, None, :] | |
| def make_packing_mask(seg_ids): | |
| mask = (seg_ids[:, :, None] == seg_ids[:, None, :]) & ( | |
| seg_ids[:, :, None] > 0 | |
| ) | |
| return mask[:, None, :, :] | |
| def make_causal_packing_mask(seg_ids): | |
| length = seg_ids.shape[1] | |
| causal = mx.tril(mx.ones((length, length), dtype=mx.bool_)) | |
| block = (seg_ids[:, :, None] == seg_ids[:, None, :]) & ( | |
| seg_ids[:, :, None] > 0 | |
| ) | |
| return (block & causal[None, :, :])[:, None, :, :] | |
| def make_cross_packing_mask(enc_seg_ids, dec_seg_ids): | |
| mask = (dec_seg_ids[:, :, None] == enc_seg_ids[:, None, :]) & ( | |
| dec_seg_ids[:, :, None] > 0 | |
| ) | |
| return mask[:, None, :, :] | |
| def _demo(): | |
| config = TransformerConfig( | |
| vocab_size=32, | |
| d_model=16, | |
| num_heads=4, | |
| num_kv_heads=2, | |
| num_encoder_layers=1, | |
| num_decoder_layers=1, | |
| d_ff=32, | |
| contrastive_dim=8, | |
| dropout_rate=0.0, | |
| no_feedforward=False, | |
| dtype="float32", | |
| ) | |
| model = SimpleAttentionNetwork(config) | |
| src = mx.array([[1, 2, 3, 0], [4, 5, 0, 0]]) | |
| tgt = mx.array([[1, 6, 7, 0], [1, 8, 0, 0]]) | |
| src_mask = make_padding_mask(src, config.pad_token_id) | |
| tgt_mask = make_causal_mask(tgt.shape[1]) & make_padding_mask( | |
| tgt, config.pad_token_id | |
| ) | |
| logits = model(src, tgt, src_mask, tgt_mask) | |
| q_emb, t_emb, _ = model.forward_contrastive(src, tgt) | |
| aux_logits, _, widths = model.forward_with_aux( | |
| src, tgt, src_mask, tgt_mask, mat_ff_widths=[16] | |
| ) | |
| mx.eval(logits, q_emb, t_emb, aux_logits, widths[0]) | |
| assert logits.shape == (2, 4, 32) | |
| assert q_emb.shape == t_emb.shape == (2, 8) | |
| assert aux_logits.shape == widths[0].shape == logits.shape | |
| assert mx.all(mx.isfinite(logits)).item() | |
| def _test_pretrained(model_dir): | |
| model = SimpleAttentionNetwork.from_pretrained(model_dir) | |
| src = mx.array([[2, 42, 314, 1]]) | |
| tgt = mx.array([[1, 27, 1]]) | |
| logits = model( | |
| src, | |
| tgt, | |
| make_padding_mask(src, model.config.pad_token_id), | |
| make_causal_mask(tgt.shape[1]), | |
| ) | |
| mx.eval(logits) | |
| assert logits.shape == (1, 3, model.config.vocab_size) | |
| assert mx.all(mx.isfinite(logits)).item() | |
| count = sum(value.size for _, value in tree_flatten(model.parameters())) | |
| print(f"model parameters: {count:,}") | |
| print(f"logits: shape={logits.shape}, dtype={logits.dtype}") | |
| if __name__ == "__main__": | |
| import sys | |
| _test_pretrained(sys.argv[1]) if len(sys.argv) > 1 else _demo() | |