Spaces:
Running
Running
| import json | |
| import os | |
| import re | |
| import random | |
| import datetime | |
| from copy import deepcopy | |
| from src.prompt_parser import ParsedPrompt | |
| from src.tag_warehouse import TagWarehouse, MAX_RATING_MAP | |
| from src.dedup_engine import smart_dedup | |
| _DATA_PATH = os.path.join(os.path.dirname(os.path.dirname(__file__)), "data", "rewrite_map.json") | |
| _REWRITE_MAP = None | |
| def _load_map() -> dict: | |
| global _REWRITE_MAP | |
| if _REWRITE_MAP is not None: | |
| return _REWRITE_MAP | |
| try: | |
| with open(_DATA_PATH, "r", encoding="utf-8") as f: | |
| _REWRITE_MAP = json.load(f) | |
| except (FileNotFoundError, json.JSONDecodeError): | |
| _REWRITE_MAP = {"tag_to_categories": {}} | |
| return _REWRITE_MAP | |
| _CATEGORY_PATTERN_CACHE: dict[str, re.Pattern] | None = None | |
| def _category_patterns() -> dict[str, re.Pattern]: | |
| global _CATEGORY_PATTERN_CACHE | |
| if _CATEGORY_PATTERN_CACHE is not None: | |
| return _CATEGORY_PATTERN_CACHE | |
| mapping = _load_map().get("tag_to_categories", {}) | |
| _CATEGORY_PATTERN_CACHE = { | |
| kw: re.compile(r"(?<![a-zA-Z0-9])" + re.escape(kw) + r"(?![a-zA-Z0-9])") | |
| for kw in mapping | |
| } | |
| return _CATEGORY_PATTERN_CACHE | |
| def get_tag_categories(tag: str) -> list[str]: | |
| mapping = _load_map().get("tag_to_categories", {}) | |
| tl = tag.lower().strip() | |
| results = [] | |
| for kw, pattern in _category_patterns().items(): | |
| if pattern.search(tl): | |
| results.extend(mapping[kw]) | |
| return list(set(results)) | |
| def extract_concepts(tags: list[str]) -> dict[str, list[str]]: | |
| concepts: dict[str, list[str]] = {} | |
| for tag in tags: | |
| tl = tag.lower().strip() | |
| categories = get_tag_categories(tag) | |
| for cat in categories: | |
| if cat not in concepts: | |
| concepts[cat] = [] | |
| if tl not in concepts[cat]: | |
| concepts[cat].append(tl) | |
| return concepts | |
| def rewrite_prompt( | |
| parsed: ParsedPrompt, | |
| selected_categories: list[str], | |
| warehouse: TagWarehouse, | |
| model: str = "anima", | |
| rating: str = "pg", | |
| creativity: str = "medium", | |
| weight_mode: str = "off", | |
| num_variations: int = 5, | |
| web_enrich: bool = False, | |
| seed: int | None = None, | |
| artist_style: str = "", | |
| selected_artists: list[str] | None = None, | |
| use_tandems: bool = False, | |
| selected_tandem: dict | None = None, | |
| ) -> list[str]: | |
| from src.variation_engine import CREATIVITY_SETTINGS, _pick_random_tandem, _pick_style_tandem, _smart_substitution, _min_diversity_index, _compute_tag_weights, _cooccurrence_bonus, _adjust_settings_by_prompt_length, _resolve_and_replace | |
| from src.semantic_coherence import (semantic_pick_tags, split_core_decorative, | |
| detect_intent, compute_theme_budget, _THEME_GROUPS) | |
| from src.tag_searcher import filter_subject_conflicts, has_human_subject | |
| from src.synonym_filter import has_synonym_conflict | |
| from src.synonym_data import _find_synonym_group | |
| seed_rng = random.Random(seed) if seed is not None else random.Random() | |
| settings = CREATIVITY_SETTINGS.get(creativity, CREATIVITY_SETTINGS["medium"]) | |
| settings = _adjust_settings_by_prompt_length(settings, parsed, len(warehouse.pools)) | |
| max_rating = MAX_RATING_MAP.get(rating, "sfw") | |
| results = [] | |
| _default_year = f"year {datetime.date.today().year}" | |
| _default_period = "newest" | |
| _skip_animal = has_human_subject(parsed.general_tags) and "animal" not in selected_categories | |
| base_general = parsed.general_tags or [] | |
| core_tags, decorative_tags = split_core_decorative(base_general) | |
| core_set = {t.lower().strip() for t in core_tags} | |
| intent = detect_intent(parsed) | |
| theme_budget = compute_theme_budget(intent, selected_categories, settings["tags_per_category"]) | |
| cat_to_theme: dict[str, str] = {} | |
| for theme, cats in _THEME_GROUPS.items(): | |
| for cat in cats: | |
| cat_to_theme[cat] = theme | |
| # Artist lookup map is loop-invariant — build once, reuse per variation. | |
| known_artists_map = {a["tag"].lower(): a["tag"] for a in warehouse.get_all_artists()} | |
| concepts = extract_concepts(parsed.general_tags) | |
| all_new_tags_per_variation: list[list[str]] = [] | |
| for var_idx in range(num_variations): | |
| per_seed = seed_rng.randint(0, 2**31 - 1) + var_idx * 7919 | |
| rng = random.Random(per_seed) | |
| variant = deepcopy(parsed) | |
| new_general = list(parsed.general_tags) if parsed.general_tags else [] | |
| resolved_categories = list(selected_categories) | |
| for concept_cat, concept_tags in concepts.items(): | |
| if concept_cat not in resolved_categories: | |
| resolved_categories.append(concept_cat) | |
| resolved_categories = list(dict.fromkeys(resolved_categories)) | |
| new_added = [] | |
| # Tandem injection (style-aware) | |
| if use_tandems and not selected_tandem and not selected_artists: | |
| tandem_tags = _pick_style_tandem(rng, warehouse, artist_style) | |
| existing_lower = {t.lower() for t in new_general} | |
| for t in tandem_tags: | |
| if t.lower() not in existing_lower: | |
| new_general.append(t) | |
| new_added.append(t) | |
| existing_lower.add(t.lower()) | |
| if use_tandems and selected_tandem and not selected_artists: | |
| tandem_artists = selected_tandem.get("artists", []) | |
| existing_lower = {t.lower() for t in new_general} | |
| for aname in tandem_artists: | |
| if aname.lower() not in existing_lower: | |
| new_general.append(aname) | |
| new_added.append(aname) | |
| existing_lower.add(aname.lower()) | |
| sig = warehouse.get_artist_signature_tags(aname) | |
| for st in sig: | |
| if warehouse.tag_exceeds_rating(st, max_rating): | |
| continue | |
| if st.lower() not in existing_lower: | |
| new_general.append(st) | |
| new_added.append(st) | |
| existing_lower.add(st.lower()) | |
| # Artist style injection | |
| if artist_style and not selected_artists and not use_tandems and not selected_tandem: | |
| style_artists = warehouse.get_artists_by_style(artist_style) | |
| if style_artists: | |
| pool = rng.sample(style_artists, min(3, len(style_artists))) | |
| for a in pool: | |
| new_general.append(a["tag"]) | |
| new_added.append(a["tag"]) | |
| # Selected artists injection | |
| if selected_artists: | |
| for aname in selected_artists: | |
| if aname not in new_general: | |
| new_general.append(aname) | |
| new_added.append(aname) | |
| sig_tags = warehouse.get_artist_signature_tags(aname) | |
| existing_lower = {t.lower() for t in new_general} | |
| for st in sig_tags: | |
| if warehouse.tag_exceeds_rating(st, max_rating): | |
| continue | |
| if st.lower() not in existing_lower: | |
| new_general.append(st) | |
| new_added.append(st) | |
| existing_lower.add(st.lower()) | |
| # Replacement rate: remove some user tags proportionally (quality tags weighted lower; core protected) | |
| if settings["replacement_rate"] > 0 and base_general: | |
| user_tags_in_new = [t for t in new_general if t in base_general and t.lower().strip() not in core_set] | |
| if user_tags_in_new: | |
| n_replace = max(1, int(len(user_tags_in_new) * settings["replacement_rate"])) | |
| quality_keywords = {"score", "masterpiece", "quality", "aesthetic", "detailed"} | |
| weights = [] | |
| for t in user_tags_in_new: | |
| tl = t.lower().strip() | |
| is_quality = any(kw in tl for kw in quality_keywords) | |
| weights.append(0.2 if is_quality else 1.0) | |
| to_remove = rng.choices(user_tags_in_new, weights=weights, k=min(n_replace, len(user_tags_in_new))) | |
| to_remove = list(dict.fromkeys(to_remove)) | |
| for t in to_remove: | |
| if t in new_general: | |
| new_general.remove(t) | |
| var_theme_usage: dict[str, int] = {} | |
| for cat in resolved_categories: | |
| if cat == "animal" and _skip_animal: | |
| continue | |
| pool = warehouse.get_pool(cat) | |
| if pool is None: | |
| continue | |
| min_t, max_t = settings["tags_per_category"] | |
| theme = cat_to_theme.get(cat, "misc") | |
| theme_max = theme_budget.get(theme, max_t * 2) | |
| used_this_theme = var_theme_usage.get(theme, 0) | |
| local_max = max(min_t, min(max_t, theme_max - used_this_theme)) | |
| pick_count = rng.randint(min_t, local_max) if local_max >= min_t else min_t | |
| candidates = pool.get_all_tags(max_rating=max_rating) | |
| if not candidates: | |
| continue | |
| rng.shuffle(candidates) | |
| preferred = [] | |
| other = [] | |
| concept_keywords = concepts.get(cat, []) | |
| seen_lower = {t.lower().strip() for t in new_general} | |
| for c in candidates: | |
| cl = c.lower().strip() | |
| if cl in seen_lower: | |
| continue | |
| matched = False | |
| for kw in concept_keywords: | |
| if kw in cl or cl in kw: | |
| matched = True | |
| break | |
| if matched: | |
| preferred.append(c) | |
| else: | |
| other.append(c) | |
| picked = preferred + other | |
| used_set = {t.lower().strip() for t in new_general} | |
| picked_tags = semantic_pick_tags(picked, pick_count, parsed, rng, used_set, new_general, intent=intent) | |
| for tag in picked_tags: | |
| if has_synonym_conflict(tag, new_general): | |
| continue | |
| new_general, was_added = _resolve_and_replace(tag, new_general, warehouse) | |
| if was_added: | |
| used_set = {t.lower().strip() for t in new_general} | |
| new_added.append(tag) | |
| var_theme_usage[theme] = var_theme_usage.get(theme, 0) + 1 | |
| # Smart substitution pass: context-aware replacement with cross-category fallback | |
| if settings["substitution_chance"] > 0 and new_general: | |
| new_general = _smart_substitution(new_general, settings["substitution_chance"], rng, | |
| parsed, warehouse, resolved_categories, core_set) | |
| # Web enrichment | |
| if web_enrich: | |
| from src.tag_searcher import enrich_prompt_tags | |
| enrich = enrich_prompt_tags(variant, warehouse, max_tags=5, | |
| user_tags=parsed.general_tags) | |
| existing_lower = {t.lower().strip() for t in new_general} | |
| for t in enrich: | |
| if t.lower().strip() not in existing_lower: | |
| new_general.append(t) | |
| existing_lower.add(t.lower().strip()) | |
| # Diversity enforcement: re-inject if too similar to prior variations | |
| for prev_tags in all_new_tags_per_variation: | |
| diversity = _min_diversity_index(new_added, prev_tags) | |
| if diversity < settings["diversity_threshold"] and num_variations > 1: | |
| extra_seed = rng.randint(0, 2**31 - 1) | |
| re_rng = random.Random(extra_seed) | |
| for cat in resolved_categories: | |
| pool = warehouse.get_pool(cat) | |
| if pool is None: | |
| continue | |
| extra_candidates = pool.get_all_tags(max_rating=max_rating) | |
| re_rng.shuffle(extra_candidates) | |
| added_count = 0 | |
| for tag in extra_candidates: | |
| if added_count >= 2: | |
| break | |
| new_general, was_added = _resolve_and_replace(tag, new_general, warehouse) | |
| if was_added: | |
| added_count += 1 | |
| break | |
| if model == "anima" and new_general: | |
| if "year_meta" in selected_categories or parsed.year_tag: | |
| variant.year_tag = parsed.year_tag or _default_year | |
| variant.period_tag = parsed.period_tag or _default_period | |
| else: | |
| variant.year_tag = "" | |
| variant.period_tag = "" | |
| elif model == "illustrious": | |
| if variant.year_tag: | |
| variant.year_tag = "" | |
| if variant.period_tag: | |
| variant.period_tag = "" | |
| all_new_tags_per_variation.append(new_added) | |
| variant.general_tags = new_general | |
| variant.general_tags = filter_subject_conflicts(variant.general_tags, parsed.general_tags) | |
| variant.general_tags = smart_dedup(variant.general_tags, model=model) | |
| if settings["shuffle_general"]: | |
| rng.shuffle(variant.general_tags) | |
| # Extract artist tags from general_tags (loop-invariant lookup map). | |
| skip = {(variant.subject or "").lower(), (variant.character or "").lower(), (variant.series or "").lower()} | |
| skip.discard("") | |
| remaining = [] | |
| for tag in variant.general_tags: | |
| tl = tag.lower() | |
| if tl in skip: | |
| continue | |
| matched = known_artists_map.get(tl) | |
| if matched and matched not in variant.artists: | |
| variant.artists.append(matched) | |
| else: | |
| remaining.append(tag) | |
| variant.general_tags = remaining | |
| variant.quality_tags = list(parsed.quality_tags) if parsed.quality_tags else [] | |
| if model == "anima": | |
| variant.meta_tags = list(parsed.meta_tags) if parsed.meta_tags else [] | |
| from src.model_formatter import format_prompt | |
| tag_weights = _compute_tag_weights(warehouse, selected_artists or []) | |
| result = format_prompt( | |
| variant, model=model, rating=rating, | |
| quality_enabled=("quality" in selected_categories), | |
| weight_mode=weight_mode, | |
| tag_weights=tag_weights, | |
| ) | |
| results.append(result) | |
| return results | |
| def reload_map(): | |
| global _REWRITE_MAP, _CATEGORY_PATTERN_CACHE | |
| _REWRITE_MAP = None | |
| _CATEGORY_PATTERN_CACHE = None | |
| _load_map() | |