summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorRose Hogenson <rhogenson@posteo.net>2022-02-14 21:22:40 -0800
committerRose Hogenson <rhogenson@posteo.net>2022-02-14 21:22:40 -0800
commit9d2be55167f84a287a2c2d9b8aff54dd6e3aade1 (patch)
tree0d91d919ec556413184d2bb9d28327d61c8b1c71
parent60857f10f5a39102047ff07268ecfd1434957668 (diff)
downloadgpt-adventure-9d2be55167f84a287a2c2d9b8aff54dd6e3aade1.tar.zst
Lint the code and clean up a bit.
-rwxr-xr-xai_dungeon.py (renamed from ai-dungeon.py)21
1 files changed, 15 insertions, 6 deletions
diff --git a/ai-dungeon.py b/ai_dungeon.py
index faa06f4..f1aabf4 100755
--- a/ai-dungeon.py
+++ b/ai_dungeon.py
@@ -1,13 +1,14 @@
#!/usr/bin/env nix-shell
#!nix-shell -i python3 -p "python3.withPackages (pkgs: with pkgs; [ transformers pytorch keras ])"
+"""AI Dungeon is a text-adventure style game powered by AI."""
import argparse
import sys
import transformers
-import warnings
def trim_sentence(tokenizer: transformers.PreTrainedTokenizer, 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):
@@ -17,6 +18,7 @@ def trim_sentence(tokenizer: transformers.PreTrainedTokenizer, message: str) ->
def generate(model: transformers.TextGenerationPipeline, prompt: str, **kwargs) -> str:
+ """Generate text from a text generator."""
out = model(
prompt,
do_sample=True,
@@ -28,6 +30,7 @@ def generate(model: transformers.TextGenerationPipeline, prompt: str, **kwargs)
def load_model(model: str) -> transformers.TextGenerationPipeline:
+ """Load a model by name."""
tokenizer = transformers.AutoTokenizer.from_pretrained(model)
return transformers.pipeline(
'text-generation',
@@ -37,6 +40,7 @@ def load_model(model: str) -> transformers.TextGenerationPipeline:
def pick_flavor() -> str:
+ """Query the user for what scenario they want to play."""
flavors = {
'fantasy': "You are a wizard named Megumin from the kingdom\n"
"of Larion. You have in your inventory a wizard's staff\n"
@@ -70,12 +74,13 @@ def pick_flavor() -> str:
def wrap(message: str) -> str:
- res = list()
+ """Wrap long lines to 72 characters."""
+ res = []
for line in message.split('\n'):
if len(line) < 72:
res.append(line)
continue
- split_line = list()
+ split_line = []
for word in line.split(' '):
if not split_line or len(' '.join(split_line)) + len(word) < 72:
split_line.append(word)
@@ -87,6 +92,7 @@ def wrap(message: str) -> str:
def main() -> None:
+ """Run the main game loop."""
parser = argparse.ArgumentParser('AI dungeon clone')
parser.add_argument(
'--model',
@@ -95,7 +101,11 @@ def main() -> None:
args = parser.parse_args()
model = load_model(args.model)
script = pick_flavor()
- prologue = generate(model, script, max_new_tokens=20, forced_eos_token_id=model.tokenizer.eos_token_id)
+ prologue = generate(
+ model,
+ script,
+ max_new_tokens=20,
+ forced_eos_token_id=model.tokenizer.eos_token_id)
print(prologue)
prompt = prologue
prev_response_length = 0
@@ -133,5 +143,4 @@ def main() -> None:
prev_response_lines = len(response.split('\n'))
-if __name__ == '__main__':
- main()
+main()