Spaces:
Running on Zero
Running on Zero
File size: 9,001 Bytes
66ee87e | 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 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 | """Arthur (byte-level decision model) for DecisionLab.
The network and its tensor preparation are copied verbatim from arthur_v0_8_0.ipynb (sha256 fb6836d8c9534e9e…):
batch_tensors() from cell 17; ngram_buckets(), rope(), RMSNorm, Block and Arthur from cell 19; autocast() from
cell 20. load_arthur() follows the notebook's load_tier(); decide() returns the same answer format as FalconDec's
decide(), so DecisionLab compares every model the same way. Model folders hold data only (config.json,
model.safetensors); no code is ever loaded from them.
"""
from __future__ import annotations
import contextlib
import json
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from safetensors.torch import load_file
from .arthur_io import (LAYOUT, PAD, QTYPES, STRIDE, VOCAB, assemble, normalize_question, probabilities,
state_text)
# The notebook's globals that its functions read. Set by load_arthur().
DEVICE = torch.device("cpu")
AMP = None
ARCH_KEYS = ("d", "heads", "e", "buckets", "recursions", "interact", "mlp")
# ----------------------------------------------------------------------------- verbatim from the notebook
def batch_tensors(decisions):
rows = [assemble(d) for d in decisions]
width = -(-max(len(r) for r, _ in rows) // STRIDE) * STRIDE
k = max(len(w) for _, w in rows)
ids = np.zeros((len(rows), width), dtype=np.int64)
win = np.zeros((len(rows), k), dtype=np.int64)
mask = np.zeros((len(rows), k), dtype=bool)
for i, (r, w) in enumerate(rows):
ids[i, :len(r)] = r
win[i, :len(w)] = w
mask[i, :len(w)] = True
qt = np.array([QTYPES[d["qtype"]] for d in decisions], dtype=np.int64)
return tuple(torch.from_numpy(a).to(DEVICE) for a in (ids, win, mask, qt))
def ngram_buckets(ids, n, buckets):
h = torch.full_like(ids, {2: 11, 3: 23, 4: 47}[n])
ok = torch.ones_like(ids, dtype=torch.bool)
for k in range(n):
s = ids if k == 0 else torch.cat([ids[:, k:], ids.new_zeros(ids.shape[0], k)], dim=1)
h = (h * 1000003 + s) % 2147483647
ok = ok & (s > 0)
return h % buckets, ok
def rope(x, cos, sin):
x1, x2 = x[..., 0::2], x[..., 1::2]
c, s = cos[None, None].to(x.dtype), sin[None, None].to(x.dtype)
return torch.stack((x1 * c - x2 * s, x1 * s + x2 * c), dim=-1).flatten(-2)
class RMSNorm(nn.Module):
def __init__(self, d):
super().__init__()
self.weight = nn.Parameter(torch.ones(d))
def forward(self, x):
xf = x.float()
return (xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + 1e-6)).to(x.dtype) * self.weight.to(x.dtype)
class Block(nn.Module):
def __init__(self, d, heads, mlp):
super().__init__()
self.heads = heads
self.n1, self.n2 = RMSNorm(d), RMSNorm(d)
self.qkv = nn.Linear(d, 3 * d, bias=False)
self.o = nn.Linear(d, d, bias=False)
self.w12 = nn.Linear(d, 2 * mlp, bias=False)
self.w3 = nn.Linear(mlp, d, bias=False)
def forward(self, x, keep, cs=None):
b, p, d = x.shape
q, k, v = self.qkv(self.n1(x)).reshape(b, p, 3, self.heads, d // self.heads).permute(2, 0, 3, 1, 4)
if cs is not None:
q, k = rope(q, *cs), rope(k, *cs)
mask = torch.zeros(b, 1, 1, p, dtype=q.dtype, device=q.device).masked_fill(~keep[:, None, None, :], float("-inf"))
x = x + self.o(F.scaled_dot_product_attention(q, k, v, attn_mask=mask).transpose(1, 2).reshape(b, p, d))
g, u = self.w12(self.n2(x)).chunk(2, dim=-1)
return x + self.w3(F.silu(g) * u)
class Arthur(nn.Module):
"""Bytes + hashed 2/3/4-byte fragments -> depthwise conv -> 4:1 pooling -> one attention block reused `recursions`
times over the text (rotary positions), then `interact` times over the option vectors -> one logit per option."""
def __init__(self, cfg):
super().__init__()
self.cfg = cfg
d, e = cfg["d"], cfg["e"]
self.byte_emb = nn.Embedding(VOCAB, d, padding_idx=PAD)
self.hash_emb = nn.Embedding(cfg["buckets"], e)
self.hash_proj = nn.Linear(e, d, bias=False)
self.mix = nn.Conv1d(d, d, 5, padding=2, groups=d)
self.norm_in, self.norm_out = RMSNorm(d), RMSNorm(d)
self.block = Block(d, cfg["heads"], cfg["mlp"])
self.step = nn.Parameter(torch.zeros(cfg["recursions"], d))
self.qtype_emb = nn.Embedding(len(QTYPES), d)
self.scorer = nn.Sequential(nn.Linear(d, d), nn.GELU(), nn.Linear(d, 1))
nn.init.normal_(self.byte_emb.weight, std=0.02)
nn.init.normal_(self.hash_emb.weight, std=0.02)
with torch.no_grad():
self.byte_emb.weight[PAD].zero_()
nn.init.zeros_(self.scorer[-1].weight)
nn.init.zeros_(self.scorer[-1].bias)
def forward(self, ids, win, wmask, qtype):
x = self.byte_emb(ids)
h = sum(self.hash_emb(b) * ok.unsqueeze(-1).float() for b, ok in (ngram_buckets(ids, n, self.cfg["buckets"]) for n in (2, 3, 4)))
x = x + self.hash_proj(h)
x = x + F.gelu(self.mix(x.transpose(1, 2))).transpose(1, 2)
m = (ids > 0).to(x.dtype).unsqueeze(-1)
b, length, d = x.shape
p = length // STRIDE
cnt = m.reshape(b, p, STRIDE, 1).sum(2)
x = self.norm_in((x * m).reshape(b, p, STRIDE, d).sum(2) / cnt.clamp(min=1.0))
keep = cnt.squeeze(-1) > 0
hd = d // self.cfg["heads"]
ang = torch.arange(p, device=x.device, dtype=torch.float32)[:, None] / (10000 ** (torch.arange(0, hd, 2, device=x.device) / hd))
for r in range(self.cfg["recursions"]):
x = self.block(x + self.step[r], keep, (ang.cos(), ang.sin()))
km = keep.to(x.dtype).unsqueeze(-1)
ctx = (x * km).sum(1) / km.sum(1).clamp(min=1.0)
o = x.gather(1, win.unsqueeze(-1).expand(-1, -1, d)) + ctx.unsqueeze(1) + self.qtype_emb(qtype).unsqueeze(1).to(x.dtype)
for _ in range(self.cfg["interact"]):
o = self.block(o, wmask)
return self.scorer(self.norm_out(o)).squeeze(-1).float().masked_fill(~wmask, -1e4)
def autocast():
return torch.autocast("cuda", dtype=AMP) if AMP is not None else contextlib.nullcontext()
# ----------------------------------------------------------------------------- DecisionLab adapters
def _bf16_supported() -> bool:
"""On ZeroGPU there is no real GPU at load time, so the question can fail; its GPUs (H200) support bfloat16."""
try:
return torch.cuda.is_bf16_supported()
except Exception:
return True
def load_arthur(folder, device) -> tuple:
"""(net, temperatures, config) from an Arthur folder, checked as the notebook's load_tier() checks it."""
global DEVICE, AMP
folder = Path(folder)
cfg = json.loads((folder / "config.json").read_text(encoding="utf-8"))
if cfg.get("layout") != LAYOUT:
raise ValueError(f"{folder} was trained with input layout {cfg.get('layout')}, but DecisionLab's Arthur code "
f"(notebook v0.8.0) uses {LAYOUT}. Retrain it with v0.8.0 or update app/arthur_io.py.")
bad = [k for k in ARCH_KEYS if not isinstance(cfg.get(k), int) or cfg[k] <= 0]
if bad:
raise ValueError(f"{folder}/config.json has missing or invalid architecture values: {', '.join(bad)}.")
temps = np.array(cfg["temperatures"], dtype=np.float64)
if temps.shape != (len(QTYPES), 4) or not np.all(temps > 0):
raise ValueError(f"{folder}/config.json temperatures must be 3 x 4 positive numbers (got shape {temps.shape}).")
DEVICE = torch.device(device)
AMP = (torch.bfloat16 if _bf16_supported() else torch.float16) if DEVICE.type == "cuda" else None
net = Arthur({k: cfg[k] for k in ARCH_KEYS})
net.load_state_dict({k: v.float() for k, v in load_file(str(folder / "model.safetensors")).items()}, strict=True)
return net.to(DEVICE).eval(), temps, cfg
@torch.no_grad()
def decide(net, temps, state, questions: dict) -> dict:
"""{question name: {"choice", "probs", "p_true"?, "expected_level"?}} for a Jev/Laya-style question dict."""
s = state_text(state)
names, metas, decisions = [], [], []
for name, q in questions.items():
qtype, text, keys, opts = normalize_question(q)
names.append(name)
metas.append((qtype, keys))
decisions.append({"question": text, "options": opts, "state": s, "qtype": qtype})
with autocast():
logits = net(*batch_tensors(decisions)).float().cpu().numpy()
out = {}
for i, (name, (qtype, keys)) in enumerate(zip(names, metas)):
p = probabilities(logits[i, :len(keys)], qtype, temps)
k = int(np.argmax(p))
r = {"choice": keys[k], "probs": {str(key): float(v) for key, v in zip(keys, p)}}
if qtype == "noul":
r["p_true"] = float(p[0])
if qtype == "score":
r["expected_level"] = float(np.dot(p, np.arange(len(p))))
out[name] = r
return out
|