From 9d2be55167f84a287a2c2d9b8aff54dd6e3aade1 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Mon, 14 Feb 2022 21:22:40 -0800 Subject: Lint the code and clean up a bit. --- ai-dungeon.py | 137 ------------------------------------------------------ ai_dungeon.py | 146 ++++++++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 146 insertions(+), 137 deletions(-) delete mode 100755 ai-dungeon.py create mode 100755 ai_dungeon.py diff --git a/ai-dungeon.py b/ai-dungeon.py deleted file mode 100755 index faa06f4..0000000 --- a/ai-dungeon.py +++ /dev/null @@ -1,137 +0,0 @@ -#!/usr/bin/env nix-shell -#!nix-shell -i python3 -p "python3.withPackages (pkgs: with pkgs; [ transformers pytorch keras ])" - -import argparse -import sys -import transformers -import warnings - - -def trim_sentence(tokenizer: transformers.PreTrainedTokenizer, message: str) -> str: - 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]) - return message - - -def generate(model: transformers.TextGenerationPipeline, prompt: str, **kwargs) -> str: - 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']) - - -def load_model(model: str) -> transformers.TextGenerationPipeline: - tokenizer = transformers.AutoTokenizer.from_pretrained(model) - return transformers.pipeline( - 'text-generation', - tokenizer=tokenizer, - model=transformers.AutoModelForCausalLM.from_pretrained(model), - pad_token_id=tokenizer.eos_token_id) - - -def pick_flavor() -> str: - flavors = { - 'fantasy': "You are a wizard named Megumin from the kingdom\n" - "of Larion. You have in your inventory a wizard's staff\n" - "and a spellbook. You are arriving after a day's travel\n" - "at an enchanted tower where there's rumors of gold.\n" - "You walk up to the entrance of the tower.", - 'post apocalyptic': "You are a machinist named Azariel, living\n" - "in the city of New New York. It's been almost 10 years\n" - "since the bombs fell, but you still remember it as if\n" - "it were yesterday. You push these thoughts out of your\n" - "mind and focus on the task at hand: finding a water\n" - "purifier for your settlement. You arrive at an\n" - "abandoned settlement to the east of your home." - } - choices = [] - for i, (scenario, script) in enumerate(flavors.items()): - print(f'{i}.\t{scenario}') - choices.append(script) - while True: - choice = input('Choose a scenario: ') - try: - choice_int = int(choice) - except ValueError: - print('Input must be an integer.') - continue - try: - return choices[choice_int] - except IndexError: - print('Index out of bounds.') - continue - - -def wrap(message: str) -> str: - res = list() - for line in message.split('\n'): - if len(line) < 72: - res.append(line) - continue - split_line = list() - for word in line.split(' '): - if not split_line or len(' '.join(split_line)) + len(word) < 72: - split_line.append(word) - continue - res.append(' '.join(split_line)) - split_line = [word] - res.append(' '.join(split_line)) - return '\n'.join(res) - - -def main() -> None: - parser = argparse.ArgumentParser('AI dungeon clone') - parser.add_argument( - '--model', - default='EleutherAI/gpt-neo-125M', - help='Model to use') - 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) - print(prologue) - prompt = prologue - prev_response_length = 0 - prev_response_lines = 0 - while True: - try: - msg = input('> You ') - except EOFError: - break - if msg == '/retry': - prompt = prompt[:-prev_response_length] - print(f'\033[{prev_response_lines}A\033[J', end='') - elif msg == '/edit': - prompt = prompt[:-prev_response_length] - print(f'\033[{prev_response_lines}A\033[J\r', end='') - new_response = sys.stdin.read() - prompt += new_response - prev_response_length = len(new_response) - prev_response_lines = len(new_response.split('\n')) - continue - else: - if not msg.endswith('.'): - msg += '.' - prompt += f'You {msg}\n' - prompt = prompt[-10000:] - response = wrap(generate( - model, - prompt, - return_full_text=False, - max_new_tokens=50).strip()) - print(response) - response += '\n' - prompt += response - prev_response_length = len(response) - prev_response_lines = len(response.split('\n')) - - -if __name__ == '__main__': - main() diff --git a/ai_dungeon.py b/ai_dungeon.py new file mode 100755 index 0000000..f1aabf4 --- /dev/null +++ b/ai_dungeon.py @@ -0,0 +1,146 @@ +#!/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 + + +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): + if tokens[i] == periodt: + return tokenizer.decode(tokens[:i+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']) + + +def load_model(model: str) -> transformers.TextGenerationPipeline: + """Load a model by name.""" + tokenizer = transformers.AutoTokenizer.from_pretrained(model) + return transformers.pipeline( + 'text-generation', + tokenizer=tokenizer, + model=transformers.AutoModelForCausalLM.from_pretrained(model), + pad_token_id=tokenizer.eos_token_id) + + +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" + "and a spellbook. You are arriving after a day's travel\n" + "at an enchanted tower where there's rumors of gold.\n" + "You walk up to the entrance of the tower.", + 'post apocalyptic': "You are a machinist named Azariel, living\n" + "in the city of New New York. It's been almost 10 years\n" + "since the bombs fell, but you still remember it as if\n" + "it were yesterday. You push these thoughts out of your\n" + "mind and focus on the task at hand: finding a water\n" + "purifier for your settlement. You arrive at an\n" + "abandoned settlement to the east of your home." + } + choices = [] + for i, (scenario, script) in enumerate(flavors.items()): + print(f'{i}.\t{scenario}') + choices.append(script) + while True: + choice = input('Choose a scenario: ') + try: + choice_int = int(choice) + except ValueError: + print('Input must be an integer.') + continue + try: + return choices[choice_int] + except IndexError: + print('Index out of bounds.') + continue + + +def wrap(message: str) -> str: + """Wrap long lines to 72 characters.""" + res = [] + for line in message.split('\n'): + if len(line) < 72: + res.append(line) + continue + split_line = [] + for word in line.split(' '): + if not split_line or len(' '.join(split_line)) + len(word) < 72: + split_line.append(word) + continue + res.append(' '.join(split_line)) + split_line = [word] + res.append(' '.join(split_line)) + return '\n'.join(res) + + +def main() -> None: + """Run the main game loop.""" + parser = argparse.ArgumentParser('AI dungeon clone') + parser.add_argument( + '--model', + default='EleutherAI/gpt-neo-125M', + help='Model to use') + 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) + print(prologue) + prompt = prologue + prev_response_length = 0 + prev_response_lines = 0 + while True: + try: + msg = input('> You ') + except EOFError: + break + if msg == '/retry': + prompt = prompt[:-prev_response_length] + print(f'\033[{prev_response_lines}A\033[J', end='') + elif msg == '/edit': + prompt = prompt[:-prev_response_length] + print(f'\033[{prev_response_lines}A\033[J\r', end='') + new_response = sys.stdin.read() + prompt += new_response + prev_response_length = len(new_response) + prev_response_lines = len(new_response.split('\n')) + continue + else: + if not msg.endswith('.'): + msg += '.' + prompt += f'You {msg}\n' + prompt = prompt[-10000:] + response = wrap(generate( + model, + prompt, + return_full_text=False, + max_new_tokens=50).strip()) + print(response) + response += '\n' + prompt += response + prev_response_length = len(response) + prev_response_lines = len(response.split('\n')) + + +main() -- cgit v1.3.1