DecisionLab / app /arthur.py
RealFalconsAI's picture
Upload 39 files
66ee87e verified
Raw History Blame Contribute Delete
9 kB
"""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