From 87cbdc4ec89fa8d54ab3123449c6397fe235f29c Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Tue, 15 Feb 2022 06:22:00 -0800 Subject: Use a naive regexp to strip extra output. The sentence tokenizer wasn't correct because not all sentences end in periods. The current regexp doesn't handle punctuation inside words such as apostrophes or abbreviations, but we're trying it. Maybe a tokenizer would be better? --- ai_dungeon.py | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) (limited to 'ai_dungeon.py') 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: -- cgit v1.3.1