mirror of
https://github.com/wassname/Clover-Edition.git
synced 2026-09-09 11:13:26 +08:00
update
This commit is contained in:
+1
-1
@@ -40,7 +40,7 @@ def play_unconstrained():
|
||||
def play_constrained():
|
||||
#generator = WebGenerator(CRED_FILE)
|
||||
generator = CTRLGenerator()
|
||||
story_start = "classic"
|
||||
story_start = "haunted"
|
||||
verbs_key = "anything"
|
||||
prompt = get_story_start(story_start)
|
||||
story_manager = ConstrainedStoryManager(generator, prompt, action_verbs_key=verbs_key)
|
||||
|
||||
@@ -137,11 +137,13 @@ class CTRLGenerator():
|
||||
self.topk = 0
|
||||
|
||||
|
||||
def generate(self, prompt):
|
||||
def generate(self, prompt, first_verb_whitelist=True):
|
||||
|
||||
if prompt[-1] != " ":
|
||||
prompt = prompt + " "
|
||||
|
||||
first_token = True
|
||||
|
||||
prompt = second_to_first_person(prompt)
|
||||
|
||||
prompt = self.control_code + prompt
|
||||
@@ -198,13 +200,15 @@ class CTRLGenerator():
|
||||
prompt_logits[_token][generated_token] /= self.penalty
|
||||
|
||||
# disallow some tokens
|
||||
prompt_logits[_token][self.word2idx['<unk>']] = -1e8
|
||||
prompt_logits[_token][self.word2idx['\n']] = -1e8
|
||||
forbidden_tokens = ['<unk>', 'Sco@@']
|
||||
for token in forbidden_tokens:
|
||||
prompt_logits[_token][self.word2idx[token]] = -1e8
|
||||
|
||||
# sometimes, when generating from reddit,
|
||||
# it tries to generate the Score (reddit Karma) immediately after generating the Title:
|
||||
# to disallow this, we can just prevent it from generating Score
|
||||
prompt_logits[_token][self.word2idx['Sco@@']] = -1e8
|
||||
# Make sure only a possible verb is chosen.
|
||||
if first_token:
|
||||
for word in get_possible_verbs():
|
||||
prompt_logits[_token][self.word2idx[word]] += 0.1
|
||||
first_token = False
|
||||
|
||||
# compute probabilities from logits
|
||||
prompt_probs = np.exp(prompt_logits[_token])
|
||||
@@ -228,7 +232,7 @@ class CTRLGenerator():
|
||||
# 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 = []
|
||||
tokens_to_disallow = ["\n\n"]
|
||||
for _ in range(len(pruned_list)):
|
||||
if 'http' in self.idx2word[pruned_list[_]]:
|
||||
tokens_to_disallow.append(_)
|
||||
|
||||
@@ -109,7 +109,7 @@ class ConstrainedStoryManager(StoryManager):
|
||||
return [self.generate_action_result(self.story_context(), phrase) for phrase in self.action_phrases]
|
||||
|
||||
def generate_action_result(self, prompt, phrase):
|
||||
action = phrase + self.generator.generate(prompt + phrase)
|
||||
action = phrase + " " + self.generator.generate(prompt + phrase)
|
||||
action_result = cut_trailing_sentence(action)
|
||||
|
||||
action, result = split_first_sentence(action_result)
|
||||
|
||||
@@ -133,6 +133,12 @@ def second_to_first_person(text):
|
||||
|
||||
return capitalize_first_letters(text)
|
||||
|
||||
|
||||
possible_verbs = ["go", "run", "open", "look", "walk", "make", "try", "say", "tell", "attack", "use", "turn", "fight", "scream", "yell"]
|
||||
|
||||
def get_possible_verbs():
|
||||
return possible_verbs
|
||||
|
||||
if __name__ == '__main__':
|
||||
f = open("test.txt", "r")
|
||||
test_text = f.read()
|
||||
|
||||
Reference in New Issue
Block a user