From 169efa63904e3850a5bd698e4f62d870488305b9 Mon Sep 17 00:00:00 2001 From: Rose Hogenson Date: Wed, 16 Feb 2022 05:24:37 +0000 Subject: Rename to GPT adventure. --- ai_dungeon.py | 172 ------------------------------------------------------- gpt_adventure.py | 172 +++++++++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 172 insertions(+), 172 deletions(-) delete mode 100755 ai_dungeon.py create mode 100755 gpt_adventure.py diff --git a/ai_dungeon.py b/ai_dungeon.py deleted file mode 100755 index b0cc64a..0000000 --- a/ai_dungeon.py +++ /dev/null @@ -1,172 +0,0 @@ -#!/usr/bin/env nix-shell -#!nix-shell -i python3 -p "python3.withPackages (pkgs: with pkgs; [ keras nltk pytorch transformers ])" -"""AI Dungeon is a text-adventure style game powered by AI.""" - -import argparse -import re -import shutil -import sys - -import nltk -from nltk import tokenize -import transformers - - -def load_nltk(): - try: - tokenize.sent_tokenize('') - except LookupError: - nltk.download("punkt") - - -def complete_sentence(snippet: str) -> bool: - sentences = tokenize.sent_tokenize(snippet) - return len(tokenize.sent_tokenize(sentences[-1] + ' extra')) == 2 - - -def trim_sentence(message: str) -> str: - """Remove extra output after the last period.""" - sentences = tokenize.sent_tokenize(message) - if complete_sentence(sentences[-1]) or len(sentences) == 1: - return message - return " ".join(sentences[:-1]) - - -def generate(model: transformers.TextGenerationPipeline, prompt: str, **kwargs) -> str: - """Generate text from a text generator.""" - out = model(prompt, do_sample=True, temperature=0.9, top_k=60, top_p=0.9, **kwargs) - return trim_sentence(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." - ), - "doctor": ( - "Your name is Amelia Plenn. Today is your doctors appointment.\n" - "You're kind of not looking forward to it because your doctor\n" - "is kind of weird, but today when you get to the doctor's office\n" - "it's a different person than usual.\n\n" - "The new doctor says to you, \"Hello, my name is Doctor Cockman.\n" - "It's my pleasure to see you today.\"\n\n" - "Doctor Cockman is very handsome, and you feel yourself blush as\n" - "he leads you to the doctor's seat. You sit in the seat and he\n" - "runs his warm hands up and down your arms.\n\n" - "\"Are you feeling alright?\" Doctor Cockman asks." - ) - } - 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 terminal width characters.""" - width = shutil.get_terminal_size().columns - res = [] - for line in message.split("\n"): - if len(line) < width: - res.append(line) - continue - split_line = [] - for word in line.split(" "): - if not split_line or len(" ".join(split_line)) + len(word) < width: - 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.""" - load_nltk() - parser = argparse.ArgumentParser("AI dungeon clone") - parser.add_argument( - "--model", default="gpt2", help="Model to use" - ) - args = parser.parse_args() - model = load_model(args.model) - script = pick_flavor() - prologue = wrap(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 complete_sentence(msg): - 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() diff --git a/gpt_adventure.py b/gpt_adventure.py new file mode 100755 index 0000000..e4dceb5 --- /dev/null +++ b/gpt_adventure.py @@ -0,0 +1,172 @@ +#!/usr/bin/env nix-shell +#!nix-shell -i python3 -p "python3.withPackages (pkgs: with pkgs; [ keras nltk pytorch transformers ])" +"""GPT adventure is a text-adventure style game powered by AI.""" + +import argparse +import re +import shutil +import sys + +import nltk +from nltk import tokenize +import transformers + + +def load_nltk(): + try: + tokenize.sent_tokenize('') + except LookupError: + nltk.download("punkt") + + +def complete_sentence(snippet: str) -> bool: + sentences = tokenize.sent_tokenize(snippet) + return len(tokenize.sent_tokenize(sentences[-1] + ' extra')) == 2 + + +def trim_sentence(message: str) -> str: + """Remove extra output after the last period.""" + sentences = tokenize.sent_tokenize(message) + if complete_sentence(sentences[-1]) or len(sentences) == 1: + return message + return " ".join(sentences[:-1]) + + +def generate(model: transformers.TextGenerationPipeline, prompt: str, **kwargs) -> str: + """Generate text from a text generator.""" + out = model(prompt, do_sample=True, temperature=0.9, top_k=60, top_p=0.9, **kwargs) + return trim_sentence(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." + ), + "doctor": ( + "Your name is Amelia Plenn. Today is your doctors appointment.\n" + "You're kind of not looking forward to it because your doctor\n" + "is kind of weird, but today when you get to the doctor's office\n" + "it's a different person than usual.\n\n" + "The new doctor says to you, \"Hello, my name is Doctor Cockman.\n" + "It's my pleasure to see you today.\"\n\n" + "Doctor Cockman is very handsome, and you feel yourself blush as\n" + "he leads you to the doctor's seat. You sit in the seat and he\n" + "runs his warm hands up and down your arms.\n\n" + "\"Are you feeling alright?\" Doctor Cockman asks." + ) + } + 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 terminal width characters.""" + width = shutil.get_terminal_size().columns + res = [] + for line in message.split("\n"): + if len(line) < width: + res.append(line) + continue + split_line = [] + for word in line.split(" "): + if not split_line or len(" ".join(split_line)) + len(word) < width: + 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.""" + load_nltk() + parser = argparse.ArgumentParser("AI dungeon clone") + parser.add_argument( + "--model", default="gpt2", help="Model to use" + ) + args = parser.parse_args() + model = load_model(args.model) + script = pick_flavor() + prologue = wrap(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 complete_sentence(msg): + 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