Quazim0t0 commited on
Commit
04814f5
·
verified ·
1 Parent(s): b9f24f1

Upload gamemaster.py

Browse files
Files changed (1) hide show
  1. game/gamemaster.py +8 -4
game/gamemaster.py CHANGED
@@ -189,15 +189,19 @@ def _run_model(system: str, user: str) -> str:
189
 
190
  messages = [{"role": "system", "content": system},
191
  {"role": "user", "content": user}]
192
- inputs = _tokenizer.apply_chat_template(
193
- messages, add_generation_prompt=True, return_tensors="pt"
 
 
 
194
  ).to(_model.device)
195
  with torch.no_grad():
196
  out = _model.generate(
197
- inputs, max_new_tokens=400, do_sample=True, temperature=0.7, top_p=0.9,
198
  pad_token_id=_tokenizer.eos_token_id,
199
  )
200
- text = _tokenizer.decode(out[0][inputs.shape[1]:], skip_special_tokens=True)
 
201
  return text
202
 
203
 
 
189
 
190
  messages = [{"role": "system", "content": system},
191
  {"role": "user", "content": user}]
192
+ # return_dict=True gives a BatchEncoding (input_ids + attention_mask) on
193
+ # every transformers version — newer releases return it by default, and
194
+ # passing it positionally to generate() crashes on `.shape`.
195
+ enc = _tokenizer.apply_chat_template(
196
+ messages, add_generation_prompt=True, return_tensors="pt", return_dict=True
197
  ).to(_model.device)
198
  with torch.no_grad():
199
  out = _model.generate(
200
+ **enc, max_new_tokens=400, do_sample=True, temperature=0.7, top_p=0.9,
201
  pad_token_id=_tokenizer.eos_token_id,
202
  )
203
+ n_in = enc["input_ids"].shape[1]
204
+ text = _tokenizer.decode(out[0][n_in:], skip_special_tokens=True)
205
  return text
206
 
207