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, 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()