From 773a2131b46d7c9cc585110a9a0c37bccbb9071f Mon Sep 17 00:00:00 2001 From: alpertunga-bile <76731692+alpertunga-bile@users.noreply.github.com> Date: Fri, 11 Aug 2023 06:33:29 +0000 Subject: [PATCH] preprocess algorithm is updated --- prompt_generator.py | 24 +++++++----------------- 1 file changed, 7 insertions(+), 17 deletions(-) diff --git a/prompt_generator.py b/prompt_generator.py index 4f8f8d5..4490a6d 100644 --- a/prompt_generator.py +++ b/prompt_generator.py @@ -81,29 +81,19 @@ class PromptGenerator: return [x for x in sequence if not (x in seen or seen.add(x))] def RemoveDuplicates(self, line : str) -> list[str]: - char_blacklist = set("():.1234567890") + from string import punctuation + + char_blacklist = set(f"{punctuation}0123456789") # remove exact prompts prompts = self.GetUniqueList(line.split(",")) pure_prompts = [] - """ - def GetCanAdd(given_substring : str, original_string : str) -> bool: - if len(given_substring) > len(original_string): - if original_string in given_substring: - return False - else: - if given_substring in original_string: - return False - - return True - """ - # remove exact keyword for prompt in prompts: can_add = True # extract the keyword - keyword = "".join(c for c in prompt if c not in char_blacklist) + keyword = "".join(c for c in prompt if c not in char_blacklist).lstrip() if keyword == "": continue @@ -113,8 +103,8 @@ class PromptGenerator: can_add = False break - extracted_pure_prompt = "".join(c for c in pure_prompt if c not in char_blacklist) - if keyword in extracted_pure_prompt or keyword == extracted_pure_prompt: + extracted_pure_prompt = "".join(c for c in pure_prompt if c not in char_blacklist).lstrip() + if keyword == extracted_pure_prompt: can_add = False break @@ -222,4 +212,4 @@ class PromptGenerator: tokens = clip.tokenize(generated_text) cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True) - return ([[cond, {"pooled_output": pooled}]], ) \ No newline at end of file + return ([[cond, {"pooled_output": pooled}]], )