diff options
| author | Rose Hogenson <rhogenson@posteo.net> | 2022-02-15 06:22:00 -0800 |
|---|---|---|
| committer | Rose Hogenson <rhogenson@posteo.net> | 2022-02-15 06:22:00 -0800 |
| commit | 87cbdc4ec89fa8d54ab3123449c6397fe235f29c (patch) | |
| tree | eb9b443a32767cc27ced1c080cda92b5e61cf309 /ai_dungeon.py | |
| parent | 6f5fc494ca360e9ffd042425b44c71d95ad3084a (diff) | |
| download | gpt-adventure-87cbdc4ec89fa8d54ab3123449c6397fe235f29c.tar.zst | |
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?
Diffstat (limited to 'ai_dungeon.py')
| -rwxr-xr-x | ai_dungeon.py | 17 |
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: |
