Whyx-PROmpTea / src /prompt_rewriter.py
ArtShumov's picture
feat(prod): rewrite pipeline + NoobAI + ensemble tagger (3xWD14+DeepDanbooru-ready) + 1000 artists + negative templates + history ext + scoring config
e6404d0
Raw
History Blame
14.6 kB
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()