Whyx-PROmpTea / src /ensemble_tagger.py
ArtShumov's picture
feat(tagger): pose extraction + Qwen3-VL optional NL captioner + WD14x3 ensemble weights recalibration
2c574fa
Raw
History Blame
12.3 kB
"""Ensemble image tagger: combines multiple WD14-family taggers (ViT / SwinV2 /
EVA02) and optionally DeepDanbooru for architecture-diverse voting.
Strategy: each model emits per-tag probabilities; the ensemble merges them with
per-model weights (WD14s share fallback "0.4 / 0.6" default; DeepDanbooru, if
present, overrides to "0.25 / 0.75" to lower its heavier corpus bias)."""
from __future__ import annotations
import hashlib
import json
import os
from typing import Optional
import numpy as np
from PIL import Image, ImageOps
from src.image_tagger import (
ImageTagger, _TAGGER_DEPS_OK, _ensure_rgb, _pad_square, _to_pil_image,
_tags_to_caption,
)
# ---------------------------------------------------------------------------
# Helpers shared with app.py (regex-based parsing of results dicts)
# ---------------------------------------------------------------------------
_IMG_TAGS_CACHE: dict[str, dict] = {}
def _normalize_tag(tag: str) -> str:
r"""WD14-style tags escape parens like \(symbol\); model output plain ones."""
return tag.replace("\\(", "(").replace("\\)", ")").strip()
def _esc_for_output(tag: str) -> str:
return tag.replace("(", "\\(").replace(")", "\\)")
# Order matters: the first available model becomes the primary and supplies
# the rating distribution. The rest contribute general/character votes.
try:
import onnxruntime as ort
_ORT_AVAILABLE = True
except Exception:
ort = None
_ORT_AVAILABLE = False
# ---------------------------------------------------------------------------
# Optional DeepDanbooru ONNX runner
# ---------------------------------------------------------------------------
_DD_REPO = "KichangKim/DeepDanbooru"
_DD_ONNX_FILES = ["deepdanbooru.onnx", "model-resnet_custom_v3.onnx"]
class _DeepDanbooruONNX:
"""Minimal DeepDanbooru v3 runner via ONNXRuntime, if the ONNX weights are
present in the HF cache. Tag set is loaded from DeepDanbooru's tags.txt.
DeepDanbooru has a ~9170-tag vocabulary that differs from WD14's 13k;
ratings are NOT predicted by this model, so ensemble falls back to the
WD14 rating distribution.
"""
def __init__(self, hf_repo: str | None = None, onnx_filename: str = "deepdanbooru.onnx"):
self._session: Optional["ort.InferenceSession"] = None
self._tags: list[str] = []
self._loaded = False
self._repo = hf_repo or _DD_REPO
self._onnx_name = onnx_filename
def ensure_loaded(self) -> bool:
if self._loaded:
return True
if not _ORT_AVAILABLE:
return False
try:
model_path = self._resolve_weight_path()
if not model_path:
return False
providers = ["CPUExecutionProvider"]
if "CUDAExecutionProvider" in ort.get_available_providers():
providers = ["CUDAExecutionProvider", "CPUExecutionProvider"]
self._session = ort.InferenceSession(model_path, providers=providers)
self._tags = self._load_tags()
self._loaded = bool(self._tags)
except Exception:
return False
return self._loaded
def _resolve_weight_path(self) -> str | None:
from huggingface_hub import hf_hub_download
try:
return hf_hub_download(repo_id=self._repo, filename=self._onnx_name)
except Exception:
return None
def _load_tags(self) -> list[str]:
from huggingface_hub import hf_hub_download
try:
p = hf_hub_download(repo_id=self._repo, filename="tags.txt")
with open(p, encoding="utf-8") as fh:
return [line.strip() for line in fh if line.strip()]
except Exception:
return []
def predict(self, pil_img: Image.Image) -> dict[str, float]:
if not self.ensure_loaded():
return {}
img = pil_img.convert("RGB").resize((512, 512), Image.LANCZOS)
arr = np.asarray(img, dtype=np.float32) / 255.0
arr = np.transpose(arr, (2, 0, 1))[None, ...] # NCHW
input_name = self._session.get_inputs()[0].name
out = self._session.run(None, {input_name: arr})[0].reshape(-1)
# DeepDanbooru outputs logits — apply sigmoid for 0..1 comparability.
probs = 1.0 / (1.0 + np.exp(-out))
idx = np.argsort(-probs)[: len(self._tags)]
tags = {}
for i in idx:
name = self._tags[i]
score = float(probs[i])
if score < 0.30:
break
tags[name] = score
return tags
# ---------------------------------------------------------------------------
# Ensemble coordinator
# ---------------------------------------------------------------------------
_ENSEMBLE_DEFAULT_MODELS = (
("SmilingWolf/wd-eva02-large-tagger-v3", 0.50, "eva02"),
("SmilingWolf/wd-swinv2-tagger-v3", 0.30, "swinv2"),
("SmilingWolf/wd-vit-tagger-v3", 0.20, "vit"),
)
class EnsembleTagger:
"""Aggregates multiple WD14-family taggers and (optionally) DeepDanbooru."""
def __init__(self):
self._wd_taggers: list[tuple[ImageTagger, float, str]] = []
for repo, weight, name in _ENSEMBLE_DEFAULT_MODELS:
self._wd_taggers.append((ImageTagger(repo_id=repo), weight, name))
self._dd = _DeepDanbooruONNX()
self._cache: dict[str, dict] = {}
def _wd_by_arch(self, arch: str) -> ImageTagger | None:
for tagger, _, name in self._wd_taggers:
if name == arch:
return tagger
return None
def _image_key(self, pil_img: Image.Image) -> str:
buf = pil_img.tobytes()[: 8192] + str(pil_img.size).encode()
return hashlib.sha256(buf).hexdigest()[:16]
# ---- single-model dispatch ------------------------------------------------
def _run_wd(self, tagger: ImageTagger, pil_img: Image.Image,
gen_threshold: float, char_threshold: float) -> dict:
return tagger.tag_image(
np.asarray(pil_img),
gen_threshold=gen_threshold,
char_threshold=char_threshold,
)
# ---- ensemble -------------------------------------------------------------
def tag_image(
self,
image,
gen_threshold: float = 0.35,
char_threshold: float = 0.75,
mode: str = "ensemble", # "ensemble" | "wd:eva02" | "wd:swinv2" | "wd:vit" | "deepdanbooru"
) -> dict:
if not _TAGGER_DEPS_OK and not _ORT_AVAILABLE:
raise RuntimeError("Tagger dependencies are not available.")
pil_img = _to_pil_image(image)
pil_img = ImageOps.exif_transpose(pil_img)
pil_img = _pad_square(_ensure_rgb(pil_img))
# Cache lookups are per-mode so changing the model doesn't reuse stale results.
key = f"{mode}:{self._image_key(pil_img)}"
cached = self._cache.get(key)
if cached is not None:
return cached
# Pose estimation always runs when available — it feeds `pose_tags`
# and `people_count` used below to enrich DeepDanbooru / NL captions.
from src.pose_tagger import get_pose_tagger
pose_info = get_pose_tagger().estimate(pil_img)
# DeepDanbooru-only path: no rating distribution, but pose tags are kept.
if mode == "deepdanbooru":
dd = self._dd.predict(pil_img)
gen = {k: round(v, 4) for k, v in dd.items() if v >= gen_threshold}
caption = ", ".join(gen.keys())
result = {
"caption": _esc_for_output(caption),
"taglist": _esc_for_output(caption.replace("_", " ")),
"ratings": {},
"characters": {},
"general": gen,
"pose_tags": pose_info.get("pose_tags", []),
"people_count": pose_info.get("people_count", 0),
"ensemble_votes": {"deepdanbooru": 1},
}
self._cache[key] = result
return result
# Single backbone by name
if mode.startswith("wd:"):
arch = mode.split(":", 1)[1]
tagger = self._wd_by_arch(arch)
return self._run_wd(tagger, pil_img, gen_threshold, char_threshold)
# Primary WD14 (EVA02) first so its ratings survive.
primary = self._run_wd(self._wd_taggers[0][0], pil_img, gen_threshold, char_threshold)
if mode != "ensemble":
result = primary
else:
result = self._merge_wd_results(primary, [
self._run_wd(tagger, pil_img, gen_threshold, char_threshold)
for tagger, _, _ in self._wd_taggers[1:]
])
result["ensemble_votes"] = {name: 1 for _, _, name in self._wd_taggers}
# Merge pose metadata (pose_tags are already validated lists).
result["pose_tags"] = pose_info.get("pose_tags", [])
result["people_count"] = pose_info.get("people_count", 0)
# Optional Qwen3-VL natural-language caption. Lazy-loaded on first call.
if mode in ("ensemble", "qwen"):
try:
from src.qwen_vl_tagger import get_qwen_tagger
caption = get_qwen_tagger().describe(pil_img)
if caption:
result["nl_caption"] = caption
except Exception:
pass # VL unavailable — result stays tag-only
# DeepDanbooru boosts general tags in ensemble. Only runs when the optional
# weights happen to be co-located with the app; the app keeps working when
# they are absent.
# Ensemble: DeepDanbooru optional secondary vote on top of merged WD14s.
if mode == "ensemble" and self._dd.ensure_loaded():
dd = self._dd.predict(pil_img)
if dd:
gen = result.get("general", {})
for tag, score in dd.items():
if score < gen_threshold:
continue
if tag in gen:
gen[tag] = round(max(gen[tag], score), 4)
else:
gen[tag] = round(score * 0.35, 4)
result["general"] = dict(
sorted(gen.items(), key=lambda kv: -kv[1])
)
result["caption"] = _esc_for_output(", ".join(result["general"].keys()))
result["taglist"] = _esc_for_output(", ".join(result["general"].keys()).replace("_", " "))
result["ensemble_votes"]["deepdanbooru"] = 1
self._cache[key] = result
return result
def _merge_wd_results(self, primary: dict, others: list[dict]) -> dict:
merged = {
"caption": primary["caption"],
"taglist": primary["taglist"],
"ratings": dict(primary.get("ratings", {})),
"characters": {},
"general": {},
}
gen_acc: dict[str, float] = {}
gen_cnt: dict[str, int] = {}
char_acc: dict[str, float] = {}
char_cnt: dict[str, int] = {}
def _accumulate(acc, cnt, tags, w):
for k, v in tags.items():
acc[k] = acc.get(k, 0.0) + v * w
cnt[k] = cnt.get(k, 0) + 1
for res, (_, w, _) in zip([primary, *others], self._wd_taggers):
_accumulate(gen_acc, gen_cnt, res.get("general", {}), w)
_accumulate(char_acc, char_cnt, res.get("characters", {}), w)
merged["general"] = dict(
sorted(
((k, round(v / max(gen_cnt[k], 1), 4)) for k, v in gen_acc.items()),
key=lambda kv: -kv[1],
)
)
merged["characters"] = dict(
sorted(
((k, round(v / max(char_cnt[k], 1), 4)) for k, v in char_acc.items()),
key=lambda kv: -kv[1],
)
)
merged["caption"] = _esc_for_output(", ".join(merged["general"].keys()))
merged["taglist"] = _esc_for_output(
", ".join(merged["general"].keys()).replace("_", " ")
)
return merged
_ensemble_instance: EnsembleTagger | None = None
def get_ensemble_tagger() -> EnsembleTagger:
global _ensemble_instance
if _ensemble_instance is None:
_ensemble_instance = EnsembleTagger()
return _ensemble_instance