From 2c6d01290d1c72f506649284f557fef8ef19cdc7 Mon Sep 17 00:00:00 2001 From: nickwalton Date: Thu, 31 Oct 2019 11:08:21 -0600 Subject: [PATCH] update --- generator/ctrl/ctrl_generator.py | 13 +++++++------ story/utils.py | 11 ++++------- 2 files changed, 11 insertions(+), 13 deletions(-) diff --git a/generator/ctrl/ctrl_generator.py b/generator/ctrl/ctrl_generator.py index 91be38a..8b89832 100644 --- a/generator/ctrl/ctrl_generator.py +++ b/generator/ctrl/ctrl_generator.py @@ -20,7 +20,7 @@ def loss(labels, logits): class CTRLGenerator(): - def __init__(self, control_code="Apocalypse ", generate_num=28, temperature=0.5, topk=40, nucleus_prob=0): + def __init__(self, control_code="Apocalypse ", generate_num=28, temperature=0.3, topk=40, nucleus_prob=0): self.generate_num=generate_num model_dir = "generator/ctrl/training_utils/seqlen256_v1.ckpt/" @@ -135,7 +135,7 @@ class CTRLGenerator(): self.temperature=temperature self.nucleusprob = nucleus_prob - self.penalty = 1.1 + self.penalty = 1.2 self.topk=topk def configure_verb_probs(self, probabilities, options): @@ -212,14 +212,15 @@ class CTRLGenerator(): penalized_so_far = set() for _ in range(token + 1): generated_token = tokens_generated[0][_] - penalized_so_far.add(generated_token) - prompt_logits[_token][generated_token] /= self.penalty + if generated_token not in penalized_so_far: + penalized_so_far.add(generated_token) + prompt_logits[_token][generated_token] /= self.penalty # disallow some tokens forbidden_tokens = ['', 'Sco@@', "&@@", "1]@@", "2]@@", "3]@@", "4]@@", "https://www.@@", "[@@", ":@@", "Edit", "&@@", "2:","1:", ":", "Edit@@", "EDI@@", "EDIT@@", "edit", "TL@@", "tl@@", ";@@", '**', "http://@@", "Redd@@", "UP@@", "mom", "Up@@", "Me:", "Update", "mom@@", "Part", - "http://www.@@", "edit@@", "*@@", "Writing", "Text@@", "\\@@", "
@@", "@@", " last_period: - text = text[0:last_exclamation+1] - elif last_period > 0: - text = text[0:last_period+1] + last_punc = max(text.rfind('.'), text.rfind("!"), text.rfind("?")) + + if last_punc > 0: + text = text[0:last_punc+1] return cut_trailing_quotes(text)