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
File size: 3,297 Bytes
b5989f0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 | """trigger_gate.py — runtime recall gate (TRIGGER_V3 head, wiring stage).
Scores the current question in the production layout [BOS][SP(kept)][raw window]
[question] with the question-position pre-RoPE queries (layers 8/14/20) and a
4608-dim logistic head. fires() == True means "the answer needs something that
was evicted from this session's window" — only then does the app pay for an
archive retrieve + verbatim re-injection.
Threshold note: the gate npz carries a TPR-first threshold (-2.6: heldout TPR
1.000, FPR 0.043) — a false fire costs one redundant retrieve (today's per-turn
status quo), a miss costs the recall itself.
"""
import numpy as np
import torch
LAYERS = (8, 14, 20)
class TriggerGate:
def __init__(self, llm, tok, pooler, head_path, thresh=None):
z = np.load(head_path)
self.coef = z["coef"].astype(np.float32).ravel()
self.b = float(np.ravel(z["intercept"])[0])
self.mean = z["mean"].astype(np.float32)
self.scale = z["scale"].astype(np.float32)
if thresh is not None:
self.thresh = float(thresh)
elif "thresh" in z.files:
self.thresh = float(z["thresh"])
else:
self.thresh = 0.0
self.llm, self.tok, self.pooler = llm, tok, pooler
self.dev = next(llm.parameters()).device
self.embT = llm.get_input_embeddings()
self.last_score = 0.0
@torch.no_grad()
def score(self, kept_ids, window_ids, q_text):
from transformers.cache_utils import DynamicCache
pre = self.tok.encode(f"<|end▁of▁sentence|><|User|>{q_text}<|Assistant|>"
f"<think>\n\n</think>\n\n", add_special_tokens=False)
cache = DynamicCache()
bos = self.tok.bos_token_id
self.llm.model(input_ids=torch.tensor([[bos]], device=self.dev),
past_key_values=cache, use_cache=True)
parts = []
if kept_ids:
sp = self.pooler(self.embT(torch.tensor([kept_ids], device=self.dev))
.float()).to(self.embT.weight.dtype)
parts.append(sp)
if window_ids:
parts.append(self.embT(torch.tensor([window_ids], device=self.dev)))
if parts:
self.llm.model(inputs_embeds=torch.cat(parts, 1), past_key_values=cache,
use_cache=True)
outs, hooks = {}, []
for li in LAYERS:
attn = self.llm.model.layers[li].self_attn
def mk(li):
def fn(_m, _i, o):
outs[li] = o.detach()
return fn
hooks.append(attn.q_proj.register_forward_hook(mk(li)))
try:
self.llm.model(inputs_embeds=self.embT(torch.tensor([pre], device=self.dev)),
past_key_values=cache, use_cache=True)
finally:
for h in hooks:
h.remove()
f = np.concatenate([outs[li][0].mean(0).float().cpu().numpy().ravel()
for li in LAYERS])
self.last_score = float(((f - self.mean) / self.scale) @ self.coef + self.b)
return self.last_score
def fires(self, kept_ids, window_ids, q_text):
return self.score(kept_ids, window_ids, q_text) > self.thresh
|