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