Spaces:
Running
Running
LMPlayer.complete: the model's raw greedy completion for the chosen move
Browse files- players.py +40 -0
players.py
CHANGED
|
@@ -16,6 +16,7 @@ MIN_PARAMS = int(os.environ.get("MIN_PARAMS", 50_000))
|
|
| 16 |
ALLOW_REMOTE_CODE = os.environ.get("ALLOW_REMOTE_CODE", "1") == "1"
|
| 17 |
MAX_CACHED_MODELS = int(os.environ.get("MAX_CACHED_MODELS", 6))
|
| 18 |
BATCH_SIZE = 16
|
|
|
|
| 19 |
|
| 20 |
PROMPTS = {
|
| 21 |
# The rules of the game are stated in the prompt: tests reading comprehension.
|
|
@@ -113,6 +114,8 @@ class LMPlayer:
|
|
| 113 |
self.n_params = n_params
|
| 114 |
self.custom_code = custom_code
|
| 115 |
self._cache: dict[tuple[str, str], float] = {}
|
|
|
|
|
|
|
| 116 |
self._lock = threading.Lock()
|
| 117 |
self._fwd_kwargs = None # discovered on first forward
|
| 118 |
self.max_len = self._max_len()
|
|
@@ -217,6 +220,43 @@ class LMPlayer:
|
|
| 217 |
table = self.score_descriptions(protocol, [c.description for c in cands])
|
| 218 |
return [table[c.description] for c in cands]
|
| 219 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 220 |
|
| 221 |
# ----------------------------------------------------------------------------
|
| 222 |
# Validation + loading
|
|
|
|
| 16 |
ALLOW_REMOTE_CODE = os.environ.get("ALLOW_REMOTE_CODE", "1") == "1"
|
| 17 |
MAX_CACHED_MODELS = int(os.environ.get("MAX_CACHED_MODELS", 6))
|
| 18 |
BATCH_SIZE = 16
|
| 19 |
+
GEN_TOKENS = 8 # length of the raw completion shown under each board
|
| 20 |
|
| 21 |
PROMPTS = {
|
| 22 |
# The rules of the game are stated in the prompt: tests reading comprehension.
|
|
|
|
| 114 |
self.n_params = n_params
|
| 115 |
self.custom_code = custom_code
|
| 116 |
self._cache: dict[tuple[str, str], float] = {}
|
| 117 |
+
self._gen_cache: dict[tuple[str, str], str] = {}
|
| 118 |
+
self._use_generate = None # decided on the first completion
|
| 119 |
self._lock = threading.Lock()
|
| 120 |
self._fwd_kwargs = None # discovered on first forward
|
| 121 |
self.max_len = self._max_len()
|
|
|
|
| 220 |
table = self.score_descriptions(protocol, [c.description for c in cands])
|
| 221 |
return [table[c.description] for c in cands]
|
| 222 |
|
| 223 |
+
@torch.inference_mode()
|
| 224 |
+
def complete(self, protocol, desc):
|
| 225 |
+
"""The model's own greedy text after the prompt of the chosen move (display only, not used to decide)."""
|
| 226 |
+
key = (protocol, desc)
|
| 227 |
+
with self._lock:
|
| 228 |
+
if key in self._gen_cache:
|
| 229 |
+
return self._gen_cache[key]
|
| 230 |
+
ids = self.prefix + self._ids(PROMPTS[protocol].format(desc=desc), False)
|
| 231 |
+
n_new = min(GEN_TOKENS, self.max_len - len(ids))
|
| 232 |
+
out = self._generate(ids, n_new) if n_new > 0 else []
|
| 233 |
+
text = self.tok.decode(out, skip_special_tokens=True).split("\n")[0]
|
| 234 |
+
self._gen_cache[key] = text
|
| 235 |
+
return text
|
| 236 |
+
|
| 237 |
+
def _generate(self, ids, n_new):
|
| 238 |
+
"""Greedy decoding: `generate()` (with KV cache) when the model supports it, else a plain loop."""
|
| 239 |
+
inp = torch.tensor([ids], dtype=torch.long)
|
| 240 |
+
if self._use_generate is not False:
|
| 241 |
+
try:
|
| 242 |
+
gen = self.model.generate(inp, attention_mask=torch.ones_like(inp), max_new_tokens=n_new, do_sample=False,
|
| 243 |
+
pad_token_id=self.pad_id, eos_token_id=self.tok.eos_token_id)
|
| 244 |
+
self._use_generate = True
|
| 245 |
+
return gen[0, len(ids):].tolist()
|
| 246 |
+
except Exception:
|
| 247 |
+
self._use_generate = False
|
| 248 |
+
eos = self.tok.eos_token_id
|
| 249 |
+
out = []
|
| 250 |
+
for _ in range(n_new):
|
| 251 |
+
inp = torch.tensor([ids + out], dtype=torch.long)
|
| 252 |
+
nxt = int(self._forward(inp, torch.ones_like(inp))[0, -1].argmax())
|
| 253 |
+
if nxt == eos:
|
| 254 |
+
break
|
| 255 |
+
out.append(nxt)
|
| 256 |
+
if "\n" in self.tok.decode(out, skip_special_tokens=True):
|
| 257 |
+
break
|
| 258 |
+
return out
|
| 259 |
+
|
| 260 |
|
| 261 |
# ----------------------------------------------------------------------------
|
| 262 |
# Validation + loading
|