From 275fac8931c46cd303bde50bc590743862c604d0 Mon Sep 17 00:00:00 2001 From: Haoming Date: Mon, 28 Jul 2025 15:57:56 +0800 Subject: [PATCH] yeet --- modules/styles.py | 52 +++++++++++++++++------------------------------ 1 file changed, 19 insertions(+), 33 deletions(-) diff --git a/modules/styles.py b/modules/styles.py index 350ea830..d061209e 100644 --- a/modules/styles.py +++ b/modules/styles.py @@ -31,7 +31,8 @@ def apply_styles_to_prompt(prompt, styles): def extract_style_text_from_prompt(style_text, prompt): - """This function extracts the text from a given prompt based on a provided style text. It checks if the style text contains the placeholder {prompt} or if it appears at the end of the prompt. If a match is found, it returns True along with the extracted text. Otherwise, it returns False and the original prompt. + """ + This function extracts the text from a given prompt based on a provided style text. It checks if the style text contains the placeholder {prompt} or if it appears at the end of the prompt. If a match is found, it returns True along with the extracted text. Otherwise, it returns False and the original prompt. extract_style_text_from_prompt("masterpiece", "1girl, art by greg, masterpiece") outputs (True, "1girl, art by greg") extract_style_text_from_prompt("masterpiece, {prompt}", "masterpiece, 1girl, art by greg") outputs (True, "1girl, art by greg") @@ -44,13 +45,13 @@ def extract_style_text_from_prompt(style_text, prompt): if "{prompt}" in stripped_style_text: left, _, right = stripped_style_text.partition("{prompt}") if stripped_prompt.startswith(left) and stripped_prompt.endswith(right): - prompt = stripped_prompt[len(left):len(stripped_prompt)-len(right)] + prompt = stripped_prompt[len(left) : len(stripped_prompt) - len(right)] return True, prompt else: if stripped_prompt.endswith(stripped_style_text): - prompt = stripped_prompt[:len(stripped_prompt)-len(stripped_style_text)] + prompt = stripped_prompt[: len(stripped_prompt) - len(stripped_style_text)] - if prompt.endswith(', '): + if prompt.endswith(", "): prompt = prompt[:-2] return True, prompt @@ -86,9 +87,9 @@ class StyleDatabase: self.all_styles_files: list[Path] = [] folder, file = os.path.split(self.paths[0]) - if '*' in file or '?' in file: + if "*" in file or "?" in file: # if the first path is a wildcard pattern, find the first match else use "folder/styles.csv" as the default path - self.default_path = next(Path(folder).glob(file), Path(os.path.join(folder, 'styles.csv'))) + self.default_path = next(Path(folder).glob(file), Path(os.path.join(folder, "styles.csv"))) self.paths.insert(0, self.default_path) else: self.default_path = Path(self.paths[0]) @@ -108,25 +109,20 @@ class StyleDatabase: all_styles_files = [] for pattern in self.paths: folder, file = os.path.split(pattern) - if '*' in file or '?' in file: + if "*" in file or "?" in file: found_files = Path(folder).glob(file) [all_styles_files.append(file) for file in found_files] else: - # if os.path.exists(pattern): - all_styles_files.append(Path(pattern)) + if os.path.exists(pattern): + all_styles_files.append(Path(pattern)) # Remove any duplicate entries seen = set() self.all_styles_files = [s for s in all_styles_files if not (s in seen or seen.add(s))] for styles_file in self.all_styles_files: - if len(all_styles_files) > 1: - # add divider when more than styles file - # '---------------- STYLES ----------------' - divider = f' {styles_file.stem.upper()} '.center(40, '-') - self.styles[divider] = PromptStyle(f"{divider}", None, None, "do_not_save") - if styles_file.is_file(): - self.load_from_csv(styles_file) + assert styles_file.is_file() + self.load_from_csv(styles_file) def load_from_csv(self, path: str | Path): try: @@ -140,11 +136,9 @@ class StyleDatabase: prompt = row["prompt"] if "prompt" in row else row["text"] negative_prompt = row.get("negative_prompt", "") # Add style to database - self.styles[row["name"]] = PromptStyle( - row["name"], prompt, negative_prompt, str(path) - ) + self.styles[row["name"]] = PromptStyle(row["name"], prompt, negative_prompt, str(path)) except Exception: - errors.report(f'Error loading styles from {path}: ', exc_info=True) + errors.report(f"Error loading styles from {path}: ", exc_info=True) def get_style_paths(self) -> set: """Returns a set of all distinct paths of files that styles are loaded from.""" @@ -172,14 +166,10 @@ class StyleDatabase: return [self.styles.get(x, self.no_style).negative_prompt for x in styles] def apply_styles_to_prompt(self, prompt, styles): - return apply_styles_to_prompt( - prompt, [self.styles.get(x, self.no_style).prompt for x in styles] - ) + return apply_styles_to_prompt(prompt, [self.styles.get(x, self.no_style).prompt for x in styles]) def apply_negative_styles_to_prompt(self, prompt, styles): - return apply_styles_to_prompt( - prompt, [self.styles.get(x, self.no_style).negative_prompt for x in styles] - ) + return apply_styles_to_prompt(prompt, [self.styles.get(x, self.no_style).negative_prompt for x in styles]) def save_styles(self, path: str = None) -> None: # The path argument is deprecated, but kept for backwards compatibility @@ -202,9 +192,7 @@ class StyleDatabase: if style.name.lower().strip("# ") in csv_names: continue # Write style fields, ignoring the path field - writer.writerow( - {k: v for k, v in style._asdict().items() if k != "path"} - ) + writer.writerow({k: v for k, v in style._asdict().items() if k != "path"}) def extract_styles_from_prompt(self, positive, negative): extracted = [] @@ -218,9 +206,7 @@ class StyleDatabase: found_style = None for style in applicable_styles: - is_match, new_positive, new_negative = extract_original_prompts( - style, positive, negative - ) + is_match, new_positive, new_negative = extract_original_prompts(style, positive, negative) if is_match: found_style = style positive = new_positive @@ -232,4 +218,4 @@ class StyleDatabase: if not found_style: break - return list(reversed(extracted)), positive, negative \ No newline at end of file + return list(reversed(extracted)), positive, negative