Spaces:
Paused
Paused
Upload gamemaster.py
Browse files- 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 |
-
|
| 193 |
-
|
|
|
|
|
|
|
|
|
|
| 194 |
).to(_model.device)
|
| 195 |
with torch.no_grad():
|
| 196 |
out = _model.generate(
|
| 197 |
-
|
| 198 |
pad_token_id=_tokenizer.eos_token_id,
|
| 199 |
)
|
| 200 |
-
|
|
|
|
| 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 |
|