summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rwxr-xr-xai_dungeon.py22
1 files changed, 14 insertions, 8 deletions
diff --git a/ai_dungeon.py b/ai_dungeon.py
index baac638..ad18c48 100755
--- a/ai_dungeon.py
+++ b/ai_dungeon.py
@@ -12,17 +12,22 @@ from nltk import tokenize
import transformers
-sentence_end = re.compile(r"[.?!\"')]$")
+def load_nltk():
+ try:
+ tokenize.sent_tokenize('')
+ except LookupError:
+ nltk.download("punkt")
+
+
+def complete_sentence(snippet: str) -> bool:
+ sentences = tokenize.sent_tokenize(snippet)
+ return len(tokenize.sent_tokenize(sentences[-1] + ' extra')) == 2
def trim_sentence(message: str) -> str:
"""Remove extra output after the last period."""
- try:
- sentences = tokenize.sent_tokenize(message)
- except LookupError:
- nltk.download("punkt")
- sentences = tokenize.sent_tokenize(message)
- if sentence_end.search(sentences[-1]) or len(sentences) == 1:
+ sentences = tokenize.sent_tokenize(message)
+ if complete_sentence(sentences[-1]) or len(sentences) == 1:
return message
return " ".join(sentences[:-1])
@@ -103,6 +108,7 @@ def wrap(message: str) -> str:
def main() -> None:
"""Run the main game loop."""
+ load_nltk()
parser = argparse.ArgumentParser("AI dungeon clone")
parser.add_argument(
"--model", default="gpt2", help="Model to use"
@@ -137,7 +143,7 @@ def main() -> None:
prev_response_lines = len(new_response.split("\n"))
continue
else:
- if not sentence_end.search(msg):
+ if not complete_sentence(msg):
msg += "."
prompt += f"You {msg}\n"
prompt = prompt[-10000:]