diff options
| author | Rose Hogenson <rhogenson@posteo.net> | 2022-02-14 21:22:40 -0800 |
|---|---|---|
| committer | Rose Hogenson <rhogenson@posteo.net> | 2022-02-14 21:22:40 -0800 |
| commit | 9d2be55167f84a287a2c2d9b8aff54dd6e3aade1 (patch) | |
| tree | 0d91d919ec556413184d2bb9d28327d61c8b1c71 /ai-dungeon.py | |
| parent | Add initial AI dungeon clone. (diff) | |
| download | gpt-adventure-9d2be55167f84a287a2c2d9b8aff54dd6e3aade1.tar.zst | |
Lint the code and clean up a bit.
Diffstat (limited to 'ai-dungeon.py')
| -rwxr-xr-x | ai-dungeon.py | 137 |
1 files changed, 0 insertions, 137 deletions
diff --git a/ai-dungeon.py b/ai-dungeon.py deleted file mode 100755 index faa06f4..0000000 --- a/ai-dungeon.py +++ /dev/null @@ -1,137 +0,0 @@ -#!/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() |
