Spaces:
Running
Running
Download players.py from DedeProGames/SLM-Tetris-Arena: direct link, hf CLI and curl.
- Browser
- Download file 15.2 kB
-
https://huggingface.co/spaces/DedeProGames/SLM-Tetris-Arena/resolve/c41ac50b9e3d2abf5fd1f8541e32bc0915a7b682/players.py
- Command line
-
hf download hf://spaces/DedeProGames/SLM-Tetris-Arena@c41ac50b9e3d2abf5fd1f8541e32bc0915a7b682/players.py
-
curl -L -o players.py https://huggingface.co/spaces/DedeProGames/SLM-Tetris-Arena/resolve/c41ac50b9e3d2abf5fd1f8541e32bc0915a7b682/players.py
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 | |
| 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] | |
| 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 | |