From 7f94079d349de13288072c8ee6f56593727a1447 Mon Sep 17 00:00:00 2001 From: Benjamin Bay Date: Mon, 9 Dec 2019 14:52:36 -0800 Subject: [PATCH 1/3] sorted imports --- data/build_training_data.py | 1 + data/make_reddit_data.py | 3 ++- data/scraper.py | 6 ++++-- generator/gpt2/download_model.py | 1 + generator/gpt2/gpt2_generator.py | 13 +++++++------ generator/gpt2/src/encoder.py | 5 +++-- generator/gpt2/src/sample.py | 1 - generator/simple/finetune.py | 4 ++-- other/cacher.py | 3 ++- play.py | 7 +++++-- play_dm.py | 11 +++++++---- story/story_manager.py | 7 ++++--- story/utils.py | 5 +++-- 13 files changed, 41 insertions(+), 26 deletions(-) diff --git a/data/build_training_data.py b/data/build_training_data.py index 6038e75..b4127cb 100644 --- a/data/build_training_data.py +++ b/data/build_training_data.py @@ -1,5 +1,6 @@ import csv import json + from story.utils import * diff --git a/data/make_reddit_data.py b/data/make_reddit_data.py index 964fc9a..3e1cafa 100644 --- a/data/make_reddit_data.py +++ b/data/make_reddit_data.py @@ -1,7 +1,8 @@ import json -from story.utils import * import os +from story.utils import * + def load_stories(file): diff --git a/data/scraper.py b/data/scraper.py index 4318ebe..041b502 100644 --- a/data/scraper.py +++ b/data/scraper.py @@ -1,7 +1,9 @@ +import json +import time + from selenium import webdriver from selenium.webdriver.chrome.options import Options -import time -import json + """ format of tree is diff --git a/generator/gpt2/download_model.py b/generator/gpt2/download_model.py index 0fbc440..e00754c 100644 --- a/generator/gpt2/download_model.py +++ b/generator/gpt2/download_model.py @@ -1,5 +1,6 @@ import os import sys + import requests from tqdm import tqdm diff --git a/generator/gpt2/gpt2_generator.py b/generator/gpt2/gpt2_generator.py index e355b75..cd74562 100644 --- a/generator/gpt2/gpt2_generator.py +++ b/generator/gpt2/gpt2_generator.py @@ -1,14 +1,15 @@ -from story.utils import * +import json +import os import warnings -warnings.filterwarnings("ignore") -import os +import numpy as np import tensorflow as tf +from generator.gpt2.src import encoder, model, sample +from story.utils import * + +warnings.filterwarnings("ignore") tf.compat.v1.logging.set_verbosity(tf.compat.v1.logging.ERROR) -from generator.gpt2.src import sample, encoder, model -import json -import numpy as np class GPT2Generator: diff --git a/generator/gpt2/src/encoder.py b/generator/gpt2/src/encoder.py index 198f19d..918e071 100644 --- a/generator/gpt2/src/encoder.py +++ b/generator/gpt2/src/encoder.py @@ -1,10 +1,11 @@ """Byte pair encoding utilities""" -import os import json -import regex as re +import os from functools import lru_cache +import regex as re + @lru_cache() def bytes_to_unicode(): diff --git a/generator/gpt2/src/sample.py b/generator/gpt2/src/sample.py index 9048fe5..d5d4983 100644 --- a/generator/gpt2/src/sample.py +++ b/generator/gpt2/src/sample.py @@ -1,5 +1,4 @@ import tensorflow as tf - from generator.gpt2.src import model diff --git a/generator/simple/finetune.py b/generator/simple/finetune.py index 35503e5..fd68ef1 100644 --- a/generator/simple/finetune.py +++ b/generator/simple/finetune.py @@ -1,7 +1,7 @@ -import tarfile import os -import gpt_2_simple as gpt2 +import tarfile +import gpt_2_simple as gpt2 model_name = "1558M" if not os.path.isdir(os.path.join("models", model_name)): diff --git a/other/cacher.py b/other/cacher.py index bb629dd..3f17a7b 100644 --- a/other/cacher.py +++ b/other/cacher.py @@ -1,6 +1,7 @@ -from google.cloud import storage import os +from google.cloud import storage + class Cacher: def __init__(self, credentials_file, bucket_name="dungeon-cache"): diff --git a/play.py b/play.py index bda63e6..c4e811d 100644 --- a/play.py +++ b/play.py @@ -1,7 +1,10 @@ -from story.story_manager import * +import os +import sys +import time + from generator.gpt2.gpt2_generator import * +from story.story_manager import * from story.utils import * -import time, sys, os os.environ["TF_CPP_MIN_LOG_LEVEL"] = "3" diff --git a/play_dm.py b/play_dm.py index 3a2a3ae..3d4eaed 100644 --- a/play_dm.py +++ b/play_dm.py @@ -1,9 +1,12 @@ -from story.story_manager import * -from generator.human_dm import * +import os +import sys +import time + from generator.gpt2.gpt2_generator import * -from story.utils import * +from generator.human_dm import * from play import * -import time, sys, os +from story.story_manager import * +from story.utils import * os.environ["TF_CPP_MIN_LOG_LEVEL"] = "3" diff --git a/story/story_manager.py b/story/story_manager.py index daa7aa2..aba3974 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -1,9 +1,10 @@ -from story.utils import * import json +import os +import subprocess import uuid from subprocess import Popen -import subprocess -import os + +from story.utils import * class Story: diff --git a/story/utils.py b/story/utils.py index ad2aee8..0a33dfc 100644 --- a/story/utils.py +++ b/story/utils.py @@ -1,11 +1,12 @@ # coding: utf-8 import re -import yaml from difflib import SequenceMatcher +import yaml +from profanityfilter import ProfanityFilter + YAML_FILE = "story/story_data.yaml" -from profanityfilter import ProfanityFilter with open("story/extra_censored_words.txt", "r") as f: more_words = [l.replace("\n", "") for l in f.readlines()] From 842d89d496b9d9c76fba47e513349630daf62dc5 Mon Sep 17 00:00:00 2001 From: Benjamin Bay <48391872+ben-bay@users.noreply.github.com> Date: Mon, 9 Dec 2019 20:50:36 -0800 Subject: [PATCH 2/3] Terminal state logic (#84) * Update AIDungeon_2.ipynb * Update AIDungeon_2.ipynb * Update AIDungeon_2.ipynb * Fix "Notebook loading error" on colab (#81) * Fix "Notebook loading error" on colab There was a missing comma after the array element on line 46, so I added it. From what I can tell, this was causing the colab to be unable to load. * Update AIDungeon_2.ipynb * Update requirements.txt * improved terminal state logic. --- AIDungeon_2.ipynb | 7 ++++--- requirements.txt | 3 ++- story/utils.py | 51 +++++++++++++++++++---------------------------- 3 files changed, 26 insertions(+), 35 deletions(-) diff --git a/AIDungeon_2.ipynb b/AIDungeon_2.ipynb index 49c66a1..12f9341 100644 --- a/AIDungeon_2.ipynb +++ b/AIDungeon_2.ipynb @@ -32,7 +32,7 @@ "If you want to help, best thing you can do is to **[download this torrent file with game files](https://github.com/nickwalton/AIDungeon/files/3935881/model_v5.torrent.zip)** and **seed it** indefinitely to the best of your ability. This will help new players download this game faster, and discover the vast worlds of AIDungeon2!\n", "\n", "- Follow @nickwalton00 on Twitter for updates on when it will be available again.\n", - "- **[Support AI Dungeon 2](https://www.patreon.com/posts/update-32193930) on Patreon**\n", + "- **[Support AI Dungeon 2](https://www.patreon.com/AIDungeon) on Patreon to help me to continue improving the game with all the awesome ideas I have for its future!**\n", "\n", "## How to play\n", "1. Click \"Tools\"-> \"Settings...\" -> \"Theme\" -> \"Dark\" (optional but recommended)\n", @@ -43,7 +43,8 @@ "\n", "## About\n", "* While you wait you can [read adventures others have had](https://aidungeon.io/)\n", - "* [Read more](https://pcc.cs.byu.edu/2019/11/21/ai-dungeon-2-creating-infinitely-generated-text-adventures-with-deep-learning-language-models/) about how AI Dungeon 2 is made." + "* [Read more](https://pcc.cs.byu.edu/2019/11/21/ai-dungeon-2-creating-infinitely-generated-text-adventures-with-deep-learning-language-models/) about how AI Dungeon 2 is made.", + "- **[Support AI Dungeon 2](https://www.patreon.com/bePatron?u=19115449) on Patreon to help me to continue improving the game with all the awesome ideas I have for its future!**\n" ] }, { @@ -80,4 +81,4 @@ "outputs": [] } ] -} \ No newline at end of file +} diff --git a/requirements.txt b/requirements.txt index 670bb7a..fcaa7a3 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,7 @@ google-cloud-storage +gsutil numpy -regex profanityfilter pyyaml +regex tensorflow==1.15 diff --git a/story/utils.py b/story/utils.py index 0a33dfc..6728cff 100644 --- a/story/utils.py +++ b/story/utils.py @@ -48,41 +48,30 @@ def get_num_options(num): def player_died(text): - - # reg_phrases = ["You[a-zA-Z ]* die.", "you[a-zA-Z ]* die.", "You[a-zA-Z ]* die ", "you[a-zA-Z ]* die ",] - # - # for phrase in reg_phrases: - # reg_expr = re.compile(phrase) - # matches = re.findall(reg_expr, text) - # if len(matches) > 0: - # return True - - dead_phrases = [ - "you die", - "You die", - "you died", - "you are dead", - "You died", - "You are dead", - "You're dead", - "you're dead", - "you have died", - "You have died", - "you bleed out", + """ + TODO: Add in more sophisticated NLP, maybe a custom classifier + trained on hand-labelled data that classifies second-person + statements as resulting in death or not. + """ + lower_text = text.lower() + you_dead_regexps = [ + "you('re| are) (dead|killed|slain|no more|nonexistent)", + "you (die|pass away|perish|suffocate|drown|bleed out)", + "you('ve| have) (died|perished|suffocated|drowned|been (killed|slain))", + "you \w* to death", + "you \w+ yourself to death", ] - for phrase in dead_phrases: - if phrase in text: - return True - return False + return any(re.search(regexp, lower_text) for regexp in you_dead_regexps) def player_won(text): - - won_phrases = ["live happily ever after", "you live forever"] - for phrase in won_phrases: - if phrase in text: - return True - return False + lower_text = text.lower() + won_phrases = [ + "live happily ever after", + "(you)? live (forever|eternally|for eternity)", + "you (are|become|turn into) (a)? (deity|god)", + ] + return any(re.search(regexp, lower_text) for regexp in won_phrases) def remove_profanity(text): From b993cb4cd76da2febe87bb22b87571afdf5e28dd Mon Sep 17 00:00:00 2001 From: Benjamin Bay <48391872+ben-bay@users.noreply.github.com> Date: Mon, 9 Dec 2019 22:43:04 -0800 Subject: [PATCH 3/3] Update utils.py --- story/utils.py | 1 + 1 file changed, 1 insertion(+) diff --git a/story/utils.py b/story/utils.py index 6728cff..31b4fbf 100644 --- a/story/utils.py +++ b/story/utils.py @@ -70,6 +70,7 @@ def player_won(text): "live happily ever after", "(you)? live (forever|eternally|for eternity)", "you (are|become|turn into) (a)? (deity|god)", + "you ((go|get) (in)?to|arrive (at|in)) (heaven|paradise)", ] return any(re.search(regexp, lower_text) for regexp in won_phrases)