diff options
| author | Rose Hogenson <rhogenson@posteo.net> | 2023-02-14 16:00:12 -0800 |
|---|---|---|
| committer | Rose Hogenson <rhogenson@posteo.net> | 2023-02-14 16:00:12 -0800 |
| commit | f33a659fb121b9d2424b7440302d11265bd508da (patch) | |
| tree | 803616bba700d18768c363f89b2e5c893f9356a4 | |
| parent | a80c7c8339e5074e7883469b3eb4739d8867f09d (diff) | |
| download | gpt-adventure-f33a659fb121b9d2424b7440302d11265bd508da.tar.zst | |
Simplify GPT adventure.
Or did I make it more complicated?
| -rwxr-xr-x | gpt_adventure.py | 248 |
1 files changed, 104 insertions, 144 deletions
diff --git a/gpt_adventure.py b/gpt_adventure.py index 65ff09c..ce77222 100755 --- a/gpt_adventure.py +++ b/gpt_adventure.py @@ -2,211 +2,171 @@ #!nix-shell -i python3 -p "python3.withPackages (pkgs: with pkgs; [ keras pytorch transformers ])" """GPT adventure is a text-adventure style game powered by AI.""" +from typing import List, Optional import argparse -import shutil -import sys import transformers -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, - return_full_text=False, - **kwargs, - )[0]["generated_text"] - first_line = out.split("\n", 1)[0].strip() - if first_line: - return first_line - return "haha i'm not sure what you mean" - - -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, - ) +class Model: + def __init__(self, name: str): + self.name = name + self.tokenizer = transformers.AutoTokenizer.from_pretrained(name) + self.model = transformers.pipeline( + "text-generation", + tokenizer=self.tokenizer, + model=transformers.AutoModelForCausalLM.from_pretrained(name), + pad_token_id=self.tokenizer.eos_token_id, + ) + def generate(self, prompt: str) -> str: + """Generate text from a text generator.""" + out = self.model( + prompt, + do_sample=True, + temperature=0.9, + top_k=60, + top_p=0.9, + return_full_text=False, + max_new_tokens=50, + )[0]["generated_text"] + first_line = out.split("\n", 1)[0].strip() + if first_line: + return first_line + return "haha i'm not sure what you mean" -class Persona: - def __init__(self, name: str, bio: str): - self.name = name - self.bio = bio - self.prologue = f"{name}'s Persona: {bio}\n<START>\n" + def token_length(self, msg: str) -> int: + return len(self.tokenizer(msg)["input_ids"]) class State: - def __init__(self, persona: Persona, history: str): - self.persona = persona + def __init__(self, model: Model, character: str, bio: str, history: List[str], filename: Optional[str]): + self.model = model + self.character = character + self.bio = bio self.history = history - self.prev_response_length = 0 + self.filename = filename + self.prologue = f"{character}'s Persona: {bio}\n<START>\n" def add_history(self, who: str, msg: str) -> None: - fmt_msg = f"{who}: {msg}\n" - self.history += fmt_msg - self.prev_response_length = len(fmt_msg) + self.history.append(f"{who}: {msg}") def rollback_history(self) -> None: - self.history = self.history[:-self.prev_response_length] - self.prev_response_length = 0 + self.history = self.history[:-1] def prompt(self) -> str: - max_prompt_size = 10000 - if len(self.history) > max_prompt_size: - self.history = self.history[-(max_prompt_size+1):].split("\n", 1)[1] - return f"{self.persona.prologue}{self.history}{self.persona.name}: " + return self.prologue + "\n".join(self.history) + f"\n{self.character}: " + + def trim_history(self): + while self.history and self.model.token_length(self.prompt()) > 1024: + self.history = self.history[1:] + + def print_context(self) -> None: + for elem in self.history[-10:]: + if elem.startswith("You: "): + print("> " + elem.removeprefix("You: ")) + else: + print(elem.removeprefix(f"{self.character}: ")) + + def autosave(self) -> None: + if self.filename: + self.save(self.filename) def save(self, filename: str) -> None: + self.filename = filename with open(filename, "w") as f: - print(self.persona.name, file=f) - print(self.persona.bio, file=f) - print(self.history, file=f, end='') + print(self.model.name, file=f) + print(self.character, file=f) + print(self.bio, file=f) + for elem in self.history: + print(elem, file=f) def load(filename: str) -> State: with open(filename, "r") as f: - persona = f.readline().strip() + model = f.readline().strip() + character = f.readline().strip() bio = f.readline().strip() - history = f.read() - return State(Persona(persona, bio), history) - - -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) - - -class Term: - def __init__(self): - self.rewind_point = 0 + history = f.read().strip().split("\n") + return State(Model(model), character, bio, history, filename) - def input(self, prompt: str) -> str: - self.rewind_point += 1 - return input(prompt) - def print(self, msg: str) -> None: - wrapped_msg = wrap(msg) - print(wrapped_msg) - self.rewind_point += wrapped_msg.count("\n") + 1 +def clear() -> None: + print(f"\033[2J", flush=True) - def set_rewind_point(self) -> None: - self.rewind_point = 0 - def rewind(self) -> None: - print(f"\033[{self.rewind_point}A\033[J\r", end="") - self.set_rewind_point() - - -flavors = ( - State(Persona("Rin", - "Tohsaka Rin is a feisty and independent mage who can come across as " - "rude and unlikeable at first. She can be sweet and caring, but it " - "takes a lot to break down her guard."), - - "Rin: What are you looking at... idiot?\n" - "You: Umm nothing...\n" - "Rin: That's right, you're nothing. You are less than a piece of trash.\n" - "You: Rin, do you want to go demon-hunting some time?\n" - "Rin: Well, maybe. But not because I like you or anything.\n" - "You: Right, we'll just go as friends.\n" - "Rin: Or maybe acquaintances...\n" - ), -) - - -def pick_flavor() -> State: - """Query the user for what scenario they want to play.""" - for i, state in enumerate(flavors): - print(f"{i}.\t{state.persona.name}") - print(f"{len(flavors)}.\tcustom") - while True: - choice = input("Choose a persona: ") - try: - choice_int = int(choice) - except ValueError: - print("Input must be an integer.") - continue - if choice_int == len(flavors): - name = input("Enter your character's name: ") - bio = input("Enter your character's persona: ") - return State(Persona(name, bio), "") - try: - return flavors[choice_int] - except IndexError: - print("Index out of bounds.") - continue +def new_character(model: Model) -> State: + print("Creating a new character.") + print("You may have to heavily edit (/edit) the AI's first few responses") + print("in order for it to learn the appropriate tone.") + name = input("Enter your character's name: ") + bio = input("Enter a few sentences describing your character: ") + return State(model, name, bio, [], None) def main() -> None: """Run the main game loop.""" - parser = argparse.ArgumentParser("AI dungeon clone") + parser = argparse.ArgumentParser("AI girlfriend") parser.add_argument("--model", default="PygmalionAI/pygmalion-6b", help="Model to use") parser.add_argument("save_file", nargs="?", default="") args = parser.parse_args() - model = load_model(args.model) - if args.save_file: state = load(args.save_file) - print("\n".join(state.history.rsplit("\n", 11)[-11:]).strip()) + state.print_context() else: - state = pick_flavor() + state = new_character(Model(args.model)) - term = Term() while True: try: - msg = term.input("> ") + msg = input("> ") except EOFError: break - if msg == "/retry": + if msg == "/help": + commands = ( + "/help", + "/retry", + "/edit", + "/undo", + "/save", + ) + print("Commands:") + for cmd in commands: + print(f"\t{cmd}") + elif msg == "/retry": state.rollback_history() - term.rewind() + clear() + state.print_context() elif msg == "/edit": state.rollback_history() - term.rewind() + clear() + state.print_context() - new_response = term.input("") - state.add_history(state.persona.name, new_response) + new_response = input("") + state.add_history(state.character, new_response) + continue + elif msg == "/undo": + state.rollback_history() + state.rollback_history() + clear() + state.print_context() continue elif msg.startswith("/save"): if " " not in msg: - term.print("Usage: /save <filename>") + print("Usage: /save <filename>") continue state.save(msg.split(" ", 1)[1]) continue else: state.add_history("You", msg) - term.set_rewind_point() + state.trim_history() - response = generate(model, state.prompt(), max_new_tokens=50) + response = state.model.generate(state.prompt()) if not response: response = "haha I'm not sure what you mean" - term.print(response) - state.add_history(state.persona.name, response) + print(response) + state.add_history(state.character, response) main() |
