diff options
| -rwxr-xr-x | ai_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() |
