MLX
Joblib
Safetensors
English
reasoning
chain-of-thought
context-compression
soft-prompt
apple-silicon
Instructions to use baya1116/hypernet-sp-distill with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use baya1116/hypernet-sp-distill with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] huggingface-cli download --local-dir hypernet-sp-distill baya1116/hypernet-sp-distill
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
| """MLX port of TSE.AttnPoolSP (forward + forward_with_mass), numerically matching the PyTorch one. | |
| Loads pooler.pt (torch state_dict) into MLX arrays. MHA implemented manually to match | |
| torch nn.MultiheadAttention (packed in_proj QKV, head-avg attn weights for eviction mass).""" | |
| import math | |
| import numpy as np | |
| import mlx.core as mx | |
| def _ln(x, w, b, eps=1e-5): | |
| mu = x.mean(axis=-1, keepdims=True) | |
| var = ((x - mu) ** 2).mean(axis=-1, keepdims=True) | |
| return (x - mu) / mx.sqrt(var + eps) * w + b | |
| def _gelu(x): | |
| return 0.5 * x * (1.0 + mx.erf(x / math.sqrt(2.0))) | |
| def _mha(q, k, v, P, heads, want_w): | |
| # q:(B,Lq,H) k,v:(B,Lk,H); P holds packed in_proj (3H,H)+bias(3H), out_proj (H,H)+bias(H) | |
| B, Lq, H = q.shape; Lk = k.shape[1]; hd = H // heads | |
| Wi, bi, Wo, bo = P["in_w"], P["in_b"], P["out_w"], P["out_b"] | |
| qp = q @ Wi[:H].T + bi[:H] | |
| kp = k @ Wi[H:2 * H].T + bi[H:2 * H] | |
| vp = v @ Wi[2 * H:].T + bi[2 * H:] | |
| def split(t, L): | |
| return t.reshape(B, L, heads, hd).transpose(0, 2, 1, 3) # (B,heads,L,hd) | |
| qh, kh, vh = split(qp, Lq), split(kp, Lk), split(vp, Lk) | |
| scores = (qh @ kh.transpose(0, 1, 3, 2)) / math.sqrt(hd) # (B,heads,Lq,Lk) | |
| attn = mx.softmax(scores, axis=-1) | |
| out = (attn @ vh).transpose(0, 2, 1, 3).reshape(B, Lq, H) # (B,Lq,H) | |
| out = out @ Wo.T + bo | |
| w = attn.mean(axis=1) if want_w else None # head-avg (B,Lq,Lk) | |
| return out, w | |
| def _sinusoidal(L, H): | |
| pos = np.arange(L)[:, None].astype(np.float32) | |
| i = np.arange(0, H, 2).astype(np.float32) | |
| div = np.exp(-math.log(10000.0) * i / H) | |
| pe = np.zeros((L, H), np.float32) | |
| pe[:, 0::2] = np.sin(pos * div); pe[:, 1::2] = np.cos(pos * div) | |
| return mx.array(pe) | |
| class PoolerMLX: | |
| def __init__(self, path): | |
| import torch | |
| ck = torch.load(path, map_location="cpu", weights_only=False) | |
| sd = {k: v.float().numpy() for k, v in ck["pooler"].items()} | |
| a = ck.get("args") or ck.get("src_args") or {} | |
| self.A = {k: mx.array(v) for k, v in sd.items()} | |
| self.query = self.A["query"] # (n_sp,H) | |
| self.n_sp, self.H = self.query.shape | |
| self.heads = a.get("heads", 8) or 8 | |
| self.layers = sum(1 for k in sd if k.endswith(".lnq1.weight")) | |
| self.out_scale = self.A["out_scale"] | |
| self._pe = None | |
| def _blk_P(self, i, which): | |
| p = f"blocks.{i}.{which}." | |
| return {"in_w": self.A[p + "in_proj_weight"], "in_b": self.A[p + "in_proj_bias"], | |
| "out_w": self.A[p + "out_proj.weight"], "out_b": self.A[p + "out_proj.bias"]} | |
| def _ln_wb(self, i, name): | |
| return self.A[f"blocks.{i}.{name}.weight"], self.A[f"blocks.{i}.{name}.bias"] | |
| def pe(self, L): | |
| if self._pe is None or self._pe.shape[0] < L: | |
| self._pe = _sinusoidal(max(L, 1024), self.H) | |
| return self._pe[:L] | |
| def _run(self, past_emb, want_mass): | |
| B, L = past_emb.shape[0], past_emb.shape[1] | |
| past = past_emb + self.pe(L)[None] if L > 0 else past_emb | |
| q = mx.broadcast_to(self.query[None], (B, self.n_sp, self.H)) | |
| mass = mx.zeros((B, L)) if want_mass else None | |
| for i in range(self.layers): | |
| w1, b1 = self._ln_wb(i, "lnq1"); wk, bk = self._ln_wb(i, "lnk") | |
| if L > 0: | |
| kn = _ln(past, wk, bk) | |
| a, w = _mha(_ln(q, w1, b1), kn, kn, self._blk_P(i, "cross"), self.heads, want_mass) | |
| q = q + a | |
| if want_mass and w is not None: | |
| mass = mass + w.sum(axis=1) | |
| w2, b2 = self._ln_wb(i, "lnq2") | |
| qn = _ln(q, w2, b2) | |
| s, _ = _mha(qn, qn, qn, self._blk_P(i, "selfa"), self.heads, False) | |
| q = q + s | |
| w3, b3 = self._ln_wb(i, "lnq3") | |
| h = _ln(q, w3, b3) | |
| ff = (h @ self.A[f"blocks.{i}.ffn.0.weight"].T + self.A[f"blocks.{i}.ffn.0.bias"]) | |
| ff = _gelu(ff) | |
| ff = ff @ self.A[f"blocks.{i}.ffn.2.weight"].T + self.A[f"blocks.{i}.ffn.2.bias"] | |
| q = q + ff | |
| sp = _ln(q, self.A["ln_out.weight"], self.A["ln_out.bias"]) | |
| norm = mx.sqrt((sp * sp).sum(axis=-1, keepdims=True)) | |
| sp = sp / mx.maximum(norm, 1e-6) * mx.abs(self.out_scale) | |
| return (sp, mass) if want_mass else sp | |
| def forward(self, past_emb): | |
| return self._run(past_emb, False) | |
| def forward_with_mass(self, past_emb): | |
| return self._run(past_emb, True) | |