preprocess algorithm is updated

This commit is contained in:
alpertunga-bile
2023-08-11 06:33:29 +00:00
committed by GitHub
parent 1612ad32b2
commit 773a2131b4
+7 -17
View File
@@ -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}]], )
return ([[cond, {"pooled_output": pooled}]], )