laya-intent-router-150m-onnx / laya_intent_router_mlx.py
vrajnotviraj's picture
Add optional MLX backend for Apple Silicon (backend='mlx' or 'auto'; ONNX stays the default)
cb7816a verified
Raw History Blame Contribute Delete
7.13 kB
"""MLX backend for laya_intent_router.py: the same model running on the Apple Silicon GPU.
Used by LayaIntentRouter(backend="mlx" or "auto"). Needs `pip install mlx` and the weights in `mlx/`.
Encoder: ModernBERT (Ettin-150M). Pre-norm layers (layer 0 has no attn_norm), bias-free Linear and
LayerNorm, GeGLU MLP with exact GELU, rotate-half RoPE with a theta per layer type, global attention
on every 3rd layer and a |i - j| <= local_attention // 2 window on the rest. Head: 2 norm-first
transformer layers (ReLU), then a scorer on each [MASK] marker. Weights are stored in fp16; the Linear
layers are quantized to 8 bits at load when mlx/config.json asks for it.
"""
import json
import os
import mlx.core as mx
import mlx.nn as nn
import numpy as np
def _mask(attention_mask, window=None, L=None):
m = None
if window is not None:
i = mx.arange(L)
m = (mx.abs(i[:, None] - i[None, :]) <= window)[None, None]
if attention_mask is not None:
pad = attention_mask.astype(mx.bool_)[:, None, None, :]
m = pad if m is None else m & pad
return m
class Attention(nn.Module):
def __init__(self, d, heads, theta):
super().__init__()
self.heads, self.theta = heads, theta
self.Wqkv = nn.Linear(d, 3 * d, bias=False)
self.Wo = nn.Linear(d, d, bias=False)
def __call__(self, x, mask):
B, L, d = x.shape
q, k, v = self.Wqkv(x).reshape(B, L, 3, self.heads, -1).transpose(2, 0, 3, 1, 4)
q = mx.fast.rope(q, q.shape[-1], traditional=False, base=self.theta, scale=1.0, offset=0)
k = mx.fast.rope(k, k.shape[-1], traditional=False, base=self.theta, scale=1.0, offset=0)
o = mx.fast.scaled_dot_product_attention(q, k, v, scale=q.shape[-1] ** -0.5, mask=mask)
return self.Wo(o.transpose(0, 2, 1, 3).reshape(B, L, d))
class MLP(nn.Module):
def __init__(self, d, inter):
super().__init__()
self.Wi = nn.Linear(d, 2 * inter, bias=False)
self.Wo = nn.Linear(inter, d, bias=False)
def __call__(self, x):
a, g = mx.split(self.Wi(x), 2, axis=-1)
return self.Wo(nn.gelu(a) * g)
class EncoderLayer(nn.Module):
def __init__(self, c, idx):
super().__init__()
d, eps = c["hidden_size"], c["norm_eps"]
self.sliding = c["layer_types"][idx] == "sliding_attention"
theta = c["rope_parameters"]["sliding_attention" if self.sliding else "full_attention"]["rope_theta"]
if idx:
self.attn_norm = nn.LayerNorm(d, eps=eps, bias=False)
self.attn = Attention(d, c["num_attention_heads"], theta)
self.mlp_norm = nn.LayerNorm(d, eps=eps, bias=False)
self.mlp = MLP(d, c["intermediate_size"])
def __call__(self, x, mask):
x = x + self.attn(self.attn_norm(x) if "attn_norm" in self else x, mask)
return x + self.mlp(self.mlp_norm(x))
class Embeddings(nn.Module):
def __init__(self, c):
super().__init__()
self.tok_embeddings = nn.Embedding(c["vocab_size"], c["hidden_size"])
self.norm = nn.LayerNorm(c["hidden_size"], eps=c["norm_eps"], bias=False)
def __call__(self, ids):
return self.norm(self.tok_embeddings(ids))
class ModernBert(nn.Module):
def __init__(self, c):
super().__init__()
self.window = c["local_attention"] // 2
self.embeddings = Embeddings(c)
self.layers = [EncoderLayer(c, i) for i in range(c["num_hidden_layers"])]
self.final_norm = nn.LayerNorm(c["hidden_size"], eps=c["norm_eps"], bias=False)
def __call__(self, ids, attention_mask=None):
x = self.embeddings(ids)
full, local = _mask(attention_mask), _mask(attention_mask, self.window, ids.shape[1])
for layer in self.layers:
x = layer(x, local if layer.sliding else full)
return self.final_norm(x)
class MultiheadAttention(nn.Module):
def __init__(self, d, heads):
super().__init__()
self.heads = heads
self.in_proj_weight, self.in_proj_bias = mx.zeros((3 * d, d)), mx.zeros((3 * d,))
self.out_proj = nn.Linear(d, d)
def __call__(self, x, mask):
B, L, d = x.shape
qkv = x @ self.in_proj_weight.T + self.in_proj_bias
q, k, v = qkv.reshape(B, L, 3, self.heads, -1).transpose(2, 0, 3, 1, 4)
o = mx.fast.scaled_dot_product_attention(q, k, v, scale=q.shape[-1] ** -0.5, mask=mask)
return self.out_proj(o.transpose(0, 2, 1, 3).reshape(B, L, d))
class HeadLayer(nn.Module):
def __init__(self, d, heads):
super().__init__()
self.self_attn = MultiheadAttention(d, heads)
self.linear1, self.linear2 = nn.Linear(d, 4 * d), nn.Linear(4 * d, d)
self.norm1, self.norm2 = nn.LayerNorm(d), nn.LayerNorm(d)
def __call__(self, x, mask):
x = x + self.self_attn(self.norm1(x), mask)
return x + self.linear2(nn.relu(self.linear1(self.norm2(x))))
class Head(nn.Module):
def __init__(self, d, n):
super().__init__()
self.layers = [HeadLayer(d, max(1, d // 64)) for _ in range(n)]
class DecisionModel(nn.Module):
def __init__(self, c):
super().__init__()
d = c["hidden_size"]
self.encoder = ModernBert(c)
self.head = Head(d, c["head_layers"])
self.type_emb = nn.Embedding(3, d)
self.scorer = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Linear(d, 1))
# the act head ships with the weights but routing never reads it
self.act_head = nn.Sequential(nn.Linear(d + 4, 256), nn.GELU(), nn.Linear(256, c["n_act"]))
self.temperature = mx.ones((3,))
def __call__(self, input_ids, attention_mask, marker_pos, marker_mask, qtype):
h = self.encoder(input_ids, attention_mask) + self.type_emb(qtype)[:, None, :]
mask = _mask(attention_mask)
for layer in self.head.layers:
h = layer(h, mask)
m = mx.take_along_axis(h, mx.maximum(marker_pos, 0)[:, :, None], axis=1)
return mx.where(marker_mask, self.scorer(m).squeeze(-1).astype(mx.float32), -1e4)
class Session:
"""Drop-in for the onnxruntime session the router calls: run(["logits"], feeds) -> [np.ndarray]."""
def __init__(self, weights_dir):
c = json.load(open(os.path.join(weights_dir, "config.json")))
self.model = DecisionModel(c)
self.model.load_weights(os.path.join(weights_dir, "model.safetensors"), strict=True)
q = c.get("quantization")
if q:
nn.quantize(self.model, group_size=q["group_size"], bits=q["bits"],
class_predicate=lambda _, m: isinstance(m, nn.Linear) and m.weight.shape[-1] % q["group_size"] == 0)
self.model.eval()
mx.eval(self.model.parameters())
def run(self, names, feeds):
assert list(names) == ["logits"], names
am = feeds["attention_mask"]
f = {k: mx.array(v) for k, v in feeds.items()}
logits = self.model(f["input_ids"], None if am.all() else f["attention_mask"], f["marker_pos"],
f["marker_mask"], f["qtype"])
return [np.array(logits)]