DedeProGames commited on
Commit
a27f7ef
·
verified ·
1 Parent(s): 8c5a888

LMPlayer.complete: the model's raw greedy completion for the chosen move

Browse files
Files changed (1) hide show
  1. 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