summaryrefslogtreecommitdiffstats
path: root/ai_dungeon.py
diff options
context:
space:
mode:
Diffstat (limited to 'ai_dungeon.py')
-rwxr-xr-xai_dungeon.py172
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()