Whyx-PROmpTea / src /tag_warehouse.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
13.2 kB
import json
import os
import random
DATA_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), "data", "tag_pools")
CATEGORY_FILES = {
"quality": "quality.json",
"lighting": "lighting.json",
"composition": "composition.json",
"effects": "effects.json",
"atmosphere": "atmosphere.json",
"framing": "framing.json",
"style": "style.json",
"colors": "color_grading.json",
"background": "background.json",
"pose": "pose.json",
"expression": "expression.json",
"clothing": "clothing.json",
"special_fx": "special_fx.json",
"nsfw": "nsfw.json",
"year_meta": "year_meta.json",
"animal": "animal.json",
"furry": "furry.json",
"object": "object.json",
"food": "food.json",
"vehicle": "vehicle.json",
"weapon": "weapon.json",
"architecture": "architecture.json",
"demon": "demon.json",
"angelic": "angelic.json",
"hair": "hair.json",
"eyes": "eyes.json",
"body": "body.json",
"accessory": "accessory.json",
"season": "season.json",
}
DISPLAY_NAMES = {
"quality": ("Quality", "Качество"),
"lighting": ("Lighting", "Освещение"),
"composition": ("Composition", "Композиция"),
"effects": ("Visual Effects", "Эффекты"),
"atmosphere": ("Atmosphere", "Атмосфера"),
"framing": ("Framing", "Кадрирование"),
"style": ("Style", "Стиль"),
"colors": ("Color Grading", "Цвета"),
"background": ("Background", "Фон"),
"pose": ("Pose", "Поза"),
"expression": ("Expression", "Выражение"),
"clothing": ("Clothing", "Одежда"),
"special_fx": ("Special FX", "Спецэффекты"),
"nsfw": ("NSFW Tags", "NSFW теги"),
"year_meta": ("Year / Meta", "Год / Мета"),
"animal": ("Animals", "Животные"),
"furry": ("Furry", "Фурри"),
"object": ("Objects", "Предметы"),
"food": ("Food", "Еда"),
"vehicle": ("Vehicles", "Транспорт"),
"weapon": ("Weapons", "Оружие"),
"architecture": ("Architecture", "Архитектура"),
"demon": ("Demon", "Демон"),
"angelic": ("Angelic", "Ангельское"),
"hair": ("Hair", "Волосы"),
"body": ("Body", "Тело"),
"accessory": ("Accessory", "Аксессуары"),
"season": ("Season", "Сезон"),
}
STYLE_ICONS = {
"vibrant": "🎨",
"dark": "🌑",
"detailed": "🔍",
"cute": "🌸",
"elegant": "💎",
"dynamic": "⚡",
"moe": "💕",
"sexy": "🔥",
"realistic": "📷",
"painterly": "🖌️",
"minimalist": "⬜",
}
from src.safety_filter import TAG_RATING_ORDER as RATING_ORDER
from src.safety_filter import ui_rating_to_tag_ceiling
# Backwards-compatible alias: MAX_RATING_MAP is derived from the canonical map.
MAX_RATING_MAP = {
ui: ui_rating_to_tag_ceiling(ui) for ui in ("pg", "pg13", "pg16", "r", "r+")
}
class TagPool:
def __init__(self, tags: list[dict], category: str):
self.category = category
self._all_tags = tags
self._by_rating: dict[str, list[dict]] = {}
for tag in tags:
rating = tag.get("rating", "sfw")
if rating not in self._by_rating:
self._by_rating[rating] = []
self._by_rating[rating].append(tag)
self._all_tags_cache: dict[str, list[str]] = {}
def get_all_tag_dicts(self) -> list[dict]:
return list(self._all_tags)
def get_random(self, count: int = 1, max_rating: str = "sfw") -> list[str]:
max_level = RATING_ORDER.get(max_rating, 0)
available = [
t for r, tags in self._by_rating.items()
if RATING_ORDER.get(r, 0) <= max_level
for t in tags
]
if not available:
return []
# random.sample picks k unique items in O(k) — the old full-shuffle
# approach was O(n) for every call (and broke rating caching).
return [t["tag"] for t in random.sample(available, k=min(count, len(available)))]
def get_all_tags(self, max_rating: str = "sfw") -> list[str]:
cached = self._all_tags_cache.get(max_rating)
if cached is not None:
return list(cached)
max_level = RATING_ORDER.get(max_rating, 0)
result = [
t["tag"] for r, tags in self._by_rating.items()
if RATING_ORDER.get(r, 0) <= max_level
for t in tags
]
self._all_tags_cache[max_rating] = result
return list(result)
class ArtistPool:
def __init__(self, artists: list[dict], tandems: list[dict] | None = None):
self._artists = artists
self._tandems = tandems or []
# O(1) name lookup — 1000 artists make the previous linear scan wasteful.
self._by_name: dict[str, dict] = {a["tag"].lower(): a for a in artists}
def get_all(self) -> list[dict]:
return sorted(self._artists, key=lambda a: a.get("popularity", 0), reverse=True)
def get_by_style(self, style: str) -> list[dict]:
return [a for a in self._artists if a.get("style") == style]
def get_all_styles(self) -> list[str]:
styles = set()
for a in self._artists:
s = a.get("style", "")
if s:
styles.add(s)
return sorted(styles)
def find_artist(self, name: str) -> dict | None:
return self._by_name.get(name.lower().strip())
def suggest_artists(self, name: str, limit: int = 3) -> list[str]:
"""Close-match suggestions for a misspelled artist name ('did you mean…')."""
import difflib
names = list(self._by_name)
close = difflib.get_close_matches(name.lower().strip(), names, n=limit, cutoff=0.75)
return [self._by_name[c]["tag"] for c in close]
def get_signature_tags(self, artist_name: str) -> list[str]:
artist = self.find_artist(artist_name)
if artist:
return artist.get("signature_tags", [])
return []
def get_tandems(self) -> list[dict]:
return self._tandems
def get_tandems_for_artist(self, artist_name: str) -> list[dict]:
name_lower = artist_name.lower()
return [
t for t in self._tandems
if name_lower in [a.lower() for a in t.get("artists", [])]
]
def get_description(self, artist_name: str, lang: str = "en") -> str:
artist = self.find_artist(artist_name)
if not artist:
return ""
key = f"description_{lang}"
return artist.get(key, artist.get("description_en", ""))
def get_danbooru_url(self, artist_name: str) -> str:
artist = self.find_artist(artist_name)
if not artist:
return ""
return artist.get("danbooru_url", "")
def get_similar(self, artist_name: str, limit: int = 3) -> list[dict]:
target = self.find_artist(artist_name)
if not target:
return []
target_style = target.get("style", "")
target_tags = set(t.lower() for t in target.get("signature_tags", []))
scores = []
for a in self._artists:
if a["tag"].lower() == artist_name.lower():
continue
score = 0
if a.get("style") == target_style:
score += 2
a_tags = set(t.lower() for t in a.get("signature_tags", []))
score += len(target_tags & a_tags)
if score > 0:
scores.append((score, a))
scores.sort(key=lambda x: (-x[0], -x[1].get("popularity", 0)))
return [a for _, a in scores[:limit]]
class TagWarehouse:
def __init__(self):
self.pools: dict[str, TagPool] = {}
self.artist_pool: ArtistPool | None = None
self._conflict_map: dict[str, set[str]] = {}
self._tag_rating_index: dict[str, int] = {}
self._load_pools()
self._load_artists()
def _load_pools(self):
for category, filename in CATEGORY_FILES.items():
filepath = os.path.join(DATA_DIR, filename)
if not os.path.exists(filepath):
continue
try:
with open(filepath, "r", encoding="utf-8") as f:
tags = json.load(f)
except (json.JSONDecodeError, OSError):
continue
if not isinstance(tags, list):
continue
self.pools[category] = TagPool(tags, category)
for tag in tags:
if not isinstance(tag, dict) or "tag" not in tag:
continue
name = tag["tag"].lower().strip()
level = RATING_ORDER.get(tag.get("rating", "sfw"), 0)
prev = self._tag_rating_index.get(name)
if prev is None or level > prev:
self._tag_rating_index[name] = level
conflicts = tag.get("conflicts", [])
if name not in self._conflict_map:
self._conflict_map[name] = set()
if isinstance(conflicts, list):
for c in conflicts:
cl = c.lower().strip()
self._conflict_map[name].add(cl)
if cl not in self._conflict_map:
self._conflict_map[cl] = set()
self._conflict_map[cl].add(name)
def reload_pools(self):
self.pools.clear()
self._conflict_map.clear()
self._tag_rating_index.clear()
self._load_pools()
def tag_exceeds_rating(self, tag: str, max_rating: str) -> bool:
"""True if a tag is only available above the requested max rating (used to
gate artist signature tags so they don't leak NSFW content at safe ratings)."""
level = self._tag_rating_index.get(tag.lower().strip())
if level is None:
return False
return level > RATING_ORDER.get(max_rating, 0)
def reload_artists(self):
self.artist_pool = None
self._load_artists()
def _load_artists(self):
filepath = os.path.join(DATA_DIR, "artist.json")
if not os.path.exists(filepath):
return
try:
with open(filepath, "r", encoding="utf-8") as f:
data = json.load(f)
except (json.JSONDecodeError, OSError):
return
if isinstance(data, dict):
artists = data.get("artists", [])
tandems = data.get("tandems", [])
else:
artists = data
tandems = []
self.artist_pool = ArtistPool(artists, tandems)
def get_pool(self, category: str) -> TagPool | None:
return self.pools.get(category)
def get_all_artists(self) -> list[dict]:
if self.artist_pool:
return self.artist_pool.get_all()
return []
def get_artist_styles(self) -> list[str]:
if self.artist_pool:
return self.artist_pool.get_all_styles()
return []
def get_artists_by_style(self, style: str) -> list[dict]:
if self.artist_pool:
return self.artist_pool.get_by_style(style)
return []
def get_artist_signature_tags(self, artist_name: str) -> list[str]:
if self.artist_pool:
return self.artist_pool.get_signature_tags(artist_name)
return []
def get_all_tandems(self) -> list[dict]:
if self.artist_pool:
return self.artist_pool.get_tandems()
return []
def get_tandems_for_artist(self, artist_name: str) -> list[dict]:
if self.artist_pool:
return self.artist_pool.get_tandems_for_artist(artist_name)
return []
def get_artist_description(self, artist_name: str, lang: str = "en") -> str:
if self.artist_pool:
return self.artist_pool.get_description(artist_name, lang)
return ""
def get_artist_danbooru_url(self, artist_name: str) -> str:
if self.artist_pool:
return self.artist_pool.get_danbooru_url(artist_name)
return ""
def get_similar_artists(self, artist_name: str, limit: int = 3) -> list[dict]:
if self.artist_pool:
return self.artist_pool.get_similar(artist_name, limit)
return []
def get_category_display_name(self, category: str, lang: str = "en") -> str:
names = DISPLAY_NAMES.get(category, (category, category))
return names[1] if lang == "ru" else names[0]
def get_style_icon(self, style: str) -> str:
return STYLE_ICONS.get(style, "")
def has_conflict(self, tag: str, selected_tags: list[str]) -> bool:
conflicts = self._conflict_map.get(tag.lower().strip(), set())
for st in selected_tags:
if st.lower().strip() in conflicts:
return True
return False
def get_available_categories(self, rating: str = "pg") -> list[str]:
max_level = RATING_ORDER.get(MAX_RATING_MAP.get(rating, "sfw"), 0)
available = []
for cat, pool in self.pools.items():
if pool.get_all_tags(max_rating=MAX_RATING_MAP.get(rating, "sfw")):
available.append(cat)
return available