From 8ab3fa8ddb535205891d85debf8e8ad0e2cbbf5c Mon Sep 17 00:00:00 2001 From: Nick Date: Thu, 26 Sep 2019 10:27:09 -0600 Subject: [PATCH] refactored --- generator/ctrl/ctrl_generator.py | 35 +------------------------------- 1 file changed, 1 insertion(+), 34 deletions(-) diff --git a/generator/ctrl/ctrl_generator.py b/generator/ctrl/ctrl_generator.py index 7917f73..a4da448 100644 --- a/generator/ctrl/ctrl_generator.py +++ b/generator/ctrl/ctrl_generator.py @@ -210,19 +210,13 @@ class CTRLGenerator(): penalized_so_far = set() for _ in range(token + 1): generated_token = tokens_generated[0][_] - # don't penalize newlines - # you could also choose not to penalize frequent words - # (which incidentally are sorted in the vocab file) - # but I don't do that - # if it prints too many new lines instead of continuing generating text, - # you might want to comment this out if generated_token in penalized_so_far: continue penalized_so_far.add(generated_token) prompt_logits[_token][generated_token] /= self.penalty # disallow some tokens - forbidden_tokens = ['', 'Sco@@'] + forbidden_tokens = ['', 'Sco@@', "UPDATE:", "EDIT", "UPDATE", "EDIT:", "[Part", "&", "*EDIT"] if num_new_lines > self.max_new_lines: forbidden_tokens.append("\n") @@ -249,27 +243,10 @@ class CTRLGenerator(): minimum_topk = 1 nucleus = max(np.where(np.cumsum(np.sort(prompt_probs)[::-1]) > self.nucleusprob)[0][0], minimum_topk) elif self.topk > 0: - # we are over-loading notation here - # if you choose to specify a topk instead of a nucleus, - # we will hardcode the nucleus to be just that nucleus = self.topk else: - # if you specify neither nucleus or topk, - # then we will use the whole list nucleus = len(pruned_list) - # if you want to disallow more complex tokens, you can do so here - # for instance, if you want to disallow anything with the phrase `http`, - # you can delete theme from the pruned_list - # you can comment this out, I'm keeping it in for demonstration purpose - tokens_to_disallow = [] - complex_to_dissalow = [] - for i in range(len(pruned_list)): - for complex in complex_to_dissalow: - if complex in self.idx2word[pruned_list[i]]: - tokens_to_disallow.append(i) - pruned_list = np.delete(pruned_list, tokens_to_disallow) - # if temperature is 0 # just pick the first (most probable) token if self.temperature == 0: @@ -281,19 +258,9 @@ class CTRLGenerator(): tf.random.categorical(np.expand_dims(prompt_logits[_token][pruned_list], 0), num_samples=1).numpy()) idx = pruned_list[chosen_idx] - # if you want to do some debugging, - # like which one was chosen, - # what the top25 were, - # here is your opportunity. - # print('chosen:', repr(self.idx2word[idx])) - # print('top25 alternatives:', pruned_list[:25]) - # assign the token for generation tokens_generated[0][token + 1] = idx - # clear screen if you want to - # os.system("clear") - return idx def generate(self, prompt, options=None):