"""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 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. ) 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) 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