SLM-Tetris-Arena / players.py
DedeProGames's picture
LMPlayer.complete: the model's raw greedy completion for the chosen move
a27f7ef verified
Raw History Blame
15.2 kB
"""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))
MIN_PARAMS = int(os.environ.get("MIN_PARAMS", 50_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
GEN_TOKENS = 8 # length of the raw completion shown under each board
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._gen_cache: dict[tuple[str, str], str] = {}
self._use_generate = None # decided on the first completion
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]
@torch.inference_mode()
def complete(self, protocol, desc):
"""The model's own greedy text after the prompt of the chosen move (display only, not used to decide)."""
key = (protocol, desc)
with self._lock:
if key in self._gen_cache:
return self._gen_cache[key]
ids = self.prefix + self._ids(PROMPTS[protocol].format(desc=desc), False)
n_new = min(GEN_TOKENS, self.max_len - len(ids))
out = self._generate(ids, n_new) if n_new > 0 else []
text = self.tok.decode(out, skip_special_tokens=True).split("\n")[0]
self._gen_cache[key] = text
return text
def _generate(self, ids, n_new):
"""Greedy decoding: `generate()` (with KV cache) when the model supports it, else a plain loop."""
inp = torch.tensor([ids], dtype=torch.long)
if self._use_generate is not False:
try:
gen = self.model.generate(inp, attention_mask=torch.ones_like(inp), max_new_tokens=n_new, do_sample=False,
pad_token_id=self.pad_id, eos_token_id=self.tok.eos_token_id)
self._use_generate = True
return gen[0, len(ids):].tolist()
except Exception:
self._use_generate = False
eos = self.tok.eos_token_id
out = []
for _ in range(n_new):
inp = torch.tensor([ids + out], dtype=torch.long)
nxt = int(self._forward(inp, torch.ones_like(inp))[0, -1].argmax())
if nxt == eos:
break
out.append(nxt)
if "\n" in self.tok.decode(out, skip_special_tokens=True):
break
return out
# ----------------------------------------------------------------------------
# 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)
if est < MIN_PARAMS:
raise ModelRejected(f"`{model_id}` has only ~{fmt_params(est)} parameters; the minimum is {fmt_params(MIN_PARAMS)}.")
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)}.")
if n_params < MIN_PARAMS:
del model
raise ModelRejected(f"`{model_id}` has only {fmt_params(n_params)} parameters; the minimum is {fmt_params(MIN_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