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"(? 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()