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