summaryrefslogtreecommitdiffstats
path: root/gpt_adventure.py
diff options
context:
space:
mode:
authorRose Hogenson <rhogenson@posteo.net>2022-02-16 05:24:37 +0000
committerRose Hogenson <rhogenson@posteo.net>2022-02-16 05:24:37 +0000
commit169efa63904e3850a5bd698e4f62d870488305b9 (patch)
treebac6c074c00b782e3de9c4c35bffeec8698a6be1 /gpt_adventure.py
parentAdd a doctor scenario. (diff)
downloadgpt-adventure-169efa63904e3850a5bd698e4f62d870488305b9.tar.zst
Rename to GPT adventure.
Diffstat (limited to 'gpt_adventure.py')
-rwxr-xr-xgpt_adventure.py172
1 files changed, 172 insertions, 0 deletions
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()