Download modeling_cerebellum.py from mkzero/Cerebellum-2B-BF16: direct link, hf CLI and curl.
- Browser
- Download file 10.9 kB
-
https://huggingface.co/mkzero/Cerebellum-2B-BF16/resolve/147775b9b3a62c14089eaff5d64e2119d430f553/modeling_cerebellum.py
- Command line
-
hf download hf://mkzero/Cerebellum-2B-BF16@147775b9b3a62c14089eaff5d64e2119d430f553/modeling_cerebellum.py
-
curl -L -o modeling_cerebellum.py https://huggingface.co/mkzero/Cerebellum-2B-BF16/resolve/147775b9b3a62c14089eaff5d64e2119d430f553/modeling_cerebellum.py
10.9 kB
| """ | |
| Cerebellum-2B (小脑-2B): Sub-30ms Non-Autoregressive Agent Decision Model | |
| Developed on Qwen3.5-2B Backbone with Symmetric Cross-Option SetPointerHead and Chow's Rejection Gate. | |
| """ | |
| import os, time, math, re | |
| from dataclasses import dataclass | |
| from typing import List, Dict, Optional, Union | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, PreTrainedModel, PretrainedConfig | |
| SPECIAL = ["<|fim_prefix|>", "<|fim_middle|>", "<|box_start|>", "<|box_end|>", "<|fim_suffix|>"] | |
| class CerebellumDecision: | |
| action: str | |
| action_index: int | |
| confidence: float | |
| probabilities: Dict[str, float] | |
| needs_escalation: bool | |
| escalate_probability: float | |
| latency_ms: float | |
| class SetPointerHead(nn.Module): | |
| """ | |
| Permutation-Equivariant Set Transformer Readout Head | |
| - Uses Cross-Option Self-Attention WITHOUT positional encoding | |
| - Mathematically guarantees 0.000 Option-Order Flip Rate! | |
| """ | |
| def __init__(self, d=2048, dp=512, n_heads=4): | |
| super().__init__() | |
| self.d = d | |
| self.dp = dp | |
| self.q_proj = nn.Sequential( | |
| nn.Linear(d, dp), | |
| nn.GELU(), | |
| nn.RMSNorm(dp), | |
| nn.Linear(dp, dp) | |
| ) | |
| self.opt_in_proj = nn.Linear(d, dp) | |
| self.set_attn = nn.MultiheadAttention(embed_dim=dp, num_heads=n_heads, batch_first=True) | |
| self.norm1 = nn.RMSNorm(dp) | |
| self.ffn = nn.Sequential( | |
| nn.Linear(dp, dp * 2), | |
| nn.GELU(), | |
| nn.Linear(dp * 2, dp) | |
| ) | |
| self.norm2 = nn.RMSNorm(dp) | |
| self.scale = 1.0 / (dp ** 0.5) | |
| def forward(self, h_decide, h_opts): | |
| q = self.q_proj(h_decide) # [dp] | |
| x_opts = self.opt_in_proj(h_opts).unsqueeze(0) # [1, K, dp] | |
| attn_out, _ = self.set_attn(x_opts, x_opts, x_opts) | |
| x_opts = self.norm1(x_opts + attn_out) | |
| x_opts = self.norm2(x_opts + self.ffn(x_opts)).squeeze(0) # [K, dp] | |
| logits = (x_opts @ q) * self.scale # [K] | |
| return logits | |
| class ActEscalateHead(nn.Module): | |
| """ | |
| Chow's Optimal Rejection Gate: | |
| Computes distribution sufficient statistics + State representation: | |
| 1. Top-1 Probability: max(p) | |
| 2. Top-2 Margin: p_top1 - p_top2 | |
| 3. Normalized Shannon Entropy: H(p) / log(K) | |
| 4. Candidate Option Budget: K / 32 | |
| Outputs: [P(Act), P(Escalate)] | |
| """ | |
| def __init__(self, d=2048, dp=128): | |
| super().__init__() | |
| self.state_proj = nn.Linear(d, dp) | |
| self.mlp = nn.Sequential( | |
| nn.Linear(dp + 4, 128), | |
| nn.GELU(), | |
| nn.Linear(128, 2) | |
| ) | |
| def forward(self, h_decide, logits): | |
| p = F.softmax(logits, dim=-1) | |
| K = p.shape[0] | |
| top1 = p.max() | |
| if K > 1: | |
| top2 = torch.topk(p, 2).values[1] | |
| margin = top1 - top2 | |
| else: | |
| margin = torch.tensor(1.0, device=p.device) | |
| entropy = -(p * torch.log(p.clamp(min=1e-9))).sum() | |
| norm_entropy = entropy / math.log(max(K, 2)) | |
| budget = torch.tensor(min(K / 32.0, 1.0), device=p.device, dtype=h_decide.dtype) | |
| stats = torch.stack([top1, margin, norm_entropy, budget]).to(h_decide.dtype) | |
| h_proj = self.state_proj(h_decide) | |
| feat = torch.cat([h_proj, stats], dim=-1) | |
| esc_logits = self.mlp(feat) | |
| return esc_logits | |
| def bidirectional_state_branch_mask_batch(segs, device, dtype=torch.bfloat16): | |
| L = max(len(s) for s in segs) | |
| s = torch.full((len(segs), L), -1, device=device, dtype=torch.long) | |
| for b, seg in enumerate(segs): | |
| s[b, :len(seg)] = torch.tensor(seg, device=device, dtype=torch.long) | |
| q_seg = s[:, :, None] | |
| k_seg = s[:, None, :] | |
| valid_key = (k_seg != -1) | |
| valid_q = (q_seg != -1) | |
| state_to_state = (q_seg == 0) & (k_seg == 0) | |
| q_to_state = (q_seg > 0) & (k_seg == 0) | |
| within_branch = (q_seg > 0) & (q_seg == k_seg) | |
| allow = (state_to_state | q_to_state | within_branch) & valid_key & valid_q | |
| allow = allow | torch.eye(L, dtype=torch.bool, device=device)[None] | |
| mask = torch.zeros((len(segs), 1, L, L), dtype=dtype, device=device) | |
| mask.masked_fill_(~allow[:, None, :, :], torch.finfo(dtype).min) | |
| return mask | |
| class CerebellumModel(nn.Module): | |
| """ | |
| Cerebellum-2B End-to-End Decision Model | |
| Provides single-forward-pass sub-30ms decision making for AI Agents. | |
| """ | |
| def __init__( | |
| self, | |
| model_dir: str, | |
| device: Optional[Union[str, torch.device]] = None, | |
| dtype: Optional[torch.dtype] = None, | |
| **kwargs | |
| ): | |
| super().__init__() | |
| # Auto-detect optimal device if not provided | |
| if device is None: | |
| if torch.cuda.is_available(): | |
| device = "cuda:0" | |
| elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available(): | |
| device = "mps" | |
| else: | |
| device = "cpu" | |
| self.device = torch.device(device) | |
| # Auto-detect optimal dtype | |
| if dtype is None: | |
| if self.device.type == "cuda" and torch.cuda.is_bf16_supported(): | |
| dtype = torch.bfloat16 | |
| elif self.device.type == "mps": | |
| dtype = torch.float16 | |
| else: | |
| dtype = torch.float32 | |
| self.dtype = dtype | |
| is_local = os.path.exists(model_dir) | |
| self.tok = AutoTokenizer.from_pretrained(model_dir, local_files_only=is_local, trust_remote_code=True) | |
| # Load merged base transformer safely across CUDA / MPS / CPU | |
| if self.device.type == "cuda": | |
| causal_lm = AutoModelForCausalLM.from_pretrained( | |
| model_dir, | |
| torch_dtype=dtype, | |
| device_map=self.device, | |
| local_files_only=is_local, | |
| trust_remote_code=True | |
| ) | |
| else: | |
| causal_lm = AutoModelForCausalLM.from_pretrained( | |
| model_dir, | |
| torch_dtype=dtype, | |
| local_files_only=is_local, | |
| trust_remote_code=True | |
| ).to(self.device) | |
| self.base_model = causal_lm.model | |
| # Load trained decision heads (support local path or HF Hub download) | |
| heads_path = os.path.join(model_dir, "heads.pt") | |
| if not os.path.exists(heads_path): | |
| try: | |
| from huggingface_hub import hf_hub_download | |
| heads_path = hf_hub_download(repo_id=model_dir, filename="heads.pt") | |
| except Exception: | |
| pass | |
| heads_data = torch.load(heads_path, map_location=self.device) | |
| self.pointer_head = SetPointerHead(2048, 512, n_heads=4).to(self.device).to(dtype) | |
| self.escalate_head = ActEscalateHead(2048, 128).to(self.device).to(dtype) | |
| self.pointer_head.load_state_dict(heads_data["pointer_head"]) | |
| self.escalate_head.load_state_dict(heads_data["escalate_head"]) | |
| self.pointer_head.eval() | |
| self.escalate_head.eval() | |
| self.base_model.eval() | |
| self.pad_id = self.tok.pad_token_id if self.tok.pad_token_id is not None else 0 | |
| def from_pretrained( | |
| cls, | |
| pretrained_model_name_or_path: str, | |
| *args, | |
| device: Optional[Union[str, torch.device]] = None, | |
| dtype: Optional[torch.dtype] = None, | |
| **kwargs | |
| ): | |
| return cls(model_dir=pretrained_model_name_or_path, device=device, dtype=dtype, **kwargs) | |
| def _encode_query(self, state: str, candidates: List[str], instruction: str = "Select the best action to execute next."): | |
| special_re = re.compile(r"<|([A-Za-z0-9_]+)|>") | |
| def utok(text): | |
| return self.tok(special_re.sub(r"<¦\1¦>", text), add_special_tokens=False).input_ids | |
| raw_state = utok(state) | |
| max_state = 1024 | |
| if len(raw_state) > max_state - 1: | |
| half = (max_state - 1) // 2 | |
| state_tokens = raw_state[:half] + raw_state[-half:] | |
| else: | |
| state_tokens = raw_state | |
| S = [self.tok.convert_tokens_to_ids(SPECIAL[0])] + state_tokens | |
| ids, seg, pos = list(S), [0] * len(S), list(range(len(S))) | |
| q_id, o_id, c_id, d_id = (self.tok.convert_tokens_to_ids(t) for t in SPECIAL[1:]) | |
| br = [q_id] + utok(instruction) | |
| oi = [] | |
| for o in candidates: | |
| o_tok = utok(o) | |
| if len(o_tok) > 64: | |
| o_tok = o_tok[:32] + o_tok[-32:] | |
| br += [o_id] + o_tok + [c_id] | |
| oi.append(len(br) - 1) | |
| br.append(d_id) | |
| base = len(ids) | |
| ids += br | |
| seg += [1] * len(br) | |
| pos += list(range(len(S), len(S) + len(br))) | |
| decide_idx = base + len(br) - 1 | |
| opt_idx = [base + i for i in oi] | |
| return { | |
| "ids": ids, | |
| "seg": seg, | |
| "pos": pos, | |
| "decide_idx": decide_idx, | |
| "opt_idx": opt_idx | |
| } | |
| def decide( | |
| self, | |
| state: str, | |
| candidates: List[str], | |
| instruction: str = "Select the best action to execute next.", | |
| escalate_threshold: float = 0.50 | |
| ) -> CerebellumDecision: | |
| t0 = time.perf_counter() | |
| enc = self._encode_query(state, candidates, instruction) | |
| ids = torch.tensor([enc["ids"]], device=self.device, dtype=torch.long) | |
| pos = torch.tensor([enc["pos"]], device=self.device, dtype=torch.long) | |
| mask = bidirectional_state_branch_mask_batch([enc["seg"]], self.device, dtype=self.dtype) | |
| hidden = self.base_model(input_ids=ids, position_ids=pos, attention_mask=mask).last_hidden_state[0] | |
| h_d = hidden[enc["decide_idx"]] | |
| h_o = hidden[torch.tensor(enc["opt_idx"], device=self.device)] | |
| logits = self.pointer_head(h_d, h_o) | |
| esc_logits = self.escalate_head(h_d, logits) | |
| probs = F.softmax(logits, dim=-1).cpu().tolist() | |
| esc_probs = F.softmax(esc_logits, dim=-1).cpu().tolist() # [0: Act, 1: Escalate] | |
| best_idx = int(torch.argmax(logits).item()) | |
| confidence = probs[best_idx] | |
| p_escalate = esc_probs[1] | |
| needs_escalation = (p_escalate >= escalate_threshold) or (confidence < 0.65) | |
| latency = (time.perf_counter() - t0) * 1000.0 | |
| prob_dict = {cand: p for cand, p in zip(candidates, probs)} | |
| return CerebellumDecision( | |
| action=candidates[best_idx], | |
| action_index=best_idx, | |
| confidence=confidence, | |
| probabilities=prob_dict, | |
| needs_escalation=needs_escalation, | |
| escalate_probability=p_escalate, | |
| latency_ms=latency | |
| ) | |