diff options
| -rwxr-xr-x | ai-dungeon.py | 137 |
1 files changed, 137 insertions, 0 deletions
diff --git a/ai-dungeon.py b/ai-dungeon.py new file mode 100755 index 0000000..faa06f4 --- /dev/null +++ b/ai-dungeon.py @@ -0,0 +1,137 @@ +#!/usr/bin/env nix-shell +#!nix-shell -i python3 -p "python3.withPackages (pkgs: with pkgs; [ transformers pytorch keras ])" + +import argparse +import sys +import transformers +import warnings + + +def trim_sentence(tokenizer: transformers.PreTrainedTokenizer, message: str) -> str: + periodt = tokenizer.encode('.')[0] + tokens = tokenizer.encode(message) + for i in range(len(tokens) - 1, -1, -1): + if tokens[i] == periodt: + return tokenizer.decode(tokens[:i+1]) + return message + + +def generate(model: transformers.TextGenerationPipeline, prompt: str, **kwargs) -> str: + out = model( + prompt, + do_sample=True, + temperature=0.8, + top_k=60, + top_p=0.9, + **kwargs) + return trim_sentence(model.tokenizer, out[0]['generated_text']) + + +def load_model(model: str) -> transformers.TextGenerationPipeline: + 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: + 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." + } + 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: + res = list() + for line in message.split('\n'): + if len(line) < 72: + res.append(line) + continue + split_line = list() + for word in line.split(' '): + if not split_line or len(' '.join(split_line)) + len(word) < 72: + 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: + parser = argparse.ArgumentParser('AI dungeon clone') + parser.add_argument( + '--model', + default='EleutherAI/gpt-neo-125M', + help='Model to use') + args = parser.parse_args() + model = load_model(args.model) + script = pick_flavor() + prologue = 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 msg.endswith('.'): + 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')) + + +if __name__ == '__main__': + main() |
