From 69d1151e868fcb2ee906b5eea13df73dadd9cd15 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 06:39:20 +0800 Subject: [PATCH 01/11] remove endoftext etc --- gpt2generator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gpt2generator.py b/gpt2generator.py index 7a9bfed..e829789 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -202,7 +202,7 @@ class GPT2Generator: out = out[:, len(context_tokens) :].tolist() for o in out: generated += 1 - text = self.tokenizer.decode(o, clean_up_tokenization_spaces=True) + text = self.tokenizer.decode(o, clean_up_tokenization_spaces=True, skip_special_tokens=True) if self.stop_token: index = text.find(self.stop_token) if index == -1: From 57432290eaeec6a4a94ced70e30662c2a5e5149b Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 07:10:23 +0800 Subject: [PATCH 02/11] add stop token to speed up --- gpt2generator.py | 1 + 1 file changed, 1 insertion(+) diff --git a/gpt2generator.py b/gpt2generator.py index e829789..59fe29d 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -116,6 +116,7 @@ class GPT2Generator: self.batch_size = 1 self.stop_token = None self.max_history_tokens = 256 + self.stop_token = '<|endoftext|>' self.model_name = "pytorch-gpt2-xl-aid2-v5" self.model_dir = "models" From 90fdedcfc21d2bacd8461d5f5cd92ee55d407094 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 07:10:57 +0800 Subject: [PATCH 03/11] avoid loop --- gpt2generator.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/gpt2generator.py b/gpt2generator.py index 59fe29d..52fd2a4 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -224,6 +224,6 @@ class GPT2Generator: result = text result = self.result_replace(result) if len(result) == 0: - return self.generate(prompt) - + logger.warn("Model generated empty text %s.", result) + # return self.generate(prompt) # Woah recursion! return result From 535c4cf9d14c5de5f2169cf85352d1e8a6b8cf2e Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 07:37:38 +0800 Subject: [PATCH 04/11] sample history when it's too long --- story/story_manager.py | 27 +++++++++++++++++---------- 1 file changed, 17 insertions(+), 10 deletions(-) diff --git a/story/story_manager.py b/story/story_manager.py index d8e05f3..4c4322e 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -64,7 +64,7 @@ class Story: def add_to_story(self, action, story_block): self.actions.append(action) self.results.append(story_block) - if (len(str(self)) > 3900): # (Fix some mem errors. From RTech, max story of 3900 characters for GTX 2080 ti 11GB + if len(self.actions) > 10000: self.actions.pop(1) self.results.pop(1) @@ -72,17 +72,24 @@ class Story: mem_ind = self.memory if len(self.results) < 2: - latest_result = self.story_start + latest_results = [self.story_start] else: - latest_result = self.context - while mem_ind > 0: + latest_results = [self.context] + latest_result = '' - if len(self.results) >= mem_ind: - latest_result += self.actions[-mem_ind] + self.results[-mem_ind] - - mem_ind -= 1 - - return latest_result + if mem_ind < len(self.results): + # When we have to much history we will take the last 10, and sample randomly from the rest + # first take last mem_ind//2 + all_inds = list(range(len(self.results))) + last = all_inds[-mem_ind//2:] + first = all_inds[:mem_ind//2] + inds = sorted(last + random.sample(first, mem_ind//2)) + else: + inds = range(len(self.results)) + logger.debug("Using history indices %s", inds) + for i in inds: + latest_result += self.actions[i] + self.results[i] + return latest_results + [latest_result] def __str__(self): story_list = [self.story_start] From 5c74a923fff4df76f9c08f9afd73a711c5379673 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 07:38:10 +0800 Subject: [PATCH 05/11] rearrange for quick dev --- play.py | 10 ++++++---- 1 file changed, 6 insertions(+), 4 deletions(-) diff --git a/play.py b/play.py index a8d0f40..9962730 100644 --- a/play.py +++ b/play.py @@ -109,8 +109,7 @@ class AIPlayer: return clean_suggested_action(result_raw, min_length=settings.getint('action-min-length')) -def play(): - generator = getGenerator() +def play(generator): story_manager = UnconstrainedStoryManager(generator) ai_player = AIPlayer(generator) print("\n") @@ -317,5 +316,8 @@ def play(): colPrint("Sorry about that...where were we?", colors["query"]) colPrint(result, colors["ai-text"]) -#TODO: there's no reason for this to be enclosed in a function -play() + +# This is here for rapid development, without reloading the model. You import play into a jupyternotebook with autoreload +if __name__ == "__main__": + generator = getGenerator() + play(generator) From 7b95fa7e7938185fb3a2ff6a230bd1c485e9de1b Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 07:39:59 +0800 Subject: [PATCH 06/11] better truncation of context --- gpt2generator.py | 24 +++++++++++++----------- play.py | 14 ++++++++------ story/story_manager.py | 8 ++++---- 3 files changed, 25 insertions(+), 21 deletions(-) diff --git a/gpt2generator.py b/gpt2generator.py index 52fd2a4..629b841 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -100,6 +100,11 @@ def sample_sequence( generated = torch.cat((generated, next_token), dim=1) return generated +def truncate_multiple_sequences(seqs, max_len=100): + """Truncate multiple sequences, longest first, removing first.""" + while sum(len(s) for s in seqs) > max_len: + longest = sorted(seqs, key=len, reverse=True)[0] + longest.pop(0) class GPT2Generator: def __init__( @@ -114,8 +119,7 @@ class GPT2Generator: self.dtype = torch.float32 if CPU else torch.float16 self.repetition_penalty = repetition_penalty self.batch_size = 1 - self.stop_token = None - self.max_history_tokens = 256 + self.max_history_tokens = 1024 - generate_num self.stop_token = '<|endoftext|>' self.model_name = "pytorch-gpt2-xl-aid2-v5" @@ -184,14 +188,12 @@ class GPT2Generator: return result def generate_raw(self, prompt, generate_num=None, temperature=None): - context_tokens = self.tokenizer.encode(prompt, add_special_tokens=False) - # TODO instead of taking last 1024, take first X and last Y - # crop context to avoid going of the GPT2 max context size of 1024 - if len(context_tokens) > self.max_history_tokens: - # FIXME it would be better to pass in a list of strings so we can cut some out, and a truncation strategy https://github.com/huggingface/transformers/blob/ce50305e5b8c8748b81b0c8f5539a337b6a995b9/src/transformers/tokenization_utils.py#L791 - first = self.max_history_tokens // 4 - last = self.max_history_tokens - first - context_tokens = context_tokens[:first] + context_tokens[-last:] + # the prompt is a list of strings, encode each one tok tokens, then truncate the longest ones + context_tokens = [self.tokenizer.encode(p, add_special_tokens=False, max_length=self.max_history_tokens) for p in prompt] + truncate_multiple_sequences(context_tokens, self.max_history_tokens) + context_tokens = list(itertools.chain(*context_tokens)) + + logger.debug("Text passing into model %s", self.tokenizer.decode(o, clean_up_tokenization_spaces=True, skip_special_tokens=True)) generated = 0 for _ in range(self.samples // self.batch_size): @@ -213,7 +215,7 @@ class GPT2Generator: def generate(self, prompt, options=None, seed=1): - prompt = self.prompt_replace(prompt) + prompt = [self.prompt_replace(p) for p in prompt] logger.debug("Prompt is: `%s`", repr(prompt)) diff --git a/play.py b/play.py index 9962730..15348d7 100644 --- a/play.py +++ b/play.py @@ -167,16 +167,14 @@ def play(generator): if settings.getint('action-alternatives') > 0: #TODO change this to two messages for different colors - action_prompt = ( - story_manager.story.results[-1] - if story_manager.story.results - else "\nWhat do you do now?" - ) + "\n>" suggested_actions = [] colPrint('Suggested actions:', colors['selection-value']) action_suggestion_lines = 1 for i in range(settings.getint('action-alternatives')): # FIXME it might be better to pass in a longer history + story_manager.story_context() # This should be within the loop as it has a random sampling element + action_prompt[-1] += '> ' + logger.debug("action_prompt %s", action_prompt) suggested_action = ai_player.get_action(action_prompt) suggested_actions.append(suggested_action) suggestion = '{}> {}'.format(i, suggested_action) @@ -238,7 +236,12 @@ def play(generator): # Options to select a suggestion action if action in [str(i) for i in range(len(suggested_actions))]: action = suggested_actions[int(action)] + + action = action.strip() + # Crop actions to a max length + action = action[:4096] + if action != "": # Roll a 20 sided dice to make things interesting @@ -259,7 +262,6 @@ def play(generator): else: action = "You say " + action else: - action = action.strip() action = first_to_second_person(action) if not action.lower().startswith("you ") and not action.lower().startswith("i "): action = action[0].lower() + action[1:] diff --git a/story/story_manager.py b/story/story_manager.py index 4c4322e..cc52a69 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -3,7 +3,7 @@ import os import subprocess import uuid from subprocess import Popen - +import random from story.utils import * @@ -164,7 +164,7 @@ class StoryManager: def start_new_story( self, story_prompt, context="", game_state=None, upload_story=False ): - block = self.generator.generate(context + story_prompt) + block = self.generator.generate([context, story_prompt]) block = cut_trailing_sentence(block) self.story = Story( context + story_prompt + block, @@ -210,7 +210,7 @@ class UnconstrainedStoryManager(StoryManager): return result def generate_result(self, action): - block = self.generator.generate(self.story_context() + action) + block = self.generator.generate(self.story_context()+[action]) return block @@ -321,7 +321,7 @@ class ConstrainedStoryManager(StoryManager): def generate_action_result(self, prompt, phrase, options=None): action_result = ( - phrase + " " + self.generator.generate(prompt + " " + phrase, options) + phrase + " " + self.generator.generate(prompt + [phrase], options) ) action, result = split_first_sentence(action_result) return action, result From 291f4c0dd070f7dcc693cb4233012b6437af1a26 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 08:02:01 +0800 Subject: [PATCH 07/11] use less GPU ram on startup --- gpt2generator.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/gpt2generator.py b/gpt2generator.py index 629b841..84b69e6 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -136,7 +136,7 @@ class GPT2Generator: model_class, tokenizer_class = MODEL_CLASSES["gpt2"] self.tokenizer = tokenizer_class.from_pretrained(self.checkpoint_path) self.model = model_class.from_pretrained(self.checkpoint_path) - self.model.to(self.device).to(self.dtype) + self.model.to(self.dtype).to(self.device) self.model.eval() def sample_sequence(self, context_tokens=None, generate_num=None, temperature=None): @@ -193,7 +193,7 @@ class GPT2Generator: truncate_multiple_sequences(context_tokens, self.max_history_tokens) context_tokens = list(itertools.chain(*context_tokens)) - logger.debug("Text passing into model %s", self.tokenizer.decode(o, clean_up_tokenization_spaces=True, skip_special_tokens=True)) + logger.debug("Text passing into model %s", self.tokenizer.decode(context_tokens, clean_up_tokenization_spaces=True, skip_special_tokens=True)) generated = 0 for _ in range(self.samples // self.batch_size): From fd6c910bd98ddec0a5b7d1b73f872fb7fec55241 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 08:08:28 +0800 Subject: [PATCH 08/11] similarity fix and comment --- story/utils.py | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/story/utils.py b/story/utils.py index 5a59fb1..087d435 100644 --- a/story/utils.py +++ b/story/utils.py @@ -24,10 +24,11 @@ def console_print(text, width=75): #TODO: get rid if pyjarowinker dependency +# (AOP) You could use a simpler method, but this has been reported by RebootTech as a much more accurate way to compare strings. It also helps clean up the history. So it will hurt ability to check for looping def get_similarity(a, b): + if len(a)==0 or len(b)==0: return 1 return distance.get_jaro_distance( - a, b, winkler=True, scaling = 0.1 - ) + a, b, winkler=True, scaling = 0.1) def get_num_options(num): From 2cc4b968072aac6942a7bacecf797e5b5ee00820 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 08:08:48 +0800 Subject: [PATCH 09/11] fixes --- gpt2generator.py | 1 + play.py | 2 +- 2 files changed, 2 insertions(+), 1 deletion(-) diff --git a/gpt2generator.py b/gpt2generator.py index 84b69e6..3baf8fa 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -1,4 +1,5 @@ import os +import itertools import torch import torch.nn.functional as F diff --git a/play.py b/play.py index 15348d7..d9c389d 100644 --- a/play.py +++ b/play.py @@ -172,7 +172,7 @@ def play(generator): action_suggestion_lines = 1 for i in range(settings.getint('action-alternatives')): # FIXME it might be better to pass in a longer history - story_manager.story_context() # This should be within the loop as it has a random sampling element + action_prompt = story_manager.story_context() # This should be within the loop as it has a random sampling element action_prompt[-1] += '> ' logger.debug("action_prompt %s", action_prompt) suggested_action = ai_player.get_action(action_prompt) From 11d54e21cf4b324c32f73616bbb211b8cd78ca69 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 08:09:35 +0800 Subject: [PATCH 10/11] fixes --- gpt2generator.py | 3 ++- story/story_manager.py | 4 ++-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/gpt2generator.py b/gpt2generator.py index 3baf8fa..802da01 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -66,7 +66,7 @@ def sample_sequence( context = torch.tensor(context, dtype=torch.long, device=device) context = context.unsqueeze(0).repeat(num_samples, 1) generated = context - USE_PAST = False + USE_PAST = True next_token = context outputs = None with torch.no_grad(): @@ -194,6 +194,7 @@ class GPT2Generator: truncate_multiple_sequences(context_tokens, self.max_history_tokens) context_tokens = list(itertools.chain(*context_tokens)) + if os.environ.get("DEBUG_GPT2", False): logger.debug("Text passing into model %s", self.tokenizer.decode(context_tokens, clean_up_tokenization_spaces=True, skip_special_tokens=True)) generated = 0 diff --git a/story/story_manager.py b/story/story_manager.py index cc52a69..73c9cdf 100644 --- a/story/story_manager.py +++ b/story/story_manager.py @@ -81,9 +81,9 @@ class Story: # When we have to much history we will take the last 10, and sample randomly from the rest # first take last mem_ind//2 all_inds = list(range(len(self.results))) + first = all_inds[:-mem_ind//2] last = all_inds[-mem_ind//2:] - first = all_inds[:mem_ind//2] - inds = sorted(last + random.sample(first, mem_ind//2)) + inds = sorted(random.sample(first, mem_ind//2)+last) else: inds = range(len(self.results)) logger.debug("Using history indices %s", inds) From 6b46f971912127d5b8150376026a0e28b17a68f6 Mon Sep 17 00:00:00 2001 From: wassname Date: Wed, 25 Dec 2019 08:25:37 +0800 Subject: [PATCH 11/11] fix --- gpt2generator.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/gpt2generator.py b/gpt2generator.py index 802da01..0cf3495 100644 --- a/gpt2generator.py +++ b/gpt2generator.py @@ -195,7 +195,7 @@ class GPT2Generator: context_tokens = list(itertools.chain(*context_tokens)) if os.environ.get("DEBUG_GPT2", False): - logger.debug("Text passing into model %s", self.tokenizer.decode(context_tokens, clean_up_tokenization_spaces=True, skip_special_tokens=True)) + logger.debug("Text passing into model %s", self.tokenizer.decode(context_tokens, clean_up_tokenization_spaces=True, skip_special_tokens=True)) generated = 0 for _ in range(self.samples // self.batch_size):