From 42f9c344f8cd688bb27d5fc6c76b1ac531901ddd Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Tue, 10 Feb 2026 16:27:40 -0800 Subject: [PATCH] feat: add WD14 tagger support (SwinV2 + ViT) with threshold controls MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Add two WD14 classifier models as alternatives to JoyCaption LLM taggers. WD14 models are ~400MB ONNX classifiers that output Danbooru tags with confidence scores — much faster inference, no prompt needed. Backend: - Add wd14-swinv2 and wd14-vit entries to MODELS dict - Add check/download/load/generate methods for ONNX-based WD14 inference - Add general_threshold and character_threshold params to generate_tags() - Filter "prepend:*" tags from exclude prefixes and skip-tagged logic Frontend: - Add WD14 radio buttons with download status in both admin and gallery - Toggle prompt textarea vs threshold sliders based on model selection - Pass thresholds to API for batch and single-image tagging - Extract _setModelStatus() helper using safe DOM methods Dependencies: onnxruntime (optional, with clear error message) --- py/api/autotag_routes.py | 36 ++++++- py/autotag.py | 218 +++++++++++++++++++++++++++++++++++++-- web/admin.html | 51 ++++++++- web/gallery.html | 51 ++++++++- web/js/admin.js | 109 +++++++++++++++----- web/js/gallery.js | 112 +++++++++++++++----- 6 files changed, 511 insertions(+), 66 deletions(-) diff --git a/py/api/autotag_routes.py b/py/api/autotag_routes.py index 69c0102..e20c5d1 100644 --- a/py/api/autotag_routes.py +++ b/py/api/autotag_routes.py @@ -54,6 +54,8 @@ class AutotagRoutesMixin: "default_prompt": service.default_prompt, "model_loaded": service.is_model_loaded(), "loaded_model_type": service.get_loaded_model_type(), + "wd14_general_threshold": service.wd14_general_threshold, + "wd14_character_threshold": service.wd14_character_threshold, } ) @@ -122,6 +124,14 @@ class AutotagRoutesMixin: keep_in_memory = request.query.get("keep_in_memory", "true").lower() == "true" use_gpu = True + # WD14 threshold params + general_threshold = request.query.get("general_threshold") + character_threshold = request.query.get("character_threshold") + if general_threshold is not None: + general_threshold = float(general_threshold) + if character_threshold is not None: + character_threshold = float(character_threshold) + async def stream_response(): try: from ..autotag import get_autotag_service @@ -203,7 +213,11 @@ class AutotagRoutesMixin: prompt_tags = [ t.strip() for t in prompt_tags.split(",") if t.strip() ] - real_tags = [t for t in prompt_tags if t != "auto-scanned"] + real_tags = [ + t + for t in prompt_tags + if t != "auto-scanned" and not t.startswith("prepend:") + ] if real_tags: tagged_prompt_ids.add(prompt_id) skipped += 1 @@ -217,7 +231,12 @@ class AutotagRoutesMixin: try: tags = await loop.run_in_executor( - None, lambda p=str(image_path): service.generate_tags(p) + None, + lambda p=str(image_path): service.generate_tags( + p, + general_threshold=general_threshold, + character_threshold=character_threshold, + ), ) processed += 1 @@ -304,6 +323,12 @@ class AutotagRoutesMixin: model_type = data.get("model_type", "gguf") custom_prompt = data.get("prompt") use_gpu = data.get("use_gpu", True) + general_threshold = data.get("general_threshold") + character_threshold = data.get("character_threshold") + if general_threshold is not None: + general_threshold = float(general_threshold) + if character_threshold is not None: + character_threshold = float(character_threshold) if not image_path: return web.json_response( @@ -328,7 +353,12 @@ class AutotagRoutesMixin: loop = asyncio.get_event_loop() tags = await loop.run_in_executor( - None, lambda: service.generate_tags(image_path) + None, + lambda: service.generate_tags( + image_path, + general_threshold=general_threshold, + character_threshold=character_threshold, + ), ) prompt_id = None diff --git a/py/autotag.py b/py/autotag.py index 32302e7..f08c84f 100644 --- a/py/autotag.py +++ b/py/autotag.py @@ -50,6 +50,24 @@ MODELS = { "repo": "fancyfeast/llama-joycaption-beta-one-hf-llava", "subdir": "llama-joycaption-beta-one-hf-llava", }, + "wd14-swinv2": { + "name": "WD14 SwinV2", + "description": "Fast classifier tagger, ~400MB, no prompt needed", + "size": "~400MB", + "repo": "SmilingWolf/wd-swinv2-tagger-v3", + "files": ["model.onnx", "selected_tags.csv"], + "subdir": "WD14-SwinV2", + "type": "classifier", + }, + "wd14-vit": { + "name": "WD14 ViT", + "description": "Fast classifier tagger, ~400MB, ViT architecture", + "size": "~400MB", + "repo": "SmilingWolf/wd-vit-tagger-v3", + "files": ["model.onnx", "selected_tags.csv"], + "subdir": "WD14-ViT", + "type": "classifier", + }, } # Default prompts @@ -112,6 +130,10 @@ class AutoTagService: self._current_model_type: Optional[str] = None self._custom_prompt: str = DEFAULT_PROMPT + # WD14 threshold defaults + self._wd14_general_threshold: float = 0.35 + self._wd14_character_threshold: float = 0.85 + @property def models_config(self) -> Dict[str, Dict[str, Any]]: """Get the models configuration dictionary.""" @@ -132,6 +154,22 @@ class AutoTagService: """Set a custom prompt for tag generation.""" self._custom_prompt = value + @property + def wd14_general_threshold(self) -> float: + return self._wd14_general_threshold + + @wd14_general_threshold.setter + def wd14_general_threshold(self, value: float): + self._wd14_general_threshold = max(0.0, min(1.0, value)) + + @property + def wd14_character_threshold(self) -> float: + return self._wd14_character_threshold + + @wd14_character_threshold.setter + def wd14_character_threshold(self, value: float): + self._wd14_character_threshold = max(0.0, min(1.0, value)) + def get_models_status(self) -> Dict[str, Dict[str, Any]]: """Get availability status for all model types. @@ -170,6 +208,12 @@ class AutoTagService: model_status["model_path"] = str( self.models_dir / config["subdir"] / config["filename"] ) + elif model_type.startswith("wd14"): + model_status["downloaded"] = self._check_wd14_model(model_type) + if model_status["downloaded"]: + model_status["model_path"] = str( + self.models_dir / config["subdir"] / "model.onnx" + ) else: # hf model_status["downloaded"] = self._check_hf_model() if model_status["downloaded"]: @@ -269,6 +313,140 @@ class AutoTagService: return None + def _check_wd14_model(self, model_type: str) -> bool: + """Check if WD14 model files exist.""" + config = MODELS[model_type] + model_dir = self.models_dir / config["subdir"] + return all((model_dir / f).exists() for f in config["files"]) + + def _download_wd14_model( + self, + model_type: str, + progress_callback: Optional[Callable[[str, float], None]] = None, + ) -> bool: + """Download WD14 model files from HuggingFace.""" + from huggingface_hub import hf_hub_download + + config = MODELS[model_type] + model_dir = self.models_dir / config["subdir"] + model_dir.mkdir(parents=True, exist_ok=True) + + files = config["files"] + for i, filename in enumerate(files): + file_path = model_dir / filename + if not file_path.exists(): + progress = int((i / len(files)) * 80) + 10 + if progress_callback: + progress_callback(f"Downloading {filename}...", progress) + + self.logger.info(f"Downloading WD14 file: {filename}") + hf_hub_download( + repo_id=config["repo"], + filename=filename, + local_dir=str(model_dir), + local_dir_use_symlinks=False, + ) + + if progress_callback: + progress_callback("Download complete", 100) + return True + + def _load_wd14_tagger(self, model_type: str, use_gpu: bool = True): + """Load WD14 ONNX model and tag list.""" + try: + import onnxruntime as ort + except ImportError: + raise RuntimeError( + "onnxruntime is required for WD14 models. " + "Install with: pip install onnxruntime (or onnxruntime-gpu for CUDA)" + ) + import csv + + config = MODELS[model_type] + model_dir = self.models_dir / config["subdir"] + model_path = model_dir / "model.onnx" + tags_path = model_dir / "selected_tags.csv" + + self.logger.info(f"Loading WD14 model from {model_dir}...") + + # Set up ONNX session + providers = [] + if use_gpu: + providers.append("CUDAExecutionProvider") + providers.append("CPUExecutionProvider") + + session = ort.InferenceSession(str(model_path), providers=providers) + + # Read tag list from CSV + tags_list = [] + with open(tags_path, "r", encoding="utf-8") as f: + reader = csv.DictReader(f) + for row in reader: + tags_list.append(row) + + self.logger.info( + f"WD14 model loaded ({len(tags_list)} tags, " + f"provider: {session.get_providers()[0]})" + ) + return ("wd14", (session, tags_list)) + + def _generate_wd14( + self, + session, + tags_list: List[Dict[str, str]], + image: Image.Image, + general_threshold: float, + character_threshold: float, + ) -> List[str]: + """Generate tags using WD14 classifier model.""" + import numpy as np + + # Get model input dimensions + input_shape = session.get_inputs()[0].shape + target_size = input_shape[2] # e.g. 448 + + # Pad to square with white background + max_dim = max(image.size) + padded = Image.new("RGB", (max_dim, max_dim), (255, 255, 255)) + padded.paste( + image, ((max_dim - image.width) // 2, (max_dim - image.height) // 2) + ) + + # Resize to model input size + padded = padded.resize((target_size, target_size), Image.Resampling.BILINEAR) + + # Convert to numpy: RGB -> BGR, float32, add batch dim + img_array = np.array(padded, dtype=np.float32) + img_array = img_array[:, :, ::-1] # RGB to BGR + img_array = np.expand_dims(img_array, axis=0) + + # Run inference + input_name = session.get_inputs()[0].name + output = session.run(None, {input_name: img_array}) + probabilities = output[0][0] # First output, first batch + + # Filter tags by category and threshold + result_tags = [] + for i, tag_info in enumerate(tags_list): + if i >= len(probabilities): + break + prob = float(probabilities[i]) + category = tag_info.get("category", "0") + tag_name = tag_info.get("name", "").strip() + + if not tag_name: + continue + + # Category 0 = general, 4 = character, 9 = rating + if category == "0" and prob >= general_threshold: + result_tags.append((tag_name, prob)) + elif category == "4" and prob >= character_threshold: + result_tags.append((tag_name, prob)) + + # Sort by confidence descending + result_tags.sort(key=lambda x: x[1], reverse=True) + return [tag for tag, _ in result_tags] + def download_model( self, model_type: str, @@ -288,7 +466,7 @@ class AutoTagService: """ if model_type not in MODELS: raise ValueError( - f"Invalid model type: {model_type}. Must be 'gguf' or 'hf'" + f"Invalid model type: {model_type}. Must be one of: {', '.join(MODELS.keys())}" ) try: @@ -302,6 +480,8 @@ class AutoTagService: try: if model_type == "gguf": return self._download_gguf_models(progress_callback) + elif model_type.startswith("wd14"): + return self._download_wd14_model(model_type, progress_callback) else: return self._download_hf_model(progress_callback) except Exception as e: @@ -413,6 +593,8 @@ class AutoTagService: try: if model_type == "gguf": self._tagger = self._load_gguf_tagger(use_gpu) + elif model_type.startswith("wd14"): + self._tagger = self._load_wd14_tagger(model_type, use_gpu) else: self._tagger = self._load_hf_tagger() @@ -527,12 +709,20 @@ class AutoTagService: """Get the type of currently loaded model.""" return self._current_model_type - def generate_tags(self, image_path: str, prompt: Optional[str] = None) -> List[str]: + def generate_tags( + self, + image_path: str, + prompt: Optional[str] = None, + general_threshold: Optional[float] = None, + character_threshold: Optional[float] = None, + ) -> List[str]: """Generate tags for an image. Args: image_path: Path to the image file - prompt: Custom prompt for tag generation. Uses default if None. + prompt: Custom prompt for tag generation (LLM models only). Uses default if None. + general_threshold: Confidence threshold for general tags (WD14 only). + character_threshold: Confidence threshold for character tags (WD14 only). Returns: List of generated tags @@ -559,15 +749,28 @@ class AutoTagService: if model_type == "gguf": raw_tags = self._generate_gguf(tagger_obj, image, use_prompt) + return self._parse_tags(raw_tags) + elif model_type == "wd14": + session, tags_list = tagger_obj + gen_thresh = ( + general_threshold + if general_threshold is not None + else self._wd14_general_threshold + ) + char_thresh = ( + character_threshold + if character_threshold is not None + else self._wd14_character_threshold + ) + return self._generate_wd14( + session, tags_list, image, gen_thresh, char_thresh + ) else: model, processor, device, compute_dtype = tagger_obj raw_tags = self._generate_hf( model, processor, device, compute_dtype, image, use_prompt ) - - # Parse tags from response - tags = self._parse_tags(raw_tags) - return tags + return self._parse_tags(raw_tags) def _generate_gguf(self, model, image: Image.Image, prompt: str) -> str: """Generate tags using GGUF model.""" @@ -670,6 +873,7 @@ class AutoTagService: "meta:", "photo_", "photo:", + "prepend:", ) # Split by common delimiters diff --git a/web/admin.html b/web/admin.html index c3b923e..ff17d2e 100644 --- a/web/admin.html +++ b/web/admin.html @@ -980,11 +980,35 @@ Checking... + + - -
+ +
+ + +
`; - } + this._setModelStatus(ggufStatus, data.models.gguf.downloaded, 'gguf'); + this._setModelStatus(hfStatus, data.models.hf.downloaded, 'hf'); + this._setModelStatus(wd14SwinV2Status, data.models['wd14-swinv2']?.downloaded, 'wd14-swinv2'); + this._setModelStatus(wd14VitStatus, data.models['wd14-vit']?.downloaded, 'wd14-vit'); - // Update HF status - if (data.models.hf.downloaded) { - hfStatus.innerHTML = '✓ Downloaded'; - } else { - hfStatus.innerHTML = ``; + // Update WD14 threshold defaults from server + if (data.wd14_general_threshold !== undefined) { + const genSlider = document.getElementById('wd14GeneralThreshold'); + genSlider.value = data.wd14_general_threshold; + document.getElementById('wd14GeneralThresholdValue').textContent = parseFloat(data.wd14_general_threshold).toFixed(2); + } + if (data.wd14_character_threshold !== undefined) { + const charSlider = document.getElementById('wd14CharacterThreshold'); + charSlider.value = data.wd14_character_threshold; + document.getElementById('wd14CharacterThresholdValue').textContent = parseFloat(data.wd14_character_threshold).toFixed(2); } } else { - ggufStatus.innerHTML = 'Error'; - hfStatus.innerHTML = 'Error'; + statusEls.forEach(el => { el.textContent = 'Error'; }); } } catch (error) { console.error('Error checking models:', error); - ggufStatus.innerHTML = 'Error'; - hfStatus.innerHTML = 'Error'; + statusEls.forEach(el => { el.textContent = 'Error'; }); } } async downloadModel(modelType) { - const modelName = modelType === 'gguf' ? 'GGUF Model' : 'HuggingFace Model'; + const modelNames = { + 'gguf': 'GGUF Model', + 'hf': 'HuggingFace Model', + 'wd14-swinv2': 'WD14 SwinV2', + 'wd14-vit': 'WD14 ViT' + }; + const modelName = modelNames[modelType] || modelType; document.getElementById('downloadModelName').textContent = modelName; document.getElementById('downloadStatus').textContent = 'Preparing...'; document.getElementById('downloadProgressPercent').textContent = '0%'; @@ -3780,9 +3822,15 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; try { const formData = new URLSearchParams(); formData.append('model_type', modelType); - formData.append('prompt', prompt); formData.append('skip_tagged', skipTagged ? 'true' : 'false'); + if (modelType.startsWith('wd14')) { + formData.append('general_threshold', document.getElementById('wd14GeneralThreshold').value); + formData.append('character_threshold', document.getElementById('wd14CharacterThreshold').value); + } else { + formData.append('prompt', prompt); + } + this.autoTagState.eventSource = new EventSource(`/prompt_manager/autotag/start?${formData.toString()}`); this.autoTagState.eventSource.onmessage = (event) => { @@ -3870,6 +3918,8 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; this.autoTagState.reviewIndex = 0; this.autoTagState.modelType = modelType; this.autoTagState.prompt = document.getElementById('autoTagPrompt').value; + this.autoTagState.generalThreshold = parseFloat(document.getElementById('wd14GeneralThreshold').value); + this.autoTagState.characterThreshold = parseFloat(document.getElementById('wd14CharacterThreshold').value); this.autoTagState.modelLoaded = false; this.autoTagState.skipAllTagged = false; // Reset skip flag for new session @@ -3892,10 +3942,10 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; } } - // Check if tags array has "real" tags (excluding auto-scanned) + // Check if tags array has "real" tags (excluding auto-scanned and prepend:* metadata) getRealTags(tags) { if (!tags || !Array.isArray(tags)) return []; - return tags.filter(tag => tag !== 'auto-scanned'); + return tags.filter(tag => tag !== 'auto-scanned' && !tag.startsWith('prepend:')); } // Show the re-tag confirmation modal @@ -3983,14 +4033,21 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; this.autoTagState.tagsExpanded = false; // Reset accordion for new image try { + const requestBody = { + image_path: image.image_path, + model_type: this.autoTagState.modelType, + }; + if (this.autoTagState.modelType.startsWith('wd14')) { + requestBody.general_threshold = this.autoTagState.generalThreshold; + requestBody.character_threshold = this.autoTagState.characterThreshold; + } else { + requestBody.prompt = this.autoTagState.prompt; + } + const response = await fetch('/prompt_manager/autotag/single', { method: 'POST', headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ - image_path: image.image_path, - model_type: this.autoTagState.modelType, - prompt: this.autoTagState.prompt - }) + body: JSON.stringify(requestBody) }); const data = await response.json(); diff --git a/web/js/gallery.js b/web/js/gallery.js index 2200d1a..97e24e6 100644 --- a/web/js/gallery.js +++ b/web/js/gallery.js @@ -68,6 +68,23 @@ document.getElementById('cancelDownloadBtn').addEventListener('click', () => this.cancelDownload()); document.getElementById('unloadModelBtn').addEventListener('click', () => this.unloadModel()); + // WD14 model selection toggle (prompt vs thresholds) + document.querySelectorAll('input[name="autoTagModel"]').forEach(radio => { + radio.addEventListener('change', (e) => { + const isWd14 = e.target.value.startsWith('wd14'); + document.getElementById('promptSection').classList.toggle('hidden', isWd14); + document.getElementById('wd14ThresholdSection').classList.toggle('hidden', !isWd14); + }); + }); + + // WD14 threshold slider value display + document.getElementById('wd14GeneralThreshold').addEventListener('input', (e) => { + document.getElementById('wd14GeneralThresholdValue').textContent = parseFloat(e.target.value).toFixed(2); + }); + document.getElementById('wd14CharacterThreshold').addEventListener('input', (e) => { + document.getElementById('wd14CharacterThresholdValue').textContent = parseFloat(e.target.value).toFixed(2); + }); + // Add Prompt modal event listeners document.getElementById('saveNewPromptBtn').addEventListener('click', () => this.saveNewPrompt()); document.getElementById('addPromptTagInput').addEventListener('input', (e) => this.handleTagInput(e)); @@ -3078,34 +3095,46 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; await this.checkAutoTagModels(); } + // _setModelStatus uses innerHTML with hardcoded model keys (not user input) + // to render download buttons or status indicators for known model types + _setModelStatus(el, downloaded, modelKey) { + if (downloaded) { + el.textContent = ''; + const span = document.createElement('span'); + span.className = 'text-pm-success text-sm'; + span.textContent = '\u2713 Downloaded'; + el.appendChild(span); + } else { + el.textContent = ''; + const btn = document.createElement('button'); + btn.className = 'px-3 py-1 bg-pm-accent hover:bg-pm-accent-hover text-pm text-xs rounded transition-colors'; + btn.textContent = 'Download'; + btn.addEventListener('click', () => window.gallery.downloadModel(modelKey)); + el.appendChild(btn); + } + } + async checkAutoTagModels() { const ggufStatus = document.getElementById("ggufModelStatus"); const hfStatus = document.getElementById("hfModelStatus"); + const wd14SwinV2Status = document.getElementById("wd14SwinV2ModelStatus"); + const wd14VitStatus = document.getElementById("wd14VitModelStatus"); const modelLoadedStatus = document.getElementById("modelLoadedStatus"); const loadedModelType = document.getElementById("loadedModelType"); const unloadModelBtn = document.getElementById("unloadModelBtn"); - ggufStatus.innerHTML = 'Checking...'; - hfStatus.innerHTML = 'Checking...'; + const statusEls = [ggufStatus, hfStatus, wd14SwinV2Status, wd14VitStatus]; + statusEls.forEach(el => { el.textContent = 'Checking...'; }); try { const response = await fetch('/prompt_manager/autotag/models'); const data = await response.json(); if (data.success) { - // Update GGUF status - if (data.models.gguf.downloaded) { - ggufStatus.innerHTML = '✓ Downloaded'; - } else { - ggufStatus.innerHTML = ``; - } - - // Update HF status - if (data.models.hf.downloaded) { - hfStatus.innerHTML = '✓ Downloaded'; - } else { - hfStatus.innerHTML = ``; - } + this._setModelStatus(ggufStatus, data.models.gguf.downloaded, 'gguf'); + this._setModelStatus(hfStatus, data.models.hf.downloaded, 'hf'); + this._setModelStatus(wd14SwinV2Status, data.models['wd14-swinv2']?.downloaded, 'wd14-swinv2'); + this._setModelStatus(wd14VitStatus, data.models['wd14-vit']?.downloaded, 'wd14-vit'); // Update model loaded status if (data.model_loaded && data.loaded_model_type) { @@ -3116,19 +3145,35 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; modelLoadedStatus.classList.add('hidden'); unloadModelBtn.classList.add('hidden'); } + + // Update WD14 threshold defaults from server + if (data.wd14_general_threshold !== undefined) { + const genSlider = document.getElementById('wd14GeneralThreshold'); + genSlider.value = data.wd14_general_threshold; + document.getElementById('wd14GeneralThresholdValue').textContent = parseFloat(data.wd14_general_threshold).toFixed(2); + } + if (data.wd14_character_threshold !== undefined) { + const charSlider = document.getElementById('wd14CharacterThreshold'); + charSlider.value = data.wd14_character_threshold; + document.getElementById('wd14CharacterThresholdValue').textContent = parseFloat(data.wd14_character_threshold).toFixed(2); + } } else { - ggufStatus.innerHTML = 'Error'; - hfStatus.innerHTML = 'Error'; + statusEls.forEach(el => { el.textContent = 'Error'; }); } } catch (error) { console.error('Error checking models:', error); - ggufStatus.innerHTML = 'Error'; - hfStatus.innerHTML = 'Error'; + statusEls.forEach(el => { el.textContent = 'Error'; }); } } async downloadModel(modelType) { - const modelName = modelType === 'gguf' ? 'GGUF Model' : 'HuggingFace Model'; + const modelNames = { + 'gguf': 'GGUF Model', + 'hf': 'HuggingFace Model', + 'wd14-swinv2': 'WD14 SwinV2', + 'wd14-vit': 'WD14 ViT' + }; + const modelName = modelNames[modelType] || modelType; document.getElementById('downloadModelName').textContent = modelName; document.getElementById('downloadStatus').textContent = 'Preparing...'; document.getElementById('downloadProgressPercent').textContent = '0%'; @@ -3241,9 +3286,15 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; try { const formData = new URLSearchParams(); formData.append('model_type', modelType); - formData.append('prompt', prompt); formData.append('keep_in_memory', keepInMemory); + if (modelType.startsWith('wd14')) { + formData.append('general_threshold', document.getElementById('wd14GeneralThreshold').value); + formData.append('character_threshold', document.getElementById('wd14CharacterThreshold').value); + } else { + formData.append('prompt', prompt); + } + this.autoTagState.eventSource = new EventSource(`/prompt_manager/autotag/start?${formData.toString()}`); this.autoTagState.eventSource.onmessage = (event) => { @@ -3321,6 +3372,8 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; this.autoTagState.reviewIndex = 0; this.autoTagState.modelType = modelType; this.autoTagState.prompt = document.getElementById('autoTagPrompt').value; + this.autoTagState.generalThreshold = parseFloat(document.getElementById('wd14GeneralThreshold').value); + this.autoTagState.characterThreshold = parseFloat(document.getElementById('wd14CharacterThreshold').value); document.getElementById('reviewTotalCount').textContent = this.images.length; @@ -3343,14 +3396,21 @@ Seed: ${this.currentMetadata.seed || 'Unknown'}`; document.getElementById('reviewTagsContainer').innerHTML = '
Generating tags...
'; try { + const requestBody = { + image_path: image.path, + model_type: this.autoTagState.modelType, + }; + if (this.autoTagState.modelType.startsWith('wd14')) { + requestBody.general_threshold = this.autoTagState.generalThreshold; + requestBody.character_threshold = this.autoTagState.characterThreshold; + } else { + requestBody.prompt = this.autoTagState.prompt; + } + const response = await fetch('/prompt_manager/autotag/single', { method: 'POST', headers: { 'Content-Type': 'application/json' }, - body: JSON.stringify({ - image_path: image.path, - model_type: this.autoTagState.modelType, - prompt: this.autoTagState.prompt - }) + body: JSON.stringify(requestBody) }); const data = await response.json();