diff options
| author | Rose Hogenson <rhogenson@posteo.net> | 2022-02-16 05:24:37 +0000 |
|---|---|---|
| committer | Rose Hogenson <rhogenson@posteo.net> | 2022-02-16 05:24:37 +0000 |
| commit | 169efa63904e3850a5bd698e4f62d870488305b9 (patch) | |
| tree | bac6c074c00b782e3de9c4c35bffeec8698a6be1 /ai_dungeon.py | |
| parent | Add a doctor scenario. (diff) | |
| download | gpt-adventure-169efa63904e3850a5bd698e4f62d870488305b9.tar.zst | |
Rename to GPT adventure.
Diffstat (limited to 'ai_dungeon.py')
| -rwxr-xr-x | ai_dungeon.py | 172 |
1 files changed, 0 insertions, 172 deletions
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() |
