summaryrefslogtreecommitdiffstats
path: root/ai_dungeon.py
diff options
context:
space:
mode:
Diffstat (limited to 'ai_dungeon.py')
-rwxr-xr-xai_dungeon.py17
1 files changed, 10 insertions, 7 deletions
diff --git a/ai_dungeon.py b/ai_dungeon.py
index ac5fe82..63f6092 100755
--- a/ai_dungeon.py
+++ b/ai_dungeon.py
@@ -3,24 +3,27 @@
"""AI Dungeon is a text-adventure style game powered by AI."""
import argparse
+import re
import sys
+
import transformers
-def trim_sentence(tokenizer: transformers.PreTrainedTokenizer, message: str) -> str:
+sentence_end = re.compile(r"[.?!(\"'][^.?!(\"']*$")
+
+
+def trim_sentence(message: str) -> str:
"""Remove extra output after the last period."""
- periodt = tokenizer.encode(".")[0]
- tokens = tokenizer.encode(message)
- for i in range(len(tokens) - 1, -1, -1):
- if tokens[i] == periodt:
- return tokenizer.decode(tokens[: i + 1])
+ match = sentence_end.search(message)
+ if match:
+ return message[: match.start() + 1]
return message
def generate(model: transformers.TextGenerationPipeline, prompt: str, **kwargs) -> str:
"""Generate text from a text generator."""
out = model(prompt, do_sample=True, temperature=0.8, top_k=60, top_p=0.9, **kwargs)
- return trim_sentence(model.tokenizer, out[0]["generated_text"])
+ return trim_sentence(out[0]["generated_text"])
def load_model(model: str) -> transformers.TextGenerationPipeline: