From 5695c4fa86e4c83dde6794704c429cf763e89456 Mon Sep 17 00:00:00 2001 From: Lev Kurbanov Date: Thu, 10 Sep 2026 01:49:24 +0700 Subject: [PATCH] Add WD14 tagger as an alternative captioning method Adds a local ONNX-based WD14 tagger (app/wd14_tagger.py) as a second captioning backend alongside the existing vision-LLM path, selectable via a new "Method" dropdown shared by both per-download captioning and the standalone "Caption an existing folder" action. Unlike the vision-LLM path, WD14 needs no server/API key -- it's a local image-classification model (the same family used by kohya_ss/taggui) that outputs a comma-separated booru-style tag list instead of a sentence. Preprocessing/inference pipeline validated live against real images before writing the final code (pad-to-square white fill, resize to the model's own input size, RGB->BGR channel order, sigmoid outputs, category-based threshold split for general vs. character tags) -- produced coherent, sensible tags on the first real try. - app/wd14_tagger.py: model loading (cached in-process, including caching a failed load so a bad model/no internet fails once instead of re-attempting a slow network call per image), preprocessing, inference, tag formatting. Four presets (vit/convnext/swinv2/eva02-large), or any compatible HF repo id typed in directly. - app/captioning.py: new Captioner interface (VisionLLMCaptioner / WD14Captioner) so jobs.py calls either backend the same way. Trigger word handling now lives here too: baked in as a prompt instruction for the vision-LLM path (unchanged), or simply prepended as the first tag for WD14 (a flat tag list has no sentence grammar for a "role" to fit into). Live-verified end to end (not just mocked): a real download job with booru search + WD14 captioning produced real, sensible tag files; the standalone caption-folder action with a trigger word correctly prepended it as the first tag on every image; the same flow driven through the actual browser UI (not just the API) produced identical results. Adds onnxruntime + huggingface_hub + numpy to requirements.txt. 27 new tests (wd14_tagger inference/caching logic, captioning.py's two backends); existing jobs tests updated for the new build_captioner() plumbing (the captioner's LLMClient now lives in app.captioning, not app.jobs). --- README.md | 42 +++++- app/captioning.py | 75 ++++++++++ app/jobs.py | 17 ++- app/models.py | 21 +++ app/wd14_tagger.py | 151 ++++++++++++++++++++ docs/landscape-and-roadmap-notes.md | 29 ++-- requirements.txt | 3 + static/app.js | 45 +++++- static/index.html | 41 +++++- tests/test_captioning.py | 134 ++++++++++++++++++ tests/test_jobs.py | 45 ++++-- tests/test_wd14_tagger.py | 204 ++++++++++++++++++++++++++++ 12 files changed, 761 insertions(+), 46 deletions(-) create mode 100644 app/captioning.py create mode 100644 app/wd14_tagger.py create mode 100644 tests/test_captioning.py create mode 100644 tests/test_wd14_tagger.py diff --git a/README.md b/README.md index fbd53f0..73ed51f 100644 --- a/README.md +++ b/README.md @@ -9,8 +9,10 @@ you choose. Optionally: - **Query LLM** — generates N additional search-query variations for each input query (Ollama / LM Studio / any OpenAI-compatible cloud API). -- **Vision LLM (captioning)** — for every downloaded image, saves a `.txt` - file with a generated description (same provider choice). +- **Captioning** — for every downloaded image, saves a `.txt` file: either + a natural-language description from a vision LLM, or a booru-style tag list + from a local **WD14 tagger** (no server/API key needed), see **Captioning + methods** below. The whole thing runs as a single local FastAPI server; the web UI (no build step) talks to it over HTTP + WebSocket. @@ -162,6 +164,40 @@ To add another search source, implement `ImageSearchProvider.search()` under `app/search/` and register it in `build_search_provider()` in `app/search/__init__.py`. +## Captioning methods + +Chosen via the **Method** dropdown at the top of the captioning card, shared +by both per-download captioning and the standalone "Caption an existing +folder" action: + +| Method | What it produces | Requirements | +|---|---|---| +| **Vision LLM** (default) | A natural-language sentence, e.g. "A rugged, jagged mountain range stretches across a deep blue sky." | A vision-capable model via Ollama/LM Studio/a cloud API -- see **Optional LLM providers** below. | +| **WD14 Tagger** | A comma-separated booru-style tag list, e.g. `solo, blue hair, outdoors, looking at viewer` -- confirmed live against real images during development. | Nothing external -- runs fully locally via ONNX Runtime (`app/wd14_tagger.py`). | + +WD14 is a local image-classification model (the same family used by +kohya_ss/taggui), not a chat model -- it outputs hundreds of candidate tags +with confidence scores, and only the ones above a threshold make it into the +caption. Character tags (a specific named character) use a separate, higher +threshold than general tags by default, since false positives there are more +disruptive to a training set. + +- **Model presets**: `wd-vit-tagger-v3` (default -- smallest/fastest, + confirmed live: ~55MB download, ~0.2s CPU inference once loaded), + `wd-convnext-tagger-v3`, `wd-swinv2-tagger-v3` (a commonly recommended + balance), `wd-eva02-large-tagger-v3` (largest/most accurate). Any other + compatible Hugging Face repo id (one with a `model.onnx` + + `selected_tags.csv`) also works, typed in directly. +- **First use downloads the model** from Hugging Face (cached afterwards by + `huggingface_hub` in its own cache dir) -- needs internet the first time for + a given model, nothing after that. +- **Trigger word + WD14**: a flat tag list doesn't have sentence grammar for a + trigger word's "role" (subject/style/action) to fit into the way a Vision + LLM prompt does, so it's simply prepended as the first tag instead -- + confirmed live: `{"enabled": true, "word": "sks_creature"}` produced + `sks_creature, animal ears, solo, ...` for every image, matching the common + LoRA/Dreambooth training convention. + ## Optional LLM providers Both blocks (query expansion and captioning) are configured independently right @@ -222,6 +258,8 @@ app/ # and standalone caption-folder job, sharing one JobState/WebSocket model download.py # async download + validation + de-dup llm_client.py # generic OpenAI-compatible API client + trigger-word prompt building + captioning.py # Captioner interface: picks vision-LLM vs WD14, bakes in the trigger word + wd14_tagger.py # local ONNX WD14-family tagger (preprocessing, inference, tag formatting) local_models.py # native folder picker + model discovery (server /v1/models + folder scan) model_registry.py # tracks which (provider, base_url, model) this process has used model_control.py # unloads a model from Ollama (native API) or LM Studio (`lms` CLI) diff --git a/app/captioning.py b/app/captioning.py new file mode 100644 index 0000000..db97166 --- /dev/null +++ b/app/captioning.py @@ -0,0 +1,75 @@ +from __future__ import annotations + +from typing import Optional, Protocol + +from . import wd14_tagger +from .llm_client import LLMClient, build_trigger_instruction +from .models import LLMConfig, TriggerWordConfig, WD14Config + + +class Captioner(Protocol): + """Common interface both captioning backends implement, so jobs.py can + call either one the same way without caring which it got.""" + + async def caption(self, image_bytes: bytes, mime: str = "image/jpeg") -> str: ... + async def aclose(self) -> None: ... + + +class VisionLLMCaptioner: + """Wraps LLMClient.caption_image() -- a vision LLM writes a natural- + language caption. Trigger word (if any) is baked in as an extra prompt + instruction at construction time.""" + + def __init__(self, llm_config: LLMConfig, trigger: Optional[TriggerWordConfig] = None): + self._client = LLMClient(llm_config) + self._instruction = build_trigger_instruction(trigger) if trigger is not None else None + + async def caption(self, image_bytes: bytes, mime: str = "image/jpeg") -> str: + return await self._client.caption_image(image_bytes, mime, extra_instruction=self._instruction) + + async def aclose(self) -> None: + await self._client.aclose() + + +class WD14Captioner: + """Wraps wd14_tagger.tag_image() -- a local ONNX classifier outputs a + comma-separated booru-style tag list instead of a sentence. A trigger + word doesn't have a "role" to play in a flat tag list the way it does in + a sentence, so it's just prepended as the first tag (the common + LoRA/Dreambooth training convention).""" + + def __init__(self, config: WD14Config, trigger: Optional[TriggerWordConfig] = None): + self._config = config + self._trigger_tag = None + if trigger is not None and trigger.enabled: + word = trigger.word.strip() + self._trigger_tag = word or None + + async def caption(self, image_bytes: bytes, mime: str = "image/jpeg") -> str: + tags = await wd14_tagger.tag_image( + image_bytes, + model=self._config.model, + general_threshold=self._config.general_threshold, + character_threshold=self._config.character_threshold, + ) + if self._trigger_tag: + return f"{self._trigger_tag}, {tags}" if tags else self._trigger_tag + return tags + + async def aclose(self) -> None: + pass # no connection to release -- inference is local and synchronous + + +def build_captioner( + method: str, + llm_config: LLMConfig, + wd14_config: WD14Config, + trigger: Optional[TriggerWordConfig] = None, +) -> Captioner: + """Takes the individual fields rather than a whole config object, since + the two call sites (a download job's nested `captioning.*` vs. the + standalone caption-folder request's top-level `method`/`llm`/`wd14`) + shape them differently.""" + if method == "wd14": + return WD14Captioner(wd14_config, trigger) + return VisionLLMCaptioner(llm_config, trigger) diff --git a/app/jobs.py b/app/jobs.py index 5032427..7aecdc9 100644 --- a/app/jobs.py +++ b/app/jobs.py @@ -9,8 +9,9 @@ import httpx +from .captioning import build_captioner from .download import download_image, sanitize_folder_name -from .llm_client import LLMClient, build_trigger_instruction +from .llm_client import LLMClient from .models import CaptionFolderRequest, JobCreateRequest from .search import build_search_provider @@ -83,7 +84,11 @@ async def run_job(self, state: JobState) -> None: await self.emit(state, "status", {"status": "running"}) expander = LLMClient(req.llm_expansion.llm) if req.llm_expansion.enabled else None - captioner = LLMClient(req.captioning.llm) if req.captioning.enabled else None + captioner = ( + build_captioner(req.captioning.method, req.captioning.llm, req.captioning.wd14) + if req.captioning.enabled + else None + ) try: search_provider = build_search_provider(req.search) allowed_formats = {f.lower().lstrip(".") for f in req.filters.formats} @@ -181,7 +186,7 @@ async def worker(): ) if captioner: try: - caption = await captioner.caption_image(res.content, res.content_type or "image/jpeg") + caption = await captioner.caption(res.content, res.content_type or "image/jpeg") txt_path = res.path.with_suffix(".txt") txt_path.write_text(caption, encoding="utf-8") state.stats["captioned"] += 1 @@ -223,7 +228,7 @@ async def run_caption_folder_job(self, state: JobState) -> None: state.status = "running" await self.emit(state, "status", {"status": "running"}) - captioner = LLMClient(req.llm) + captioner = build_captioner(req.method, req.llm, req.wd14, trigger=req.trigger) try: root = Path(req.folder) if not root.is_dir(): @@ -236,8 +241,6 @@ async def run_caption_folder_job(self, state: JobState) -> None: if not image_files: await self.emit(state, "warning", {"message": "No images found in this folder."}) - instruction = build_trigger_instruction(req.trigger) - for path in image_files: if state.cancel_requested: break @@ -254,7 +257,7 @@ async def run_caption_folder_job(self, state: JobState) -> None: try: content = path.read_bytes() mime = mimetypes.guess_type(path.name)[0] or "image/jpeg" - caption = await captioner.caption_image(content, mime, extra_instruction=instruction) + caption = await captioner.caption(content, mime) except Exception as e: state.stats["errors"] += 1 await self.emit(state, "warning", {"message": f"Captioning failed for {path.name}: {e}"}) diff --git a/app/models.py b/app/models.py index 3df5421..0538650 100644 --- a/app/models.py +++ b/app/models.py @@ -32,9 +32,28 @@ class LLMExpansionConfig(BaseModel): llm: LLMConfig = Field(default_factory=LLMConfig) +class WD14Config(BaseModel): + """Config for WD14-family tagger models (ONNX image classifiers that + output booru-style tags with confidence scores, as an alternative to a + vision LLM writing a natural-language caption). + + `model` accepts either one of the built-in presets (see + app/wd14_tagger.py MODEL_PRESETS) or a raw Hugging Face repo id for any + compatible WD14-family tagger. The first use of a given model downloads + it (cached afterwards by huggingface_hub in its own cache dir) -- no + local server/API key needed, unlike the vision-LLM path. + """ + + model: str = "wd-vit-tagger-v3" + general_threshold: float = Field(default=0.35, ge=0, le=1) + character_threshold: float = Field(default=0.85, ge=0, le=1) + + class CaptioningConfig(BaseModel): enabled: bool = False + method: Literal["vision_llm", "wd14"] = "vision_llm" llm: LLMConfig = Field(default_factory=lambda: LLMConfig(model="moondream")) + wd14: WD14Config = Field(default_factory=WD14Config) class FilterConfig(BaseModel): @@ -104,7 +123,9 @@ class CaptionFolderRequest(BaseModel): download job -- e.g. an existing dataset you want captions for.""" folder: str + method: Literal["vision_llm", "wd14"] = "vision_llm" llm: LLMConfig = Field(default_factory=lambda: LLMConfig(model="moondream")) + wd14: WD14Config = Field(default_factory=WD14Config) recursive: bool = False overwrite: bool = False trigger: TriggerWordConfig = Field(default_factory=TriggerWordConfig) diff --git a/app/wd14_tagger.py b/app/wd14_tagger.py new file mode 100644 index 0000000..75bb66a --- /dev/null +++ b/app/wd14_tagger.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +import asyncio +import csv +import io +import logging + +import numpy as np +import onnxruntime as ort +from huggingface_hub import hf_hub_download +from PIL import Image + +logger = logging.getLogger(__name__) + +# Curated presets -- HF repo ids for SmilingWolf's WD14 v3 tagger family. +# "vit" is the default: smallest download (~55MB) and fastest, while still +# giving coherent tags (confirmed live against a real image during +# development: 0.22s inference on CPU once loaded). The others trade +# download size/speed for accuracy; any other WD14-compatible HF repo id +# (one with a model.onnx + selected_tags.csv) also works, typed in directly. +MODEL_PRESETS: dict[str, str] = { + "wd-vit-tagger-v3": "SmilingWolf/wd-vit-tagger-v3", + "wd-convnext-tagger-v3": "SmilingWolf/wd-convnext-tagger-v3", + "wd-swinv2-tagger-v3": "SmilingWolf/wd-swinv2-tagger-v3", + "wd-eva02-large-tagger-v3": "SmilingWolf/wd-eva02-large-tagger-v3", +} +DEFAULT_MODEL = "wd-vit-tagger-v3" + +# A handful of emoticon-style tags conventionally keep their underscores +# (e.g. "^_^") when every other tag has "_" replaced with a space for +# readability -- this list matches the one kohya_ss/taggui use. +_KAOMOJI = { + "0_0", "(o)_(o)", "+_+", "+_-", "._.", "_", "<|>_<|>", "=_=", ">_<", + "3_3", "6_9", ">_o", "@_@", "^_^", "o_o", "u_u", "x_x", "|_|", "||_||", +} + +# selected_tags.csv "category" column: 9=rating, 4=character, 0=general +# (everything else, e.g. copyright/artist tags on some models, is treated +# like "general" here rather than special-cased). +_RATING_CATEGORY = "9" +_CHARACTER_CATEGORY = "4" + +# Loaded ONNX sessions + tag tables are cached in-process, keyed by resolved +# model repo id -- (re)loading is the slow part (a network fetch on first +# use, real init cost every time after), and one job can tag many images in +# a row. A failed load is cached too (as the exception itself), so a bad +# model/no internet fails once loudly instead of re-attempting a slow +# network call for every single image in the job. +_cache: dict[str, tuple | Exception] = {} + + +def resolve_model_repo(model: str) -> str: + """Accepts either one of the short presets above or a raw HF repo id.""" + return MODEL_PRESETS.get(model, model) + + +def _load(model: str) -> tuple: + repo_id = resolve_model_repo(model) + cached = _cache.get(repo_id) + if isinstance(cached, Exception): + raise cached + if cached is not None: + return cached + + try: + model_path = hf_hub_download(repo_id, "model.onnx") + csv_path = hf_hub_download(repo_id, "selected_tags.csv") + session = ort.InferenceSession(model_path, providers=["CPUExecutionProvider"]) + input_info = session.get_inputs()[0] + _, height, width, _ = input_info.shape + output_name = session.get_outputs()[0].name + with open(csv_path, encoding="utf-8-sig") as f: + rows = list(csv.DictReader(f)) + loaded = (session, input_info.name, output_name, int(width), int(height), rows) + except Exception as e: + logger.exception("Failed to load WD14 model %r", repo_id) + _cache[repo_id] = e + raise + _cache[repo_id] = loaded + return loaded + + +def _preprocess(image_bytes: bytes, width: int, height: int) -> np.ndarray: + img = Image.open(io.BytesIO(image_bytes)) + if img.mode in ("RGBA", "LA") or (img.mode == "P" and "transparency" in img.info): + img = img.convert("RGBA") + bg = Image.new("RGBA", img.size, (255, 255, 255, 255)) + img = Image.alpha_composite(bg, img).convert("RGB") + else: + img = img.convert("RGB") + + # Pad to square (centered, white fill) before resizing, so the aspect + # ratio isn't distorted by a naive stretch to the model's fixed input size. + w, h = img.size + size = max(w, h) + canvas = Image.new("RGB", (size, size), (255, 255, 255)) + canvas.paste(img, ((size - w) // 2, (size - h) // 2)) + canvas = canvas.resize((width, height), Image.LANCZOS) + + arr = np.asarray(canvas, dtype=np.float32) + arr = arr[:, :, ::-1] # RGB -> BGR: these models were trained on BGR input + return np.expand_dims(arr, axis=0) + + +def _format_tag(name: str) -> str: + if name in _KAOMOJI: + return name + return name.replace("_", " ") + + +def _run_sync(image_bytes: bytes, model: str, general_threshold: float, character_threshold: float) -> str: + session, input_name, output_name, width, height, rows = _load(model) + arr = _preprocess(image_bytes, width, height) + probs = session.run([output_name], {input_name: arr})[0][0] + + general: list[tuple[str, float]] = [] + character: list[tuple[str, float]] = [] + for row, p in zip(rows, probs): + p = float(p) + cat = row["category"] + if cat == _RATING_CATEGORY: + continue + if cat == _CHARACTER_CATEGORY: + if p >= character_threshold: + character.append((row["name"], p)) + elif p >= general_threshold: + general.append((row["name"], p)) + + character.sort(key=lambda x: -x[1]) + general.sort(key=lambda x: -x[1]) + tags = [_format_tag(n) for n, _ in character] + [_format_tag(n) for n, _ in general] + return ", ".join(tags) + + +async def tag_image( + image_bytes: bytes, + model: str = DEFAULT_MODEL, + general_threshold: float = 0.35, + character_threshold: float = 0.85, +) -> str: + """Tag an image with a WD14-family ONNX classifier, returning a + comma-separated booru-style tag string (character tags first, then + general tags, both ranked by confidence -- highest first). + + The first call for a given model downloads it from Hugging Face (cached + by huggingface_hub afterwards) and loads the ONNX session, which can take + a while; every call after that is fast (confirmed live: ~0.2s CPU + inference for the default "vit" preset once loaded). Runs in a thread + since both the model load and inference are blocking/CPU-bound. + """ + return await asyncio.to_thread(_run_sync, image_bytes, model, general_threshold, character_threshold) diff --git a/docs/landscape-and-roadmap-notes.md b/docs/landscape-and-roadmap-notes.md index 13b490b..f694376 100644 --- a/docs/landscape-and-roadmap-notes.md +++ b/docs/landscape-and-roadmap-notes.md @@ -54,32 +54,31 @@ The closer competition, especially for the trigger-word feature: ## Honest weak spots vs. the mature alternatives -- **WD14 tagger** (taggui/kohya_ss) usually gives more consistent, - "trainable" tags for anime/booru-style work than a generic vision-LLM - caption. That's a specialized, widely-used, fine-tuned tool -- DatasetForge - doesn't try to compete with it head-on yet. -- **gallery-dl** is far more robust and covers vastly more sites. DatasetForge - is pinned to one engine (DuckDuckGo via `ddgs`), not an affiliated API, and - can break if they change something on their end. +- **gallery-dl** is far more robust and covers vastly more sites than our + four search providers combined, and doesn't rely on scraping search-engine + result pages the way DuckDuckGo/Yandex/Google here do. - No dataset-wide batch tag editing (find/replace across already-written `.txt` files). - No similarity-based de-dup, only exact content-hash de-dup. - No integration with the actual training step (kohya_ss trains; DatasetForge stops at "captioned folder ready for a trainer"). +## Shipped since these notes were written + +- ✅ **WD14/booru tagger mode** -- `app/wd14_tagger.py` + `app/captioning.py`. + Local ONNX inference, no server needed; trigger word prepended as the first + tag. Confirmed live against real images. +- ✅ **A second (third, fourth...) search backend** -- Yandex, Google, and + booru boards (e621/gelbooru/rule34/danbooru) alongside DuckDuckGo, via + `build_search_provider()` in `app/search/__init__.py`. + ## Candidate next steps (not decided -- discuss before building any of these) Roughly in the order the weak spots above suggest, not a priority order: -- **WD14/booru tagger mode** as an alternative captioning backend alongside - the vision-LLM path, for anime/booru-style datasets where tag lists train - better than natural-language captions. - **Batch find/replace across existing captions** -- edit tags/captions for a - folder that's already been captioned, without re-running the vision model. -- **A second search backend** (e.g. Google Custom Search API, Bing, or an - aggregator like SerpAPI) so the tool isn't a single point of failure tied - to one unofficial engine -- `ImageSearchProvider` in `app/search/` was - already designed to make this a drop-in addition. + folder that's already been captioned, without re-running the vision model + or tagger. - **Similarity-based de-dup** (perceptual hash) in addition to exact content-hash de-dup, to catch near-duplicate images from different sources. - **Dataset-level review UI** -- browse/filter/bulk-delete already-downloaded diff --git a/requirements.txt b/requirements.txt index be7a2b8..e5335ca 100644 --- a/requirements.txt +++ b/requirements.txt @@ -4,3 +4,6 @@ httpx>=0.27 pillow>=10.4 pydantic>=2.8 ddgs>=9.0 +numpy>=1.26 +onnxruntime>=1.18 +huggingface_hub>=0.24 diff --git a/static/app.js b/static/app.js index b3b9191..9e25bce 100644 --- a/static/app.js +++ b/static/app.js @@ -14,15 +14,33 @@ function wireProviderDefaults(providerSelect, baseUrlInput) { wireProviderDefaults($("llm_provider"), $("llm_base_url")); wireProviderDefaults($("cap_provider"), $("cap_base_url")); -function wireToggle(checkbox, block) { - const sync = () => block.classList.toggle("disabled", !checkbox.checked); +function wireToggle(checkbox, ...blocks) { + const sync = () => blocks.forEach((b) => b.classList.toggle("disabled", !checkbox.checked)); checkbox.addEventListener("change", sync); sync(); } wireToggle($("llm_enabled"), $("llm_block")); -wireToggle($("cap_enabled"), $("cap_block")); +wireToggle($("cap_enabled"), $("cap_vision_block"), $("cap_wd14_block")); wireToggle($("trigger_enabled"), $("trigger_block")); +// ---- captioning method: vision LLM (a sentence) vs WD14 tagger (a booru- +// style tag list, local ONNX, no server needed) -- only one settings block +// is relevant at a time, and the trigger-word note below reads differently +// depending on which one is active (role-aware prompt vs. flat prepend). ---- +const CAP_METHOD_NOTES = { + vision_llm: "Needs a vision-capable model reachable via Ollama/LM Studio/a cloud API -- configure it in Settings below.", + wd14: "Runs locally via ONNX -- no model server needed. Outputs a comma-separated tag list (e.g. \"solo, blue hair, outdoors\") instead of a sentence.", +}; + +function syncCapMethodUI() { + const method = $("cap_method").value; + $("cap_vision_block").hidden = method !== "vision_llm"; + $("cap_wd14_block").hidden = method !== "wd14"; + $("cap-method-note").textContent = CAP_METHOD_NOTES[method] || ""; +} +$("cap_method").addEventListener("change", syncCapMethodUI); +syncCapMethodUI(); + function syncGenerateBtn() { $("generate-btn").disabled = !$("llm_enabled").checked; } @@ -195,7 +213,8 @@ const FIELD_IDS = [ "search_provider", "booru_site", "booru_api_key", "booru_user_id", "booru_login", "min_width", "min_height", "safesearch", "llm_enabled", "llm_provider", "llm_variations", "llm_base_url", "llm_model", "llm_api_key", "llm_timeout", "llm_models_folder", "llm_disable_reasoning", - "cap_enabled", "cap_provider", "cap_model", "cap_base_url", "cap_api_key", "cap_timeout", "cap_models_folder", "cap_disable_reasoning", + "cap_enabled", "cap_method", "cap_provider", "cap_model", "cap_base_url", "cap_api_key", "cap_timeout", "cap_models_folder", "cap_disable_reasoning", + "cap_wd14_model", "cap_wd14_general_threshold", "cap_wd14_character_threshold", "capfolder_path", "capfolder_recursive", "capfolder_overwrite", "trigger_enabled", "trigger_word", "trigger_role", "trigger_custom", ]; @@ -231,13 +250,15 @@ function loadForm() { }); } $("llm_block").classList.toggle("disabled", !$("llm_enabled").checked); - $("cap_block").classList.toggle("disabled", !$("cap_enabled").checked); + $("cap_vision_block").classList.toggle("disabled", !$("cap_enabled").checked); + $("cap_wd14_block").classList.toggle("disabled", !$("cap_enabled").checked); $("trigger_block").classList.toggle("disabled", !$("trigger_enabled").checked); } loadForm(); syncGenerateBtn(); syncTriggerCustomField(); syncSearchProviderUI(); +syncCapMethodUI(); document.getElementById("job-form").addEventListener("change", saveForm); // ---- query generation: expand queries with the LLM and write the result back @@ -353,6 +374,7 @@ function buildRequest() { }, captioning: { enabled: $("cap_enabled").checked, + method: $("cap_method").value, llm: { provider: $("cap_provider").value, base_url: $("cap_base_url").value.trim(), @@ -361,10 +383,21 @@ function buildRequest() { timeout_seconds: Number($("cap_timeout").value) || 180, disable_reasoning: $("cap_disable_reasoning").checked, }, + wd14: buildWD14Config(), }, }; } +// Shared by the per-download captioning config and the standalone +// caption-folder request -- both use the same WD14 fields. +function buildWD14Config() { + return { + model: $("cap_wd14_model").value.trim() || "wd-vit-tagger-v3", + general_threshold: Number($("cap_wd14_general_threshold").value) || 0.35, + character_threshold: Number($("cap_wd14_character_threshold").value) || 0.85, + }; +} + // ---- job run + websocket progress (shared by the download job and the // standalone "caption a folder" action -- only one can run at a time) ---- let ws = null; @@ -586,6 +619,7 @@ document.getElementById("job-form").addEventListener("submit", async (e) => { function buildCaptionFolderRequest() { return { folder: $("capfolder_path").value.trim(), + method: $("cap_method").value, recursive: $("capfolder_recursive").checked, overwrite: $("capfolder_overwrite").checked, llm: { @@ -596,6 +630,7 @@ function buildCaptionFolderRequest() { timeout_seconds: Number($("cap_timeout").value) || 180, disable_reasoning: $("cap_disable_reasoning").checked, }, + wd14: buildWD14Config(), trigger: { enabled: $("trigger_enabled").checked, word: $("trigger_word").value.trim(), diff --git a/static/index.html b/static/index.html index 3f4d4fb..4a0606d 100644 --- a/static/index.html +++ b/static/index.html @@ -190,12 +190,21 @@

Filters

+
+ + +
+

+
Settings -
+
+ + + + + + +
+
+
+ + +
+
+ + +
+
+

Runs fully locally via ONNX Runtime -- no Ollama/LM Studio/cloud API needed. First use downloads the model from Hugging Face (roughly 50–800MB depending on the preset) and caches it; every run after that is fast (confirmed live: ~0.2s per image on CPU for the default preset once loaded).

+

Caption an existing folder

-

Captions every image already in a folder, independent of any download above — uses the vision LLM settings from "Settings" above.

+

Captions every image already in a folder, independent of any download above — uses the Method and Settings configured above (Vision LLM or WD14 Tagger).

@@ -270,7 +303,7 @@

Caption an existing folder

Use a trigger word
-

A trigger word is a token every caption should consistently include — common for LoRA/Dreambooth training sets. What it stands for changes how it should read in a sentence, so tell us its role:

+

A trigger word is a token every caption should consistently include — common for LoRA/Dreambooth training sets. With the Vision LLM method, what it stands for changes how it should read in a sentence, so tell us its role below. With WD14 Tagger, it's simply prepended as the first tag regardless of role (a flat tag list doesn't have sentence grammar to fit a role into).

diff --git a/tests/test_captioning.py b/tests/test_captioning.py new file mode 100644 index 0000000..cf1eba9 --- /dev/null +++ b/tests/test_captioning.py @@ -0,0 +1,134 @@ +from __future__ import annotations + +from app import captioning as captioning_module +from app.captioning import VisionLLMCaptioner, WD14Captioner, build_captioner +from app.models import LLMConfig, TriggerWordConfig, WD14Config + + +class FakeLLMClient: + instances: list["FakeLLMClient"] = [] + + def __init__(self, config): + self.config = config + self.closed = False + self.caption_calls: list[tuple[bytes, str, str | None]] = [] + FakeLLMClient.instances.append(self) + + async def caption_image(self, image_bytes, mime="image/jpeg", extra_instruction=None): + self.caption_calls.append((image_bytes, mime, extra_instruction)) + return "a vision caption" + + async def aclose(self): + self.closed = True + + +class TestBuildCaptioner: + def test_vision_llm_is_the_default(self): + captioner = build_captioner("vision_llm", LLMConfig(), WD14Config()) + assert isinstance(captioner, VisionLLMCaptioner) + + def test_wd14_method_returns_wd14_captioner(self): + captioner = build_captioner("wd14", LLMConfig(), WD14Config()) + assert isinstance(captioner, WD14Captioner) + + +class TestVisionLLMCaptioner: + async def test_delegates_to_llm_client_caption_image(self, monkeypatch): + FakeLLMClient.instances.clear() + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) + + captioner = VisionLLMCaptioner(LLMConfig()) + result = await captioner.caption(b"bytes", "image/png") + + assert result == "a vision caption" + client = FakeLLMClient.instances[0] + assert client.caption_calls == [(b"bytes", "image/png", None)] + + async def test_trigger_instruction_is_computed_once_and_reused(self, monkeypatch): + FakeLLMClient.instances.clear() + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) + + trigger = TriggerWordConfig(enabled=True, word="zog", role="subject") + captioner = VisionLLMCaptioner(LLMConfig(), trigger) + await captioner.caption(b"a") + await captioner.caption(b"b") + + client = FakeLLMClient.instances[0] + instructions = [c[2] for c in client.caption_calls] + assert instructions[0] is not None and "zog" in instructions[0] + assert instructions[0] == instructions[1] # same instruction both times, not recomputed + + async def test_no_trigger_means_no_instruction(self, monkeypatch): + FakeLLMClient.instances.clear() + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) + + captioner = VisionLLMCaptioner(LLMConfig(), TriggerWordConfig(enabled=False)) + await captioner.caption(b"a") + + assert FakeLLMClient.instances[0].caption_calls[0][2] is None + + async def test_aclose_delegates_to_the_underlying_client(self, monkeypatch): + FakeLLMClient.instances.clear() + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) + + captioner = VisionLLMCaptioner(LLMConfig()) + await captioner.aclose() + + assert FakeLLMClient.instances[0].closed is True + + +class TestWD14Captioner: + async def test_delegates_to_tag_image_with_configured_thresholds(self, monkeypatch): + captured = {} + + async def fake_tag_image(image_bytes, model, general_threshold, character_threshold): + captured["args"] = (image_bytes, model, general_threshold, character_threshold) + return "solo, blue hair" + + monkeypatch.setattr(captioning_module.wd14_tagger, "tag_image", fake_tag_image) + + config = WD14Config(model="wd-swinv2-tagger-v3", general_threshold=0.4, character_threshold=0.8) + captioner = WD14Captioner(config) + result = await captioner.caption(b"bytes", "image/jpeg") + + assert result == "solo, blue hair" + assert captured["args"] == (b"bytes", "wd-swinv2-tagger-v3", 0.4, 0.8) + + async def test_trigger_word_is_prepended_as_the_first_tag(self, monkeypatch): + async def fake_tag_image(image_bytes, model, general_threshold, character_threshold): + return "solo, blue hair" + + monkeypatch.setattr(captioning_module.wd14_tagger, "tag_image", fake_tag_image) + + trigger = TriggerWordConfig(enabled=True, word="zog", role="style") # role is irrelevant for WD14 + captioner = WD14Captioner(WD14Config(), trigger) + result = await captioner.caption(b"bytes") + + assert result == "zog, solo, blue hair" + + async def test_trigger_word_alone_when_no_tags_pass_threshold(self, monkeypatch): + async def fake_tag_image(image_bytes, model, general_threshold, character_threshold): + return "" + + monkeypatch.setattr(captioning_module.wd14_tagger, "tag_image", fake_tag_image) + + trigger = TriggerWordConfig(enabled=True, word="zog") + captioner = WD14Captioner(WD14Config(), trigger) + result = await captioner.caption(b"bytes") + + assert result == "zog" + + async def test_disabled_trigger_does_not_prepend_anything(self, monkeypatch): + async def fake_tag_image(image_bytes, model, general_threshold, character_threshold): + return "solo" + + monkeypatch.setattr(captioning_module.wd14_tagger, "tag_image", fake_tag_image) + + captioner = WD14Captioner(WD14Config(), TriggerWordConfig(enabled=False, word="zog")) + result = await captioner.caption(b"bytes") + + assert result == "solo" + + async def test_aclose_is_a_no_op(self): + captioner = WD14Captioner(WD14Config()) + await captioner.aclose() # must not raise diff --git a/tests/test_jobs.py b/tests/test_jobs.py index 6bcd66f..bdaf1ea 100644 --- a/tests/test_jobs.py +++ b/tests/test_jobs.py @@ -4,6 +4,7 @@ import pytest +from app import captioning as captioning_module from app import jobs as jobs_module from app.download import DownloadResult from app.jobs import JobManager @@ -67,6 +68,26 @@ async def aclose(self): self.closed = True +class FakeCaptioner: + """Stand-in for a Captioner (app/captioning.py) passed directly to + _process_query -- unlike FakeLLMClient, this matches the `.caption()` + protocol jobs.py actually calls, without going through build_captioner().""" + + def __init__(self, caption_text: str = "a fake caption", raise_error: Exception | None = None): + self.caption_text = caption_text + self.raise_error = raise_error + self.calls: list[bytes] = [] + + async def caption(self, image_bytes, mime="image/jpeg"): + self.calls.append(image_bytes) + if self.raise_error: + raise self.raise_error + return self.caption_text + + async def aclose(self): + pass + + @pytest.fixture(autouse=True) def reset_fake_llm_instances(): FakeLLMClient.instances.clear() @@ -246,7 +267,7 @@ async def fake_download_image(client, result, dest_dir, index, name_prefix, req = make_request(n_per_query=1, captioning=CaptioningConfig(enabled=True)) state = manager.create_job(req) provider = FakeSearchProvider() - captioner = FakeLLMClient(req.captioning.llm) + captioner = FakeCaptioner() await manager._process_query( state, "cats", req, provider, None, captioner, {"jpg"}, tmp_path, None @@ -272,14 +293,10 @@ async def fake_download_image(client, result, dest_dir, index, name_prefix, monkeypatch.setattr(jobs_module, "download_image", fake_download_image) - class BrokenCaptioner(FakeLLMClient): - async def caption_image(self, *a, **kw): - raise RuntimeError("vision model down") - req = make_request(n_per_query=1, captioning=CaptioningConfig(enabled=True)) state = manager.create_job(req) provider = FakeSearchProvider() - captioner = BrokenCaptioner(req.captioning.llm) + captioner = FakeCaptioner(raise_error=RuntimeError("vision model down")) await manager._process_query( state, "cats", req, provider, None, captioner, {"jpg"}, tmp_path, None @@ -302,7 +319,8 @@ async def fake_download_image(client, result, dest_dir, index, name_prefix, monkeypatch.setattr(jobs_module, "download_image", fake_download_image) monkeypatch.setattr(jobs_module, "build_search_provider", lambda config: FakeSearchProvider()) - monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) + monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) # the query-expander + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) # inside build_captioner() req = make_request( queries=["cats", "dogs"], @@ -344,6 +362,7 @@ async def fake_download_image(client, result, dest_dir, index, name_prefix, async def test_llm_clients_are_closed_even_when_the_job_errors(self, manager, monkeypatch, tmp_path): monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) def broken_provider(config): raise RuntimeError("provider init failed") @@ -371,7 +390,7 @@ def _make_image(self, path: Path) -> None: async def test_captions_every_image_in_a_flat_folder(self, manager, monkeypatch, tmp_path): self._make_image(tmp_path / "a.jpg") self._make_image(tmp_path / "b.png") - monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) req = CaptionFolderRequest(folder=str(tmp_path), llm=LLMConfig()) state = manager.create_caption_folder_job(req) @@ -387,7 +406,7 @@ async def test_captions_every_image_in_a_flat_folder(self, manager, monkeypatch, async def test_skips_images_that_already_have_a_caption(self, manager, monkeypatch, tmp_path): self._make_image(tmp_path / "a.jpg") (tmp_path / "a.txt").write_text("existing caption", encoding="utf-8") - monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) req = CaptionFolderRequest(folder=str(tmp_path), llm=LLMConfig(), overwrite=False) state = manager.create_caption_folder_job(req) @@ -402,7 +421,7 @@ async def test_skips_images_that_already_have_a_caption(self, manager, monkeypat async def test_overwrite_forces_recaptioning(self, manager, monkeypatch, tmp_path): self._make_image(tmp_path / "a.jpg") (tmp_path / "a.txt").write_text("stale caption", encoding="utf-8") - monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) req = CaptionFolderRequest(folder=str(tmp_path), llm=LLMConfig(), overwrite=True) state = manager.create_caption_folder_job(req) @@ -416,7 +435,7 @@ async def test_recursive_flag_includes_nested_images(self, manager, monkeypatch, nested = tmp_path / "sub" nested.mkdir() self._make_image(nested / "a.jpg") - monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) non_recursive_req = CaptionFolderRequest(folder=str(tmp_path), llm=LLMConfig(), recursive=False) state = manager.create_caption_folder_job(non_recursive_req) @@ -429,7 +448,7 @@ async def test_recursive_flag_includes_nested_images(self, manager, monkeypatch, assert state2.stats["captioned"] == 1 async def test_missing_folder_reports_error_status(self, manager, monkeypatch, tmp_path): - monkeypatch.setattr(jobs_module, "LLMClient", FakeLLMClient) + monkeypatch.setattr(captioning_module, "LLMClient", FakeLLMClient) req = CaptionFolderRequest(folder=str(tmp_path / "nope"), llm=LLMConfig()) state = manager.create_caption_folder_job(req) @@ -448,7 +467,7 @@ async def caption_image(self, image_bytes, mime="image/jpeg", extra_instruction= received["instruction"] = extra_instruction return await super().caption_image(image_bytes, mime, extra_instruction) - monkeypatch.setattr(jobs_module, "LLMClient", RecordingCaptioner) + monkeypatch.setattr(captioning_module, "LLMClient", RecordingCaptioner) from app.models import TriggerWordConfig diff --git a/tests/test_wd14_tagger.py b/tests/test_wd14_tagger.py new file mode 100644 index 0000000..ab6184e --- /dev/null +++ b/tests/test_wd14_tagger.py @@ -0,0 +1,204 @@ +from __future__ import annotations + +import numpy as np +import pytest + +from app import wd14_tagger + + +class FakeSession: + """Stand-in for onnxruntime.InferenceSession -- returns a fixed + probability vector regardless of the actual preprocessed input, so tests + can focus on the threshold/category/sort/format logic around it.""" + + def __init__(self, probs: list[float]): + self._probs = probs + + def run(self, output_names, input_feed): + return [np.array([self._probs], dtype=np.float32)] + + +def make_rows(*entries: tuple[str, str]) -> list[dict]: + """entries: (name, category) pairs, in the same order as the fake probs.""" + return [{"tag_id": str(i), "name": name, "category": cat, "count": "0"} for i, (name, cat) in enumerate(entries)] + + +@pytest.fixture(autouse=True) +def clear_cache(): + wd14_tagger._cache.clear() + yield + wd14_tagger._cache.clear() + + +class TestResolveModelRepo: + def test_known_preset_resolves_to_full_repo_id(self): + assert wd14_tagger.resolve_model_repo("wd-vit-tagger-v3") == "SmilingWolf/wd-vit-tagger-v3" + + def test_unknown_string_passes_through_unchanged(self): + assert wd14_tagger.resolve_model_repo("someone/custom-wd14-repo") == "someone/custom-wd14-repo" + + def test_default_model_is_a_known_preset(self): + assert wd14_tagger.DEFAULT_MODEL in wd14_tagger.MODEL_PRESETS + + +class TestFormatTag: + def test_underscore_replaced_with_space(self): + assert wd14_tagger._format_tag("blue_hair") == "blue hair" + + def test_kaomoji_keeps_its_underscore(self): + assert wd14_tagger._format_tag("^_^") == "^_^" + assert wd14_tagger._format_tag("o_o") == "o_o" + + def test_tag_with_no_underscore_is_unchanged(self): + assert wd14_tagger._format_tag("solo") == "solo" + + +class TestTagImage: + async def test_general_tags_above_threshold_are_included(self, monkeypatch, jpeg_bytes): + rows = make_rows(("rating", "9"), ("solo", "0"), ("background", "0")) + session = FakeSession([0.9, 0.5, 0.1]) # rating ignored, solo passes 0.35, background doesn't + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session, "in", "out", 448, 448, rows)) + + result = await wd14_tagger.tag_image(jpeg_bytes, general_threshold=0.35) + assert result == "solo" + + async def test_character_tags_use_their_own_higher_threshold(self, monkeypatch, jpeg_bytes): + rows = make_rows(("some_character", "4"), ("generic_tag", "0")) + session = FakeSession([0.6, 0.4]) # character: 0.6 < default 0.85 -> excluded; general: 0.4 >= 0.35 -> included + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session, "in", "out", 448, 448, rows)) + + result = await wd14_tagger.tag_image(jpeg_bytes) + assert result == "generic tag" + + session2 = FakeSession([0.9, 0.4]) # now character clears 0.85 too + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session2, "in", "out", 448, 448, rows)) + result2 = await wd14_tagger.tag_image(jpeg_bytes) + assert result2 == "some character, generic tag" # character tags sorted before general + + async def test_rating_category_is_never_included_in_output(self, monkeypatch, jpeg_bytes): + rows = make_rows(("explicit", "9")) + session = FakeSession([0.99]) + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session, "in", "out", 448, 448, rows)) + + result = await wd14_tagger.tag_image(jpeg_bytes) + assert result == "" + + async def test_general_tags_sorted_by_confidence_descending(self, monkeypatch, jpeg_bytes): + rows = make_rows(("low_conf", "0"), ("high_conf", "0"), ("mid_conf", "0")) + session = FakeSession([0.4, 0.9, 0.6]) + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session, "in", "out", 448, 448, rows)) + + result = await wd14_tagger.tag_image(jpeg_bytes) + assert result == "high conf, mid conf, low conf" + + async def test_no_tags_above_threshold_returns_empty_string(self, monkeypatch, jpeg_bytes): + rows = make_rows(("obscure", "0")) + session = FakeSession([0.01]) + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session, "in", "out", 448, 448, rows)) + + result = await wd14_tagger.tag_image(jpeg_bytes) + assert result == "" + + async def test_custom_thresholds_are_respected(self, monkeypatch, jpeg_bytes): + rows = make_rows(("tag_a", "0")) + session = FakeSession([0.5]) + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session, "in", "out", 448, 448, rows)) + + assert await wd14_tagger.tag_image(jpeg_bytes, general_threshold=0.6) == "" + assert await wd14_tagger.tag_image(jpeg_bytes, general_threshold=0.4) == "tag a" + + async def test_works_with_png_and_transparency(self, monkeypatch, png_bytes): + rows = make_rows(("solo", "0")) + session = FakeSession([0.9]) + monkeypatch.setattr(wd14_tagger, "_load", lambda model: (session, "in", "out", 448, 448, rows)) + + result = await wd14_tagger.tag_image(png_bytes) + assert result == "solo" + + +class TestLoadCaching: + def test_successful_load_is_cached_across_calls(self, monkeypatch, tmp_path): + calls = {"download": 0, "session": 0} + + def fake_hf_hub_download(repo_id, filename): + calls["download"] += 1 + return str(tmp_path / filename) + + class FakeOrtSession: + def __init__(self, *a, **kw): + calls["session"] += 1 + + def get_inputs(self): + class Info: + name = "input" + shape = ["batch", 448, 448, 3] + + return [Info()] + + def get_outputs(self): + class Info: + name = "output" + + return [Info()] + + (tmp_path / "selected_tags.csv").write_text("tag_id,name,category,count\n0,solo,0,1\n", encoding="utf-8") + + monkeypatch.setattr(wd14_tagger, "hf_hub_download", fake_hf_hub_download) + monkeypatch.setattr(wd14_tagger.ort, "InferenceSession", FakeOrtSession) + + wd14_tagger._load("wd-vit-tagger-v3") + wd14_tagger._load("wd-vit-tagger-v3") + + assert calls["download"] == 2 # model.onnx + selected_tags.csv, once + assert calls["session"] == 1 # not recreated on the second _load() + + def test_failed_load_is_cached_and_reraised_without_retrying(self, monkeypatch): + attempts = {"n": 0} + + def failing_download(repo_id, filename): + attempts["n"] += 1 + raise RuntimeError("network unreachable") + + monkeypatch.setattr(wd14_tagger, "hf_hub_download", failing_download) + + with pytest.raises(RuntimeError, match="network unreachable"): + wd14_tagger._load("wd-vit-tagger-v3") + with pytest.raises(RuntimeError, match="network unreachable"): + wd14_tagger._load("wd-vit-tagger-v3") + + assert attempts["n"] == 1 # second call hit the cached exception, not the network again + + def test_different_models_are_cached_separately(self, monkeypatch, tmp_path): + seen_repos = [] + + def fake_hf_hub_download(repo_id, filename): + seen_repos.append(repo_id) + return str(tmp_path / filename) + + class FakeOrtSession: + def __init__(self, *a, **kw): + pass + + def get_inputs(self): + class Info: + name = "input" + shape = ["batch", 448, 448, 3] + + return [Info()] + + def get_outputs(self): + class Info: + name = "output" + + return [Info()] + + (tmp_path / "selected_tags.csv").write_text("tag_id,name,category,count\n0,solo,0,1\n", encoding="utf-8") + + monkeypatch.setattr(wd14_tagger, "hf_hub_download", fake_hf_hub_download) + monkeypatch.setattr(wd14_tagger.ort, "InferenceSession", FakeOrtSession) + + wd14_tagger._load("wd-vit-tagger-v3") + wd14_tagger._load("wd-swinv2-tagger-v3") + + assert "SmilingWolf/wd-vit-tagger-v3" in seen_repos + assert "SmilingWolf/wd-swinv2-tagger-v3" in seen_repos