Whyx-PROmpTea / scripts /enrich_artist_signatures.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
5.3 kB
"""Enrich artist.json entries that have empty signature_tags.
For each such artist, fetch a page of their Danbooru posts and count
general-tag frequencies across that sample; keep the top N as signature_tags
and re-guess the style bucket from them.
Resumable: entries that already have signature_tags are skipped, so the
script can be interrupted and re-run safely.
"""
import json
import os
import sys
import time
import urllib.parse
import urllib.request
from collections import Counter
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from scripts.fetch_top_artists import (
ARTIST_PATH, _atomic_write_json, _guess_style, _load_json,
)
POSTS_URL = "https://danbooru.donmai.us/posts.json"
DEFAULT_POST_PAGES = 2 # 2 pages x 50 posts = up to 100 posts per artist
TOP_SIGNATURES = 5
# Filter ubiquitous noise tags out of signature extraction.
_NOISE = {
"1girl", "1boy", "solo", "highres", "absurdres", "commentary", "translation",
"signed", "artist name", "artist request", "commentary request", "translated",
"safe", "questionable", "nsfw", "explicit", "sensitive",
"blush", "smile", "open mouth", "looking at viewer", "long hair",
}
# Ratings we allow to contribute tags (g = general, s = sensitive).
_ALLOWED_POST_RATINGS = {"g", "s"}
def _load_nsfw_like_tags() -> set[str]:
"""All pool tags rated nsfw/explicit, in spaced lower-case form."""
pools_dir = os.path.join(
os.path.dirname(os.path.dirname(os.path.abspath(__file__))),
"data", "tag_pools",
)
bad: set[str] = set()
import glob as _glob
for path in _glob.glob(os.path.join(pools_dir, "*.json")):
try:
entries = _load_json(path)
except Exception:
continue
if isinstance(entries, dict): # artist.json shape
entries = entries.get("artists", [])
for e in entries:
if isinstance(e, dict) and e.get("rating") in ("nsfw", "explicit"):
t = e.get("tag", "").lower().strip()
if t:
bad.add(t)
return bad
_NSFW_TAGS = None
def _valid_signature(tag: str) -> bool:
global _NSFW_TAGS
if _NSFW_TAGS is None:
_NSFW_TAGS = _load_nsfw_like_tags()
if not tag or len(tag) < 3:
return False
if tag in _NOISE or tag in _NSFW_TAGS:
return False
# Reject symbol-heavy/emoticon tags (:3, :d, ...).
if any(ch in tag for ch in ":;<>^-_"):
return False
return tag.isprintable()
def _post_count_for(tag: str) -> int:
q = urllib.parse.urlencode({"tags": tag.replace(" ", "_"), "limit": 0})
url = f"{POSTS_URL}?{q}"
req = urllib.request.Request(url, headers={"User-Agent": "Whyx-PROptimizer/1.0"})
try:
with urllib.request.urlopen(req, timeout=20) as r:
return int(r.headers.get("X-Total-Count", "0") or 0)
except Exception:
return 0
def _top_tags_for(tag: str, pages: int = DEFAULT_POST_PAGES) -> list[str]:
counts: Counter = Counter()
for page in range(1, pages + 1):
q = urllib.parse.urlencode({
"tags": tag.replace(" ", "_"),
"limit": 50,
"page": page,
})
url = f"{POSTS_URL}?{q}"
req = urllib.request.Request(url, headers={"User-Agent": "Whyx-PROptimizer/1.0"})
try:
with urllib.request.urlopen(req, timeout=25) as r:
posts = json.loads(r.read().decode("utf-8"))
except Exception:
break
if not posts:
break
for post in posts:
if post.get("rating") not in _ALLOWED_POST_RATINGS:
continue
tags = (post.get("tag_string_general") or "").split()
for t in tags:
t = t.replace("_", " ").strip().lower()
if _valid_signature(t):
counts[t] += 1
time.sleep(0.35)
return [t for t, _ in counts.most_common(TOP_SIGNATURES)]
def main(limit: int | None = None, throttle: float = 0.6):
data = _load_json(ARTIST_PATH)
artists = data["artists"]
todo = [a for a in artists if not a.get("signature_tags")]
if limit:
todo = todo[:limit]
total = len(todo)
if total == 0:
print("nothing to enrich")
return
print(f"enriching {total} artists...")
changed = 0
for i, artist in enumerate(todo, 1):
tag = artist["tag"]
sig = _top_tags_for(tag)
if sig:
artist["signature_tags"] = sig
if artist.get("style", "detailed") == "detailed":
artist["style"] = _guess_style(sig)
changed += 1
if i % 25 == 0 or i == total:
print(f" {i}/{total} done, {changed} updated")
time.sleep(throttle)
_atomic_write_json(ARTIST_PATH, data)
still_empty = sum(1 for a in artists if not a.get("signature_tags"))
print(f"done: {changed} enriched, {still_empty} still empty")
if __name__ == "__main__":
import argparse
ap = argparse.ArgumentParser()
ap.add_argument("--limit", type=int, default=None, help="max artists to enrich this run")
ap.add_argument("--throttle", type=float, default=0.6, help="seconds between artists")
args = ap.parse_args()
main(args.limit, args.throttle)