File size: 12,873 Bytes
7b0b288
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
"""Players for the LM Tetris Arena: causal LMs scored zero-shot, plus two baselines."""
from __future__ import annotations

import json
import os
import re
import threading
from collections import OrderedDict

import torch
from huggingface_hub import HfApi, hf_hub_download
from huggingface_hub.utils import GatedRepoError, RepositoryNotFoundError

MAX_PARAMS = int(os.environ.get("MAX_PARAMS", 250_000_000))
ALLOW_REMOTE_CODE = os.environ.get("ALLOW_REMOTE_CODE", "1") == "1"
MAX_CACHED_MODELS = int(os.environ.get("MAX_CACHED_MODELS", 6))
BATCH_SIZE = 16

PROMPTS = {
    # The rules of the game are stated in the prompt: tests reading comprehension.
    "guided": (
        "In Tetris, the goal is to clear lines, avoid holes and keep the stack low.\n"
        "This move {desc}.\n"
        "It is a"
    ),
    # No rules: the model must already know from pre-training what is good in Tetris.
    "blind": (
        "Here is a move from a game of Tetris.\n"
        "This move {desc}.\n"
        "It is a"
    ),
}
GOOD, BAD = " good move", " bad move"

REPO_RE = re.compile(r"^[A-Za-z0-9][\w.\-]*/[\w.\-]+$")
WEIGHT_EXT = (".safetensors", ".bin", ".pt", ".pth")

RANDOM_ID = "baseline/random"
ORACLE_ID = "baseline/oracle-reader"
BASELINES = {
    RANDOM_ID: "🎲 Random (baseline)",
    ORACLE_ID: "📏 Oracle reader (baseline)",
}


class ModelRejected(Exception):
    """Raised with a user-facing message when a model can't enter the arena."""


def fmt_params(n: int | None) -> str:
    if not n:
        return "?"
    if n >= 1e9:
        return f"{n / 1e9:.2f}B"
    if n >= 1e6:
        return f"{n / 1e6:.1f}M"
    return f"{n / 1e3:.0f}K"


# ----------------------------------------------------------------------------
# Baselines
# ----------------------------------------------------------------------------
class RandomPlayer:
    model_id = RANDOM_ID
    display = BASELINES[RANDOM_ID]
    n_params = 0
    sha = None
    custom_code = False

    def values(self, protocol, cands):
        return [0.0] * len(cands)  # everything ties -> uniform random choice


def _buckets(c):
    holes = -1 if c.new_holes < 0 else 0 if c.new_holes == 0 else 1 if c.new_holes == 1 else 2
    mh = c.max_height
    height = 0 if mh <= 4 else 1 if mh <= 8 else 2 if mh <= 12 else 3 if mh <= 16 else 4
    bump = 0 if c.bump <= 4 else 1 if c.bump <= 10 else 2
    landing = 0 if c.landing <= 0 else 1 if c.landing <= 2 else 2
    return holes, c.lines, height, bump, landing


class OracleReaderPlayer:
    """Reads the exact same descriptions the LMs see and ranks them with fixed
    common-sense priorities (holes > lines > height > surface > landing).
    It is a reference point for 'a perfect reader of the text', not a strong bot."""

    model_id = ORACLE_ID
    display = BASELINES[ORACLE_ID]
    n_params = 0
    sha = None
    custom_code = False

    def values(self, protocol, cands):
        out = []
        for c in cands:
            holes, lines, height, bump, landing = _buckets(c)
            out.append(float(((-holes + 5) * 10 + lines + 5) * 1000 + (-height + 5) * 100 + (-bump + 5) * 10 + (-landing + 5)))
        return out


# ----------------------------------------------------------------------------
# Language-model player
# ----------------------------------------------------------------------------
class LMPlayer:
    def __init__(self, model_id, sha, model, tokenizer, n_params, custom_code):
        self.model_id = model_id
        self.display = model_id
        self.sha = sha
        self.model = model
        self.tok = tokenizer
        self.n_params = n_params
        self.custom_code = custom_code
        self._cache: dict[tuple[str, str], float] = {}
        self._lock = threading.Lock()
        self._fwd_kwargs = None  # discovered on first forward
        self.max_len = self._max_len()
        # Reproduce whatever the tokenizer prepends by default (e.g. <s>)
        with_special = self._ids("a", True)
        without = self._ids("a", False)
        n = len(with_special) - len(without)
        self.prefix = with_special[:n] if n > 0 and with_special[n:n + len(without)] == without else []
        pad = tokenizer.pad_token_id
        if pad is None:
            pad = tokenizer.eos_token_id if tokenizer.eos_token_id is not None else 0
        self.pad_id = int(pad)

    def _max_len(self):
        cfg = self.model.config
        for k in ("max_position_embeddings", "n_positions", "max_seq_len", "seq_length", "block_size", "n_ctx"):
            v = getattr(cfg, k, None)
            if isinstance(v, int) and v > 0:
                return v
        return 2048

    def _ids(self, text, special):
        return list(self.tok(text, add_special_tokens=special)["input_ids"])

    def _build(self, context, continuation):
        ctx = self._ids(context, False)
        full = self._ids(context + continuation, False)
        if len(full) > len(ctx) and full[: len(ctx)] == ctx:
            cont = full[len(ctx):]
        else:  # tokenizer merged across the boundary: encode separately
            cont = self._ids(continuation, False)
        ids = self.prefix + ctx + cont
        return ids, len(self.prefix) + len(ctx)

    def _forward(self, input_ids, attention_mask):
        attempts = (
            [self._fwd_kwargs]
            if self._fwd_kwargs is not None
            else [{"attention_mask": True, "use_cache": False}, {"attention_mask": True}, {}]
        )
        last_err = None
        for kw in attempts:
            call = {}
            if kw.get("attention_mask"):
                call["attention_mask"] = attention_mask
            if "use_cache" in kw:
                call["use_cache"] = False
            try:
                out = self.model(input_ids=input_ids, **call)
                self._fwd_kwargs = kw
                break
            except TypeError as e:  # custom forward() without these kwargs
                last_err = e
        else:
            raise last_err
        if isinstance(out, dict) and "logits" in out:
            return out["logits"]
        if hasattr(out, "logits"):
            return out.logits
        if isinstance(out, (tuple, list)):
            return out[0]
        return out

    @torch.inference_mode()
    def _logprob_batch(self, items):
        """items: list of (ids, start). Returns sum log p(ids[start:] | ids[:start])."""
        L = max(len(ids) for ids, _ in items)
        inp = torch.full((len(items), L), self.pad_id, dtype=torch.long)
        mask = torch.zeros((len(items), L), dtype=torch.long)
        for i, (ids, _) in enumerate(items):
            inp[i, : len(ids)] = torch.tensor(ids)
            mask[i, : len(ids)] = 1
        logits = self._forward(inp, mask)
        out = []
        for i, (ids, start) in enumerate(items):
            pos = torch.arange(start - 1, len(ids) - 1)
            lp = torch.log_softmax(logits[i, pos].float(), dim=-1)
            tgt = torch.tensor(ids[start:])
            out.append(lp.gather(1, tgt[:, None]).sum().item())
        return out

    def score_descriptions(self, protocol, descs):
        template = PROMPTS[protocol]
        with self._lock:
            todo = [d for d in dict.fromkeys(descs) if (protocol, d) not in self._cache]
            items = []
            for d in todo:
                ctx = template.format(desc=d)
                for cont in (GOOD, BAD):
                    ids, start = self._build(ctx, cont)
                    if len(ids) > self.max_len:
                        raise ModelRejected(f"{self.model_id}: prompt ({len(ids)} tokens) exceeds context ({self.max_len}).")
                    items.append((ids, start))
            scores = []
            for i in range(0, len(items), BATCH_SIZE):
                scores.extend(self._logprob_batch(items[i : i + BATCH_SIZE]))
            for k, d in enumerate(todo):
                self._cache[(protocol, d)] = scores[2 * k] - scores[2 * k + 1]  # log P(good) - log P(bad)
            return {d: self._cache[(protocol, d)] for d in descs}

    def values(self, protocol, cands):
        table = self.score_descriptions(protocol, [c.description for c in cands])
        return [table[c.description] for c in cands]


# ----------------------------------------------------------------------------
# Validation + loading
# ----------------------------------------------------------------------------
_api = HfApi()
_models: "OrderedDict[tuple[str, str], LMPlayer]" = OrderedDict()
_models_lock = threading.Lock()


def precheck(model_id: str) -> dict:
    """Cheap checks before downloading any weights."""
    model_id = model_id.strip()
    if not REPO_RE.match(model_id):
        raise ModelRejected(f"`{model_id}` is not a valid repo id (expected `owner/name`).")
    try:
        info = _api.model_info(model_id, files_metadata=True)
    except (RepositoryNotFoundError, GatedRepoError):
        raise ModelRejected(f"`{model_id}` was not found, is private or gated.")
    except Exception as e:  # network etc.
        raise ModelRejected(f"`{model_id}`: could not read repo info ({type(e).__name__}).")
    if getattr(info, "gated", False):
        raise ModelRejected(f"`{model_id}` is gated; only public models can play.")
    files = {s.rfilename: (s.size or 0) for s in (info.siblings or [])}
    if "config.json" not in files:
        raise ModelRejected(f"`{model_id}` has no config.json (GGUF/ONNX-only repos are not supported).")
    try:
        cfg = json.load(open(hf_hub_download(model_id, "config.json", revision=info.sha)))
    except Exception:
        raise ModelRejected(f"`{model_id}`: config.json could not be read.")
    if cfg.get("is_encoder_decoder"):
        raise ModelRejected(f"`{model_id}` is an encoder-decoder model; only decoder-only models can play.")
    weights = {f: s for f, s in files.items() if f.endswith(WEIGHT_EXT) and "/" not in f.strip("./")}
    st = {f: s for f, s in weights.items() if f.endswith(".safetensors")}
    use = st or weights
    if not use:
        raise ModelRejected(f"`{model_id}` has no PyTorch/safetensors weights at the repo root.")
    est = None
    if getattr(info, "safetensors", None) and getattr(info.safetensors, "total", None):
        est = int(info.safetensors.total)
    else:
        dtype = str(cfg.get("dtype") or cfg.get("torch_dtype") or "float32")
        est = int(sum(use.values()) / (2 if ("16" in dtype) else 4))
    # generous margin: tied embeddings are often stored twice on disk; the exact
    # count after loading is what decides
    if est > MAX_PARAMS * 1.5:
        raise ModelRejected(f"`{model_id}` has ~{fmt_params(est)} parameters; the limit is {fmt_params(MAX_PARAMS)}.")
    # canonical id (fixes casing) so the leaderboard has one entry per repo
    return {"id": info.id or model_id, "sha": info.sha, "custom_code": "auto_map" in cfg, "est_params": est}


def load_player(model_id: str, meta: dict) -> LMPlayer:
    from transformers import AutoModelForCausalLM, AutoTokenizer

    key = (model_id, meta["sha"])
    with _models_lock:
        if key in _models:
            _models.move_to_end(key)
            return _models[key]
    if meta["custom_code"] and not ALLOW_REMOTE_CODE:
        raise ModelRejected(f"`{model_id}` needs custom code, which is disabled on this Space.")
    kw = dict(revision=meta["sha"], trust_remote_code=ALLOW_REMOTE_CODE)
    try:
        tok = AutoTokenizer.from_pretrained(model_id, **kw)
    except Exception as e:
        raise ModelRejected(f"`{model_id}`: tokenizer failed to load ({type(e).__name__}: {str(e)[:200]}).")
    try:
        model = AutoModelForCausalLM.from_pretrained(model_id, dtype=torch.float32, **kw)
    except Exception as e:
        raise ModelRejected(f"`{model_id}`: could not load as a causal LM ({type(e).__name__}: {str(e)[:200]}).")
    if getattr(model.config, "is_encoder_decoder", False):
        raise ModelRejected(f"`{model_id}` is an encoder-decoder model.")
    model.eval()
    n_params = sum(p.numel() for p in model.parameters())  # tied weights counted once
    if n_params > MAX_PARAMS:
        del model
        raise ModelRejected(f"`{model_id}` has {fmt_params(n_params)} parameters; the limit is {fmt_params(MAX_PARAMS)}.")
    player = LMPlayer(model_id, meta["sha"], model, tok, n_params, meta["custom_code"])
    # smoke test: one forward pass on a real prompt
    try:
        player.score_descriptions("guided", ["drops the piece into the lowest part of the board, clears one line, creates no new holes, keeps the stack very low and leaves the surface flat"])
    except ModelRejected:
        raise
    except Exception as e:
        raise ModelRejected(f"`{model_id}`: forward pass failed ({type(e).__name__}: {str(e)[:200]}).")
    with _models_lock:
        _models[key] = player
        while len(_models) > MAX_CACHED_MODELS:
            _models.popitem(last=False)
    return player