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..31b4fbf 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()] @@ -47,41 +48,31 @@ 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)", + "you ((go|get) (in)?to|arrive (at|in)) (heaven|paradise)", + ] + return any(re.search(regexp, lower_text) for regexp in won_phrases) def remove_profanity(text):