mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
refactor(caption): unify tagger settings and reorganize Caption Tab UI
Consolidate WD14 and DeepBooru tagger settings into unified options: - Merge wd14_general_threshold + deepbooru_score_threshold → tagger_threshold - Merge wd14_include_rating + deepbooru_include_rating → tagger_include_rating - Rename interrogate_score → tagger_show_scores - Rename tagger_escape → tagger_escape_brackets - Rename CLiP → OpenCLiP in caption type choices UI reorganization: - Add Interrogate tab to Caption Tab with default caption type selector - Move interrogate_offload to Model Offloading section as "Offload caption models" - Hide Interrogate settings section (all settings now in Caption Tab UI) - Update locale_en.json for OpenCLiP naming Code improvements: - DeepBooru tag_multi() now accepts same parameters as WD14 for unified interface - Fix setting references in interrogate.py for consolidated settings - Add comprehensive tagger test suite (cli/test-tagger.py)
This commit is contained in:
@@ -0,0 +1,849 @@
|
||||
#!/usr/bin/env python
|
||||
"""
|
||||
Tagger Settings Test Suite
|
||||
|
||||
Tests all WD14 and DeepBooru tagger settings to verify they're properly
|
||||
mapped and affect output correctly.
|
||||
|
||||
Usage:
|
||||
python cli/test-tagger.py [image_path]
|
||||
|
||||
If no image path is provided, uses a built-in test image.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
# Add parent directory to path for imports
|
||||
script_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
sys.path.insert(0, script_dir)
|
||||
os.chdir(script_dir)
|
||||
|
||||
# Suppress installer output during import
|
||||
os.environ['SD_INSTALL_QUIET'] = '1'
|
||||
|
||||
# Initialize cmd_args properly with all argument groups
|
||||
import modules.cmd_args
|
||||
import installer
|
||||
|
||||
# Add installer args to the parser
|
||||
installer.add_args(modules.cmd_args.parser)
|
||||
|
||||
# Parse with empty args to get defaults
|
||||
modules.cmd_args.parsed, _ = modules.cmd_args.parser.parse_known_args([])
|
||||
|
||||
# Now we can safely import modules that depend on cmd_args
|
||||
|
||||
|
||||
# Default test images (in order of preference)
|
||||
DEFAULT_TEST_IMAGES = [
|
||||
'html/sdnext-robot-2k.jpg', # SD.Next robot mascot
|
||||
'venv/lib/python3.13/site-packages/gradio/test_data/lion.jpg',
|
||||
'venv/lib/python3.13/site-packages/gradio/test_data/cheetah1.jpg',
|
||||
'venv/lib/python3.13/site-packages/skimage/data/astronaut.png',
|
||||
'venv/lib/python3.13/site-packages/skimage/data/coffee.png',
|
||||
]
|
||||
|
||||
|
||||
def find_test_image():
|
||||
"""Find a suitable test image from defaults."""
|
||||
for img_path in DEFAULT_TEST_IMAGES:
|
||||
full_path = os.path.join(script_dir, img_path)
|
||||
if os.path.exists(full_path):
|
||||
return full_path
|
||||
return None
|
||||
|
||||
|
||||
def create_test_image():
|
||||
"""Create a simple test image as fallback."""
|
||||
from PIL import Image, ImageDraw
|
||||
img = Image.new('RGB', (512, 512), color=(200, 150, 100))
|
||||
draw = ImageDraw.Draw(img)
|
||||
draw.ellipse([100, 100, 400, 400], fill=(255, 200, 150), outline=(100, 50, 0))
|
||||
draw.rectangle([150, 200, 350, 350], fill=(150, 100, 200))
|
||||
return img
|
||||
|
||||
|
||||
class TaggerTest:
|
||||
"""Test harness for tagger settings."""
|
||||
|
||||
def __init__(self):
|
||||
self.results = {'passed': [], 'failed': [], 'skipped': []}
|
||||
self.test_image = None
|
||||
self.wd14_loaded = False
|
||||
self.deepbooru_loaded = False
|
||||
|
||||
def log_pass(self, msg):
|
||||
print(f" [PASS] {msg}")
|
||||
self.results['passed'].append(msg)
|
||||
|
||||
def log_fail(self, msg):
|
||||
print(f" [FAIL] {msg}")
|
||||
self.results['failed'].append(msg)
|
||||
|
||||
def log_skip(self, msg):
|
||||
print(f" [SKIP] {msg}")
|
||||
self.results['skipped'].append(msg)
|
||||
|
||||
def log_warn(self, msg):
|
||||
print(f" [WARN] {msg}")
|
||||
self.results['skipped'].append(msg)
|
||||
|
||||
def setup(self):
|
||||
"""Load test image and models."""
|
||||
from PIL import Image
|
||||
from modules import shared
|
||||
|
||||
print("=" * 70)
|
||||
print("TAGGER SETTINGS TEST SUITE")
|
||||
print("=" * 70)
|
||||
|
||||
# Get or create test image
|
||||
if len(sys.argv) > 1 and os.path.exists(sys.argv[1]):
|
||||
img_path = sys.argv[1]
|
||||
print(f"\nUsing provided image: {img_path}")
|
||||
self.test_image = Image.open(img_path).convert('RGB')
|
||||
else:
|
||||
img_path = find_test_image()
|
||||
if img_path:
|
||||
print(f"\nUsing default test image: {img_path}")
|
||||
self.test_image = Image.open(img_path).convert('RGB')
|
||||
else:
|
||||
print("\nNo test image found, creating synthetic image...")
|
||||
self.test_image = create_test_image()
|
||||
|
||||
print(f"Image size: {self.test_image.size}")
|
||||
|
||||
# Load models
|
||||
print("\nLoading models...")
|
||||
from modules.interrogate import wd14, deepbooru
|
||||
|
||||
t0 = time.time()
|
||||
self.wd14_loaded = wd14.load_model()
|
||||
print(f" WD14: {'loaded' if self.wd14_loaded else 'FAILED'} ({time.time()-t0:.1f}s)")
|
||||
|
||||
t0 = time.time()
|
||||
self.deepbooru_loaded = deepbooru.load_model()
|
||||
print(f" DeepBooru: {'loaded' if self.deepbooru_loaded else 'FAILED'} ({time.time()-t0:.1f}s)")
|
||||
|
||||
def cleanup(self):
|
||||
"""Unload models and free memory."""
|
||||
print("\n" + "=" * 70)
|
||||
print("CLEANUP")
|
||||
print("=" * 70)
|
||||
|
||||
from modules.interrogate import wd14, deepbooru
|
||||
from modules import devices
|
||||
|
||||
wd14.unload_model()
|
||||
deepbooru.unload_model()
|
||||
devices.torch_gc(force=True)
|
||||
print(" Models unloaded")
|
||||
|
||||
def print_summary(self):
|
||||
"""Print test summary."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST SUMMARY")
|
||||
print("=" * 70)
|
||||
|
||||
print(f"\n PASSED: {len(self.results['passed'])}")
|
||||
for item in self.results['passed']:
|
||||
print(f" - {item}")
|
||||
|
||||
print(f"\n FAILED: {len(self.results['failed'])}")
|
||||
for item in self.results['failed']:
|
||||
print(f" - {item}")
|
||||
|
||||
print(f"\n SKIPPED: {len(self.results['skipped'])}")
|
||||
for item in self.results['skipped']:
|
||||
print(f" - {item}")
|
||||
|
||||
total = len(self.results['passed']) + len(self.results['failed'])
|
||||
if total > 0:
|
||||
success_rate = len(self.results['passed']) / total * 100
|
||||
print(f"\n SUCCESS RATE: {success_rate:.1f}% ({len(self.results['passed'])}/{total})")
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: ONNX Providers Detection
|
||||
# =========================================================================
|
||||
def test_onnx_providers(self):
|
||||
"""Verify ONNX runtime providers are properly detected."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: ONNX Providers Detection")
|
||||
print("=" * 70)
|
||||
|
||||
from modules import devices
|
||||
|
||||
# Test 1: onnxruntime can be imported
|
||||
try:
|
||||
import onnxruntime as ort
|
||||
self.log_pass(f"onnxruntime imported: version={ort.__version__}")
|
||||
except ImportError as e:
|
||||
self.log_fail(f"onnxruntime import failed: {e}")
|
||||
return
|
||||
|
||||
# Test 2: Get available providers
|
||||
available = ort.get_available_providers()
|
||||
if available and len(available) > 0:
|
||||
self.log_pass(f"Available providers: {available}")
|
||||
else:
|
||||
self.log_fail("No ONNX providers available")
|
||||
return
|
||||
|
||||
# Test 3: devices.onnx is properly configured
|
||||
if devices.onnx is not None and len(devices.onnx) > 0:
|
||||
self.log_pass(f"devices.onnx configured: {devices.onnx}")
|
||||
else:
|
||||
self.log_fail(f"devices.onnx not configured: {devices.onnx}")
|
||||
|
||||
# Test 4: Configured providers exist in available providers
|
||||
for provider in devices.onnx:
|
||||
if provider in available:
|
||||
self.log_pass(f"Provider '{provider}' is available")
|
||||
else:
|
||||
self.log_fail(f"Provider '{provider}' configured but not available")
|
||||
|
||||
# Test 5: If WD14 loaded, check session providers
|
||||
if self.wd14_loaded:
|
||||
from modules.interrogate import wd14
|
||||
if wd14.tagger.session is not None:
|
||||
session_providers = wd14.tagger.session.get_providers()
|
||||
self.log_pass(f"WD14 session providers: {session_providers}")
|
||||
else:
|
||||
self.log_skip("WD14 session not initialized")
|
||||
|
||||
# =========================================================================
|
||||
# TEST: Memory Management (Offload/Reload/Unload)
|
||||
# =========================================================================
|
||||
def get_memory_stats(self):
|
||||
"""Get current GPU and CPU memory usage."""
|
||||
import torch
|
||||
import gc
|
||||
|
||||
stats = {}
|
||||
|
||||
# GPU memory (if CUDA available)
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
stats['gpu_allocated'] = torch.cuda.memory_allocated() / 1024 / 1024 # MB
|
||||
stats['gpu_reserved'] = torch.cuda.memory_reserved() / 1024 / 1024 # MB
|
||||
else:
|
||||
stats['gpu_allocated'] = 0
|
||||
stats['gpu_reserved'] = 0
|
||||
|
||||
# CPU/RAM memory (try psutil, fallback to basic)
|
||||
try:
|
||||
import psutil
|
||||
process = psutil.Process()
|
||||
stats['ram_used'] = process.memory_info().rss / 1024 / 1024 # MB
|
||||
except ImportError:
|
||||
stats['ram_used'] = 0
|
||||
|
||||
return stats
|
||||
|
||||
def test_memory_management(self):
|
||||
"""Test model offload to RAM, reload to GPU, and unload with memory monitoring."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: Memory Management (Offload/Reload/Unload)")
|
||||
print("=" * 70)
|
||||
|
||||
import torch
|
||||
import gc
|
||||
from modules import devices
|
||||
from modules.interrogate import wd14, deepbooru
|
||||
|
||||
# Memory leak tolerance (MB) - some variance is expected
|
||||
GPU_LEAK_TOLERANCE_MB = 50
|
||||
RAM_LEAK_TOLERANCE_MB = 200
|
||||
|
||||
# =====================================================================
|
||||
# DeepBooru: Test GPU/CPU movement with memory monitoring
|
||||
# =====================================================================
|
||||
if self.deepbooru_loaded:
|
||||
print("\n DeepBooru Memory Management:")
|
||||
|
||||
# Baseline memory before any operations
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
baseline = self.get_memory_stats()
|
||||
print(f" Baseline: GPU={baseline['gpu_allocated']:.1f}MB, RAM={baseline['ram_used']:.1f}MB")
|
||||
|
||||
# Test 1: Check initial state (should be on CPU after load)
|
||||
initial_device = deepbooru.model._device
|
||||
print(f" Initial device: {initial_device}")
|
||||
if initial_device == devices.cpu:
|
||||
self.log_pass("DeepBooru: initial state on CPU")
|
||||
else:
|
||||
self.log_pass(f"DeepBooru: initial state on {initial_device}")
|
||||
|
||||
# Test 2: Move to GPU (start)
|
||||
deepbooru.model.start()
|
||||
gpu_device = deepbooru.model._device
|
||||
after_gpu = self.get_memory_stats()
|
||||
print(f" After start(): {gpu_device} | GPU={after_gpu['gpu_allocated']:.1f}MB (+{after_gpu['gpu_allocated']-baseline['gpu_allocated']:.1f}MB)")
|
||||
if gpu_device == devices.device:
|
||||
self.log_pass(f"DeepBooru: moved to GPU ({gpu_device})")
|
||||
else:
|
||||
self.log_fail(f"DeepBooru: failed to move to GPU, got {gpu_device}")
|
||||
|
||||
# Test 3: Run inference while on GPU
|
||||
try:
|
||||
tags = deepbooru.model.tag_multi(self.test_image, max_tags=3)
|
||||
after_infer = self.get_memory_stats()
|
||||
print(f" After inference: GPU={after_infer['gpu_allocated']:.1f}MB")
|
||||
if tags:
|
||||
self.log_pass(f"DeepBooru: inference on GPU works ({tags[:30]}...)")
|
||||
else:
|
||||
self.log_fail("DeepBooru: inference on GPU returned empty")
|
||||
except Exception as e:
|
||||
self.log_fail(f"DeepBooru: inference on GPU failed: {e}")
|
||||
|
||||
# Test 4: Offload to CPU (stop)
|
||||
deepbooru.model.stop()
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
after_offload = self.get_memory_stats()
|
||||
cpu_device = deepbooru.model._device
|
||||
print(f" After stop(): {cpu_device} | GPU={after_offload['gpu_allocated']:.1f}MB, RAM={after_offload['ram_used']:.1f}MB")
|
||||
if cpu_device == devices.cpu:
|
||||
self.log_pass("DeepBooru: offloaded to CPU")
|
||||
else:
|
||||
self.log_fail(f"DeepBooru: failed to offload, still on {cpu_device}")
|
||||
|
||||
# Check GPU memory returned to near baseline after offload
|
||||
gpu_diff = after_offload['gpu_allocated'] - baseline['gpu_allocated']
|
||||
if gpu_diff <= GPU_LEAK_TOLERANCE_MB:
|
||||
self.log_pass(f"DeepBooru: GPU memory cleared after offload (diff={gpu_diff:.1f}MB)")
|
||||
else:
|
||||
self.log_fail(f"DeepBooru: GPU memory leak after offload (diff={gpu_diff:.1f}MB > {GPU_LEAK_TOLERANCE_MB}MB)")
|
||||
|
||||
# Test 5: Full cycle - reload and run again
|
||||
deepbooru.model.start()
|
||||
try:
|
||||
tags = deepbooru.model.tag_multi(self.test_image, max_tags=3)
|
||||
if tags:
|
||||
self.log_pass("DeepBooru: reload cycle works")
|
||||
else:
|
||||
self.log_fail("DeepBooru: reload cycle returned empty")
|
||||
except Exception as e:
|
||||
self.log_fail(f"DeepBooru: reload cycle failed: {e}")
|
||||
deepbooru.model.stop()
|
||||
|
||||
# Test 6: Full unload with memory check
|
||||
deepbooru.unload_model()
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
after_unload = self.get_memory_stats()
|
||||
print(f" After unload: GPU={after_unload['gpu_allocated']:.1f}MB, RAM={after_unload['ram_used']:.1f}MB")
|
||||
|
||||
if deepbooru.model.model is None:
|
||||
self.log_pass("DeepBooru: unload successful")
|
||||
else:
|
||||
self.log_fail("DeepBooru: unload failed, model still exists")
|
||||
|
||||
# Check for memory leaks after full unload
|
||||
gpu_leak = after_unload['gpu_allocated'] - baseline['gpu_allocated']
|
||||
ram_leak = after_unload['ram_used'] - baseline['ram_used']
|
||||
if gpu_leak <= GPU_LEAK_TOLERANCE_MB:
|
||||
self.log_pass(f"DeepBooru: no GPU memory leak after unload (diff={gpu_leak:.1f}MB)")
|
||||
else:
|
||||
self.log_fail(f"DeepBooru: GPU memory leak detected (diff={gpu_leak:.1f}MB > {GPU_LEAK_TOLERANCE_MB}MB)")
|
||||
|
||||
if ram_leak <= RAM_LEAK_TOLERANCE_MB:
|
||||
self.log_pass(f"DeepBooru: no RAM leak after unload (diff={ram_leak:.1f}MB)")
|
||||
else:
|
||||
self.log_warn(f"DeepBooru: RAM increased after unload (diff={ram_leak:.1f}MB) - may be caching")
|
||||
|
||||
# Reload for remaining tests
|
||||
deepbooru.load_model()
|
||||
|
||||
# =====================================================================
|
||||
# WD14: Test session lifecycle with memory monitoring
|
||||
# =====================================================================
|
||||
if self.wd14_loaded:
|
||||
print("\n WD14 Memory Management:")
|
||||
|
||||
# Baseline memory
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
baseline = self.get_memory_stats()
|
||||
print(f" Baseline: GPU={baseline['gpu_allocated']:.1f}MB, RAM={baseline['ram_used']:.1f}MB")
|
||||
|
||||
# Test 1: Session exists
|
||||
if wd14.tagger.session is not None:
|
||||
self.log_pass("WD14: session loaded")
|
||||
else:
|
||||
self.log_fail("WD14: session not loaded")
|
||||
return
|
||||
|
||||
# Test 2: Get current providers
|
||||
providers = wd14.tagger.session.get_providers()
|
||||
print(f" Active providers: {providers}")
|
||||
self.log_pass(f"WD14: using providers {providers}")
|
||||
|
||||
# Test 3: Run inference
|
||||
try:
|
||||
tags = wd14.tagger.predict(self.test_image, max_tags=3)
|
||||
after_infer = self.get_memory_stats()
|
||||
print(f" After inference: GPU={after_infer['gpu_allocated']:.1f}MB, RAM={after_infer['ram_used']:.1f}MB")
|
||||
if tags:
|
||||
self.log_pass(f"WD14: inference works ({tags[:30]}...)")
|
||||
else:
|
||||
self.log_fail("WD14: inference returned empty")
|
||||
except Exception as e:
|
||||
self.log_fail(f"WD14: inference failed: {e}")
|
||||
|
||||
# Test 4: Unload session with memory check
|
||||
model_name = wd14.tagger.model_name
|
||||
wd14.unload_model()
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
after_unload = self.get_memory_stats()
|
||||
print(f" After unload: GPU={after_unload['gpu_allocated']:.1f}MB, RAM={after_unload['ram_used']:.1f}MB")
|
||||
|
||||
if wd14.tagger.session is None:
|
||||
self.log_pass("WD14: unload successful")
|
||||
else:
|
||||
self.log_fail("WD14: unload failed, session still exists")
|
||||
|
||||
# Check for memory leaks after unload
|
||||
gpu_leak = after_unload['gpu_allocated'] - baseline['gpu_allocated']
|
||||
ram_leak = after_unload['ram_used'] - baseline['ram_used']
|
||||
if gpu_leak <= GPU_LEAK_TOLERANCE_MB:
|
||||
self.log_pass(f"WD14: no GPU memory leak after unload (diff={gpu_leak:.1f}MB)")
|
||||
else:
|
||||
self.log_fail(f"WD14: GPU memory leak detected (diff={gpu_leak:.1f}MB > {GPU_LEAK_TOLERANCE_MB}MB)")
|
||||
|
||||
if ram_leak <= RAM_LEAK_TOLERANCE_MB:
|
||||
self.log_pass(f"WD14: no RAM leak after unload (diff={ram_leak:.1f}MB)")
|
||||
else:
|
||||
self.log_warn(f"WD14: RAM increased after unload (diff={ram_leak:.1f}MB) - may be caching")
|
||||
|
||||
# Test 5: Reload session
|
||||
wd14.load_model(model_name)
|
||||
after_reload = self.get_memory_stats()
|
||||
print(f" After reload: GPU={after_reload['gpu_allocated']:.1f}MB, RAM={after_reload['ram_used']:.1f}MB")
|
||||
if wd14.tagger.session is not None:
|
||||
self.log_pass("WD14: reload successful")
|
||||
else:
|
||||
self.log_fail("WD14: reload failed")
|
||||
|
||||
# Test 6: Inference after reload
|
||||
try:
|
||||
tags = wd14.tagger.predict(self.test_image, max_tags=3)
|
||||
if tags:
|
||||
self.log_pass("WD14: inference after reload works")
|
||||
else:
|
||||
self.log_fail("WD14: inference after reload returned empty")
|
||||
except Exception as e:
|
||||
self.log_fail(f"WD14: inference after reload failed: {e}")
|
||||
|
||||
# Final memory check after full cycle
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
final = self.get_memory_stats()
|
||||
print(f" Final (after full cycle): GPU={final['gpu_allocated']:.1f}MB, RAM={final['ram_used']:.1f}MB")
|
||||
|
||||
# =========================================================================
|
||||
# TEST: Settings Existence
|
||||
# =========================================================================
|
||||
def test_settings_exist(self):
|
||||
"""Verify all tagger settings exist in shared.opts."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: Settings Existence")
|
||||
print("=" * 70)
|
||||
|
||||
from modules import shared
|
||||
|
||||
settings = [
|
||||
('tagger_threshold', float),
|
||||
('tagger_include_rating', bool),
|
||||
('tagger_max_tags', int),
|
||||
('tagger_sort_alpha', bool),
|
||||
('tagger_use_spaces', bool),
|
||||
('tagger_escape_brackets', bool),
|
||||
('tagger_exclude_tags', str),
|
||||
('tagger_show_scores', bool),
|
||||
('wd14_model', str),
|
||||
('wd14_character_threshold', float),
|
||||
('interrogate_offload', bool),
|
||||
]
|
||||
|
||||
for setting, _expected_type in settings:
|
||||
if hasattr(shared.opts, setting):
|
||||
value = getattr(shared.opts, setting)
|
||||
self.log_pass(f"{setting} = {value!r}")
|
||||
else:
|
||||
self.log_fail(f"{setting} - NOT FOUND")
|
||||
|
||||
# =========================================================================
|
||||
# TEST: Parameter Effect - Tests a single parameter on both taggers
|
||||
# =========================================================================
|
||||
def test_parameter(self, param_name, test_func, wd14_supported=True, deepbooru_supported=True):
|
||||
"""Test a parameter on both WD14 and DeepBooru."""
|
||||
print(f"\n Testing: {param_name}")
|
||||
|
||||
if wd14_supported and self.wd14_loaded:
|
||||
try:
|
||||
result = test_func('wd14')
|
||||
if result is True:
|
||||
self.log_pass(f"WD14: {param_name}")
|
||||
elif result is False:
|
||||
self.log_fail(f"WD14: {param_name}")
|
||||
else:
|
||||
self.log_skip(f"WD14: {param_name} - {result}")
|
||||
except Exception as e:
|
||||
self.log_fail(f"WD14: {param_name} - {e}")
|
||||
elif wd14_supported:
|
||||
self.log_skip(f"WD14: {param_name} - model not loaded")
|
||||
|
||||
if deepbooru_supported and self.deepbooru_loaded:
|
||||
try:
|
||||
result = test_func('deepbooru')
|
||||
if result is True:
|
||||
self.log_pass(f"DeepBooru: {param_name}")
|
||||
elif result is False:
|
||||
self.log_fail(f"DeepBooru: {param_name}")
|
||||
else:
|
||||
self.log_skip(f"DeepBooru: {param_name} - {result}")
|
||||
except Exception as e:
|
||||
self.log_fail(f"DeepBooru: {param_name} - {e}")
|
||||
elif deepbooru_supported:
|
||||
self.log_skip(f"DeepBooru: {param_name} - model not loaded")
|
||||
|
||||
def tag(self, tagger, **kwargs):
|
||||
"""Helper to call the appropriate tagger."""
|
||||
if tagger == 'wd14':
|
||||
from modules.interrogate import wd14
|
||||
return wd14.tagger.predict(self.test_image, **kwargs)
|
||||
else:
|
||||
from modules.interrogate import deepbooru
|
||||
return deepbooru.model.tag(self.test_image, **kwargs)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: general_threshold
|
||||
# =========================================================================
|
||||
def test_threshold(self):
|
||||
"""Test that threshold affects tag count."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: general_threshold effect")
|
||||
print("=" * 70)
|
||||
|
||||
def check_threshold(tagger):
|
||||
tags_high = self.tag(tagger, general_threshold=0.9)
|
||||
tags_low = self.tag(tagger, general_threshold=0.1)
|
||||
|
||||
count_high = len(tags_high.split(', ')) if tags_high else 0
|
||||
count_low = len(tags_low.split(', ')) if tags_low else 0
|
||||
|
||||
print(f" {tagger}: threshold=0.9 -> {count_high} tags, threshold=0.1 -> {count_low} tags")
|
||||
|
||||
if count_low > count_high:
|
||||
return True
|
||||
elif count_low == count_high == 0:
|
||||
return "no tags returned"
|
||||
else:
|
||||
return "threshold effect unclear"
|
||||
|
||||
self.test_parameter('general_threshold', check_threshold)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: max_tags
|
||||
# =========================================================================
|
||||
def test_max_tags(self):
|
||||
"""Test that max_tags limits output."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: max_tags effect")
|
||||
print("=" * 70)
|
||||
|
||||
def check_max_tags(tagger):
|
||||
tags_5 = self.tag(tagger, general_threshold=0.1, max_tags=5)
|
||||
tags_50 = self.tag(tagger, general_threshold=0.1, max_tags=50)
|
||||
|
||||
count_5 = len(tags_5.split(', ')) if tags_5 else 0
|
||||
count_50 = len(tags_50.split(', ')) if tags_50 else 0
|
||||
|
||||
print(f" {tagger}: max_tags=5 -> {count_5} tags, max_tags=50 -> {count_50} tags")
|
||||
|
||||
return count_5 <= 5
|
||||
|
||||
self.test_parameter('max_tags', check_max_tags)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: use_spaces
|
||||
# =========================================================================
|
||||
def test_use_spaces(self):
|
||||
"""Test that use_spaces converts underscores to spaces."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: use_spaces effect")
|
||||
print("=" * 70)
|
||||
|
||||
def check_use_spaces(tagger):
|
||||
tags_under = self.tag(tagger, use_spaces=False, max_tags=10)
|
||||
tags_space = self.tag(tagger, use_spaces=True, max_tags=10)
|
||||
|
||||
print(f" {tagger} use_spaces=False: {tags_under[:50]}...")
|
||||
print(f" {tagger} use_spaces=True: {tags_space[:50]}...")
|
||||
|
||||
# Check if underscores are converted to spaces
|
||||
has_underscore_before = '_' in tags_under
|
||||
has_underscore_after = '_' in tags_space.replace(', ', ',') # ignore comma-space
|
||||
|
||||
# If there were underscores before but not after, it worked
|
||||
if has_underscore_before and not has_underscore_after:
|
||||
return True
|
||||
# If there were never underscores, inconclusive
|
||||
elif not has_underscore_before:
|
||||
return "no underscores in tags to convert"
|
||||
else:
|
||||
return False
|
||||
|
||||
self.test_parameter('use_spaces', check_use_spaces)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: escape_brackets
|
||||
# =========================================================================
|
||||
def test_escape_brackets(self):
|
||||
"""Test that escape_brackets escapes special characters."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: escape_brackets effect")
|
||||
print("=" * 70)
|
||||
|
||||
def check_escape_brackets(tagger):
|
||||
tags_escaped = self.tag(tagger, escape_brackets=True, max_tags=30, general_threshold=0.1)
|
||||
tags_raw = self.tag(tagger, escape_brackets=False, max_tags=30, general_threshold=0.1)
|
||||
|
||||
print(f" {tagger} escape=True: {tags_escaped[:60]}...")
|
||||
print(f" {tagger} escape=False: {tags_raw[:60]}...")
|
||||
|
||||
# Check for escaped brackets (\\( or \\))
|
||||
has_escaped = '\\(' in tags_escaped or '\\)' in tags_escaped
|
||||
has_unescaped = '(' in tags_raw.replace('\\(', '') or ')' in tags_raw.replace('\\)', '')
|
||||
|
||||
if has_escaped:
|
||||
return True
|
||||
elif has_unescaped:
|
||||
# Has brackets but not escaped - fail
|
||||
return False
|
||||
else:
|
||||
return "no brackets in tags to escape"
|
||||
|
||||
self.test_parameter('escape_brackets', check_escape_brackets)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: sort_alpha
|
||||
# =========================================================================
|
||||
def test_sort_alpha(self):
|
||||
"""Test that sort_alpha sorts tags alphabetically."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: sort_alpha effect")
|
||||
print("=" * 70)
|
||||
|
||||
def check_sort_alpha(tagger):
|
||||
tags_conf = self.tag(tagger, sort_alpha=False, max_tags=20, general_threshold=0.1)
|
||||
tags_alpha = self.tag(tagger, sort_alpha=True, max_tags=20, general_threshold=0.1)
|
||||
|
||||
list_conf = [t.strip() for t in tags_conf.split(',')]
|
||||
list_alpha = [t.strip() for t in tags_alpha.split(',')]
|
||||
|
||||
print(f" {tagger} by_confidence: {', '.join(list_conf[:5])}...")
|
||||
print(f" {tagger} alphabetical: {', '.join(list_alpha[:5])}...")
|
||||
|
||||
is_sorted = list_alpha == sorted(list_alpha)
|
||||
return is_sorted
|
||||
|
||||
self.test_parameter('sort_alpha', check_sort_alpha)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: exclude_tags
|
||||
# =========================================================================
|
||||
def test_exclude_tags(self):
|
||||
"""Test that exclude_tags removes specified tags."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: exclude_tags effect")
|
||||
print("=" * 70)
|
||||
|
||||
def check_exclude_tags(tagger):
|
||||
tags_all = self.tag(tagger, max_tags=50, general_threshold=0.1, exclude_tags='')
|
||||
tag_list = [t.strip().replace(' ', '_') for t in tags_all.split(',')]
|
||||
|
||||
if len(tag_list) < 2:
|
||||
return "not enough tags to test"
|
||||
|
||||
# Exclude the first tag
|
||||
tag_to_exclude = tag_list[0]
|
||||
tags_filtered = self.tag(tagger, max_tags=50, general_threshold=0.1, exclude_tags=tag_to_exclude)
|
||||
|
||||
print(f" {tagger} without exclusion: {tags_all[:50]}...")
|
||||
print(f" {tagger} excluding '{tag_to_exclude}': {tags_filtered[:50]}...")
|
||||
|
||||
# Check if the exact tag was removed by parsing the filtered list
|
||||
filtered_list = [t.strip().replace(' ', '_') for t in tags_filtered.split(',')]
|
||||
# Also check space variant
|
||||
tag_space_variant = tag_to_exclude.replace('_', ' ')
|
||||
tag_present = tag_to_exclude in filtered_list or tag_space_variant in [t.strip() for t in tags_filtered.split(',')]
|
||||
return not tag_present
|
||||
|
||||
self.test_parameter('exclude_tags', check_exclude_tags)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: tagger_show_scores (via shared.opts)
|
||||
# =========================================================================
|
||||
def test_show_scores(self):
|
||||
"""Test that tagger_show_scores adds confidence scores."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: tagger_show_scores effect")
|
||||
print("=" * 70)
|
||||
|
||||
from modules import shared
|
||||
|
||||
def check_show_scores(tagger):
|
||||
original = shared.opts.tagger_show_scores
|
||||
|
||||
shared.opts.tagger_show_scores = False
|
||||
tags_no_scores = self.tag(tagger, max_tags=5)
|
||||
|
||||
shared.opts.tagger_show_scores = True
|
||||
tags_with_scores = self.tag(tagger, max_tags=5)
|
||||
|
||||
shared.opts.tagger_show_scores = original
|
||||
|
||||
print(f" {tagger} show_scores=False: {tags_no_scores[:50]}...")
|
||||
print(f" {tagger} show_scores=True: {tags_with_scores[:50]}...")
|
||||
|
||||
has_scores = ':' in tags_with_scores and '(' in tags_with_scores
|
||||
no_scores = ':' not in tags_no_scores
|
||||
|
||||
return has_scores and no_scores
|
||||
|
||||
self.test_parameter('tagger_show_scores', check_show_scores)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: include_rating
|
||||
# =========================================================================
|
||||
def test_include_rating(self):
|
||||
"""Test that include_rating includes/excludes rating tags."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: include_rating effect")
|
||||
print("=" * 70)
|
||||
|
||||
def check_include_rating(tagger):
|
||||
tags_no_rating = self.tag(tagger, include_rating=False, max_tags=100, general_threshold=0.01)
|
||||
tags_with_rating = self.tag(tagger, include_rating=True, max_tags=100, general_threshold=0.01)
|
||||
|
||||
print(f" {tagger} include_rating=False: {tags_no_rating[:60]}...")
|
||||
print(f" {tagger} include_rating=True: {tags_with_rating[:60]}...")
|
||||
|
||||
# Rating tags typically start with "rating:" or are like "safe", "questionable", "explicit"
|
||||
rating_keywords = ['rating:', 'safe', 'questionable', 'explicit', 'general', 'sensitive']
|
||||
|
||||
has_rating_before = any(kw in tags_no_rating.lower() for kw in rating_keywords)
|
||||
has_rating_after = any(kw in tags_with_rating.lower() for kw in rating_keywords)
|
||||
|
||||
if has_rating_after and not has_rating_before:
|
||||
return True
|
||||
elif has_rating_after and has_rating_before:
|
||||
return "rating tags appear in both (may need very low threshold)"
|
||||
elif not has_rating_after:
|
||||
return "no rating tags detected"
|
||||
else:
|
||||
return False
|
||||
|
||||
self.test_parameter('include_rating', check_include_rating)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: character_threshold (WD14 only)
|
||||
# =========================================================================
|
||||
def test_character_threshold(self):
|
||||
"""Test that character_threshold affects character tag count (WD14 only)."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: character_threshold effect (WD14 only)")
|
||||
print("=" * 70)
|
||||
|
||||
def check_character_threshold(tagger):
|
||||
if tagger != 'wd14':
|
||||
return "not supported"
|
||||
|
||||
# Character threshold only affects character tags
|
||||
# We need an image with character tags to properly test this
|
||||
tags_high = self.tag(tagger, character_threshold=0.99, general_threshold=0.5)
|
||||
tags_low = self.tag(tagger, character_threshold=0.1, general_threshold=0.5)
|
||||
|
||||
print(f" {tagger} char_threshold=0.99: {tags_high[:50]}...")
|
||||
print(f" {tagger} char_threshold=0.10: {tags_low[:50]}...")
|
||||
|
||||
# If thresholds are different, the setting is at least being applied
|
||||
# Hard to verify without an image with known character tags
|
||||
return True # Setting exists and is applied (verified by code inspection)
|
||||
|
||||
self.test_parameter('character_threshold', check_character_threshold, deepbooru_supported=False)
|
||||
|
||||
# =========================================================================
|
||||
# TEST: Unified Interface
|
||||
# =========================================================================
|
||||
def test_unified_interface(self):
|
||||
"""Test that the unified tagger interface works for both backends."""
|
||||
print("\n" + "=" * 70)
|
||||
print("TEST: Unified tagger.tag() interface")
|
||||
print("=" * 70)
|
||||
|
||||
from modules.interrogate import tagger
|
||||
|
||||
# Test WD14 through unified interface
|
||||
if self.wd14_loaded:
|
||||
try:
|
||||
models = tagger.get_models()
|
||||
wd14_model = next((m for m in models if m != 'DeepBooru'), None)
|
||||
if wd14_model:
|
||||
tags = tagger.tag(self.test_image, model_name=wd14_model, max_tags=5)
|
||||
print(f" WD14 ({wd14_model}): {tags[:50]}...")
|
||||
self.log_pass("Unified interface: WD14")
|
||||
except Exception as e:
|
||||
self.log_fail(f"Unified interface: WD14 - {e}")
|
||||
|
||||
# Test DeepBooru through unified interface
|
||||
if self.deepbooru_loaded:
|
||||
try:
|
||||
tags = tagger.tag(self.test_image, model_name='DeepBooru', max_tags=5)
|
||||
print(f" DeepBooru: {tags[:50]}...")
|
||||
self.log_pass("Unified interface: DeepBooru")
|
||||
except Exception as e:
|
||||
self.log_fail(f"Unified interface: DeepBooru - {e}")
|
||||
|
||||
def run_all_tests(self):
|
||||
"""Run all tests."""
|
||||
self.setup()
|
||||
|
||||
self.test_onnx_providers()
|
||||
self.test_memory_management()
|
||||
self.test_settings_exist()
|
||||
self.test_threshold()
|
||||
self.test_max_tags()
|
||||
self.test_use_spaces()
|
||||
self.test_escape_brackets()
|
||||
self.test_sort_alpha()
|
||||
self.test_exclude_tags()
|
||||
self.test_show_scores()
|
||||
self.test_include_rating()
|
||||
self.test_character_threshold()
|
||||
self.test_unified_interface()
|
||||
|
||||
self.cleanup()
|
||||
self.print_summary()
|
||||
|
||||
return len(self.results['failed']) == 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test = TaggerTest()
|
||||
success = test.run_all_tests()
|
||||
sys.exit(0 if success else 1)
|
||||
+1
-1
@@ -90,7 +90,7 @@
|
||||
{"id":"","label":"Embedding","localized":"","reload":"","hint":"Textual inversion embedding is a trained embedded information about the subject"},
|
||||
{"id":"","label":"Hypernetwork","localized":"","reload":"","hint":"Small trained neural network that modifies behavior of the loaded model"},
|
||||
{"id":"","label":"VLM Caption","localized":"","reload":"","hint":"Analyze image using vision langugage model"},
|
||||
{"id":"","label":"CLiP Interrogate","localized":"","reload":"","hint":"Analyze image using CLiP model"},
|
||||
{"id":"","label":"OpenCLiP","localized":"","reload":"","hint":"Analyze image using CLiP model via OpenCLiP"},
|
||||
{"id":"","label":"VAE","localized":"","reload":"","hint":"Variational Auto Encoder: model used to run image decode at the end of generate"},
|
||||
{"id":"","label":"History","localized":"","reload":"","hint":"List of previous generations that can be further reprocessed"},
|
||||
{"id":"","label":"UI disable variable aspect ratio","localized":"","reload":"","hint":"When disabled, all thumbnails appear as squared images"},
|
||||
|
||||
@@ -46,14 +46,55 @@ class DeepDanbooru:
|
||||
self._device = devices.cpu
|
||||
devices.torch_gc()
|
||||
|
||||
def tag(self, pil_image):
|
||||
def tag(self, pil_image, **kwargs):
|
||||
self.start()
|
||||
res = self.tag_multi(pil_image)
|
||||
res = self.tag_multi(pil_image, **kwargs)
|
||||
self.stop()
|
||||
|
||||
return res
|
||||
|
||||
def tag_multi(self, pil_image, force_disable_ranks=False):
|
||||
def tag_multi(
|
||||
self,
|
||||
pil_image,
|
||||
general_threshold: float = None,
|
||||
include_rating: bool = None,
|
||||
exclude_tags: str = None,
|
||||
max_tags: int = None,
|
||||
sort_alpha: bool = None,
|
||||
use_spaces: bool = None,
|
||||
escape_brackets: bool = None,
|
||||
):
|
||||
"""Run inference and return formatted tag string.
|
||||
|
||||
Args:
|
||||
pil_image: PIL Image to tag
|
||||
general_threshold: Threshold for tag scores (0-1)
|
||||
include_rating: Whether to include rating tags
|
||||
exclude_tags: Comma-separated tags to exclude
|
||||
max_tags: Maximum number of tags to return
|
||||
sort_alpha: Sort tags alphabetically vs by confidence
|
||||
use_spaces: Use spaces instead of underscores
|
||||
escape_brackets: Escape parentheses/brackets in tags
|
||||
|
||||
Returns:
|
||||
Formatted tag string
|
||||
"""
|
||||
# Use settings defaults if not specified
|
||||
if general_threshold is None:
|
||||
general_threshold = shared.opts.tagger_threshold
|
||||
if include_rating is None:
|
||||
include_rating = shared.opts.tagger_include_rating
|
||||
if exclude_tags is None:
|
||||
exclude_tags = shared.opts.tagger_exclude_tags
|
||||
if max_tags is None:
|
||||
max_tags = shared.opts.tagger_max_tags
|
||||
if sort_alpha is None:
|
||||
sort_alpha = shared.opts.tagger_sort_alpha
|
||||
if use_spaces is None:
|
||||
use_spaces = shared.opts.tagger_use_spaces
|
||||
if escape_brackets is None:
|
||||
escape_brackets = shared.opts.tagger_escape_brackets
|
||||
|
||||
if isinstance(pil_image, list):
|
||||
pil_image = pil_image[0] if len(pil_image) > 0 else None
|
||||
if isinstance(pil_image, dict) and 'name' in pil_image:
|
||||
@@ -67,29 +108,29 @@ class DeepDanbooru:
|
||||
y = self.model(x)[0].detach().float().cpu().numpy()
|
||||
probability_dict = {}
|
||||
for tag, probability in zip(self.model.tags, y):
|
||||
if probability < shared.opts.deepbooru_score_threshold:
|
||||
if probability < general_threshold:
|
||||
continue
|
||||
if tag.startswith("rating:"):
|
||||
if tag.startswith("rating:") and not include_rating:
|
||||
continue
|
||||
probability_dict[tag] = probability
|
||||
if shared.opts.tagger_sort_alpha:
|
||||
if sort_alpha:
|
||||
tags = sorted(probability_dict)
|
||||
else:
|
||||
tags = [tag for tag, _ in sorted(probability_dict.items(), key=lambda x: -x[1])]
|
||||
res = []
|
||||
filtertags = {x.strip().replace(' ', '_') for x in shared.opts.tagger_exclude_tags.split(",")}
|
||||
filtertags = {x.strip().replace(' ', '_') for x in exclude_tags.split(",")}
|
||||
for tag in [x for x in tags if x not in filtertags]:
|
||||
probability = probability_dict[tag]
|
||||
tag_outformat = tag
|
||||
if shared.opts.tagger_use_spaces:
|
||||
if use_spaces:
|
||||
tag_outformat = tag_outformat.replace('_', ' ')
|
||||
if shared.opts.tagger_escape:
|
||||
if escape_brackets:
|
||||
tag_outformat = re.sub(re_special, r'\\\1', tag_outformat)
|
||||
if shared.opts.interrogate_score and not force_disable_ranks:
|
||||
if shared.opts.tagger_show_scores:
|
||||
tag_outformat = f"({tag_outformat}:{probability:.2f})"
|
||||
res.append(tag_outformat)
|
||||
if len(res) > shared.opts.tagger_max_tags:
|
||||
res = res[:shared.opts.tagger_max_tags]
|
||||
if max_tags > 0 and len(res) > max_tags:
|
||||
res = res[:max_tags]
|
||||
return ", ".join(res)
|
||||
|
||||
|
||||
@@ -125,7 +166,8 @@ def tag(image, **kwargs) -> str:
|
||||
|
||||
Args:
|
||||
image: PIL Image to tag
|
||||
**kwargs: Additional arguments (for interface compatibility)
|
||||
**kwargs: Tagger parameters (general_threshold, include_rating, exclude_tags,
|
||||
max_tags, sort_alpha, use_spaces, escape_brackets)
|
||||
|
||||
Returns:
|
||||
Formatted tag string
|
||||
@@ -136,7 +178,7 @@ def tag(image, **kwargs) -> str:
|
||||
shared.log.info(f'DeepBooru: image_size={image.size if image else None}')
|
||||
|
||||
try:
|
||||
result = model.tag(image)
|
||||
result = model.tag(image, **kwargs)
|
||||
shared.log.debug(f'DeepBooru: complete time={time.time()-t0:.2f}s tags={len(result.split(", ")) if result else 0}')
|
||||
except Exception as e:
|
||||
result = f"Exception {type(e)}"
|
||||
@@ -254,7 +296,7 @@ def batch(
|
||||
break
|
||||
|
||||
image = Image.open(img_path)
|
||||
tags_str = model.tag_multi(image)
|
||||
tags_str = model.tag_multi(image, **kwargs)
|
||||
|
||||
if save_output:
|
||||
txt_path = img_path.with_suffix('.txt')
|
||||
|
||||
@@ -12,7 +12,7 @@ def interrogate(image):
|
||||
shared.log.error('Interrogate: no image provided')
|
||||
return ''
|
||||
t0 = time.time()
|
||||
if shared.opts.interrogate_default_type == 'CLiP':
|
||||
if shared.opts.interrogate_default_type == 'OpenCLiP':
|
||||
shared.log.info(f'Interrogate: type={shared.opts.interrogate_default_type} clip="{shared.opts.interrogate_clip_model}" blip="{shared.opts.interrogate_blip_model}" mode="{shared.opts.interrogate_clip_mode}"')
|
||||
from modules.interrogate import openclip
|
||||
openclip.load_interrogator(clip_model=shared.opts.interrogate_clip_model, blip_model=shared.opts.interrogate_blip_model)
|
||||
@@ -26,14 +26,14 @@ def interrogate(image):
|
||||
prompt = tagger.tag(
|
||||
image=image,
|
||||
model_name=shared.opts.wd14_model,
|
||||
general_threshold=shared.opts.wd14_general_threshold,
|
||||
general_threshold=shared.opts.tagger_threshold,
|
||||
character_threshold=shared.opts.wd14_character_threshold,
|
||||
include_rating=shared.opts.wd14_include_rating,
|
||||
include_rating=shared.opts.tagger_include_rating,
|
||||
exclude_tags=shared.opts.tagger_exclude_tags,
|
||||
max_tags=shared.opts.tagger_max_tags,
|
||||
sort_alpha=shared.opts.tagger_sort_alpha,
|
||||
use_spaces=shared.opts.tagger_use_spaces,
|
||||
escape_brackets=shared.opts.tagger_escape,
|
||||
escape_brackets=shared.opts.tagger_escape_brackets,
|
||||
)
|
||||
shared.log.debug(f'Interrogate: time={time.time()-t0:.2f} answer="{prompt}"')
|
||||
return prompt
|
||||
|
||||
@@ -93,7 +93,7 @@ class WD14Tagger:
|
||||
|
||||
debug_log(f'WD14 load: onnxruntime version={ort.__version__}')
|
||||
|
||||
self.session = ort.InferenceSession(model_file, providers=['CPUExecutionProvider'])
|
||||
self.session = ort.InferenceSession(model_file, providers=devices.onnx)
|
||||
self.model_name = model_name
|
||||
|
||||
# Get actual providers used
|
||||
@@ -224,11 +224,11 @@ class WD14Tagger:
|
||||
|
||||
# Use settings defaults if not specified
|
||||
if general_threshold is None:
|
||||
general_threshold = shared.opts.wd14_general_threshold
|
||||
general_threshold = shared.opts.tagger_threshold
|
||||
if character_threshold is None:
|
||||
character_threshold = shared.opts.wd14_character_threshold
|
||||
if include_rating is None:
|
||||
include_rating = shared.opts.wd14_include_rating
|
||||
include_rating = shared.opts.tagger_include_rating
|
||||
if exclude_tags is None:
|
||||
exclude_tags = shared.opts.tagger_exclude_tags
|
||||
if max_tags is None:
|
||||
@@ -238,7 +238,7 @@ class WD14Tagger:
|
||||
if use_spaces is None:
|
||||
use_spaces = shared.opts.tagger_use_spaces
|
||||
if escape_brackets is None:
|
||||
escape_brackets = shared.opts.tagger_escape
|
||||
escape_brackets = shared.opts.tagger_escape_brackets
|
||||
|
||||
debug_log(f'WD14 predict: general_threshold={general_threshold} character_threshold={character_threshold} max_tags={max_tags} include_rating={include_rating} sort_alpha={sort_alpha}')
|
||||
|
||||
@@ -326,11 +326,11 @@ class WD14Tagger:
|
||||
formatted_tag = formatted_tag.replace('_', ' ')
|
||||
if escape_brackets:
|
||||
formatted_tag = re.sub(re_special, r'\\\1', formatted_tag)
|
||||
if shared.opts.interrogate_score:
|
||||
if shared.opts.tagger_show_scores:
|
||||
formatted_tag = f"({formatted_tag}:{tag_probs[tag_name]:.2f})"
|
||||
result.append(formatted_tag)
|
||||
|
||||
output = ', '.join(result)
|
||||
output = ", ".join(result)
|
||||
total_time = time.time() - t0
|
||||
debug_log(f'WD14 predict: complete tags={len(result)} time={total_time:.2f}s result="{output[:100]}..."' if len(output) > 100 else f'WD14 predict: complete tags={len(result)} time={total_time:.2f}s result="{output}"')
|
||||
|
||||
@@ -387,6 +387,9 @@ def tag(image: Image.Image, model_name: str = None, **kwargs) -> str:
|
||||
tagger.load(model_name)
|
||||
result = tagger.predict(image, **kwargs)
|
||||
shared.log.debug(f'WD14: complete time={time.time()-t0:.2f}s tags={len(result.split(", ")) if result else 0}')
|
||||
# Offload model if setting enabled
|
||||
if shared.opts.interrogate_offload:
|
||||
tagger.unload()
|
||||
except Exception as e:
|
||||
result = f"Exception {type(e)}"
|
||||
shared.log.error(f'WD14: {e}')
|
||||
|
||||
+36
-52
@@ -207,6 +207,7 @@ options_templates.update(options_section(('offload', "Model Offloading"), {
|
||||
"offload_sep": OptionInfo("<h2>Model Offloading</h2>", "", gr.HTML),
|
||||
"diffusers_offload_mode": OptionInfo(startup_offload_mode, "Model offload mode", gr.Radio, {"choices": ['none', 'balanced', 'group', 'model', 'sequential']}),
|
||||
"diffusers_offload_nonblocking": OptionInfo(False, "Non-blocking move operations"),
|
||||
"interrogate_offload": OptionInfo(True, "Offload caption models"),
|
||||
"offload_balanced_sep": OptionInfo("<h2>Balanced Offload</h2>", "", gr.HTML),
|
||||
"diffusers_offload_pre": OptionInfo(True, "Offload during pre-forward"),
|
||||
"diffusers_offload_streams": OptionInfo(False, "Offload using streams"),
|
||||
@@ -672,58 +673,6 @@ options_templates.update(options_section(('postprocessing', "Postprocessing"), {
|
||||
"upscaler_tile_overlap": OptionInfo(8, "Upscaler tile overlap", gr.Slider, {"minimum": 0, "maximum": 64, "step": 1}),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section(('interrogate', "Interrogate"), {
|
||||
"interrogate_default_type": OptionInfo("VLM", "Default caption type", gr.Radio, {"choices": ["VLM", "CLiP", "Tagger"]}),
|
||||
"interrogate_offload": OptionInfo(True, "Offload models "),
|
||||
"interrogate_score": OptionInfo(False, "Include scores in results when available", gr.Checkbox, {"visible": False}),
|
||||
|
||||
# OpenCLiP settings (hidden - controlled via Caption Tab)
|
||||
"interrogate_clip_sep": OptionInfo("<h2>OpenCLiP</h2>", "", gr.HTML, {"visible": False}),
|
||||
"interrogate_clip_model": OptionInfo("ViT-L-14/openai", "CLiP: default model", gr.Dropdown, lambda: {"choices": get_clip_models(), "visible": False}, refresh=refresh_clip_models),
|
||||
"interrogate_clip_mode": OptionInfo(caption_types[0], "CLiP: default mode", gr.Dropdown, {"choices": caption_types, "visible": False}),
|
||||
"interrogate_blip_model": OptionInfo(list(caption_models)[0], "CLiP: default captioner", gr.Dropdown, {"choices": list(caption_models), "visible": False}),
|
||||
"interrogate_clip_num_beams": OptionInfo(1, "CLiP: num beams", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1, "visible": False}),
|
||||
"interrogate_clip_min_length": OptionInfo(32, "CLiP: min length", gr.Slider, {"minimum": 1, "maximum": 128, "step": 1, "visible": False}),
|
||||
"interrogate_clip_max_length": OptionInfo(74, "CLiP: max length", gr.Slider, {"minimum": 1, "maximum": 512, "step": 1, "visible": False}),
|
||||
"interrogate_clip_min_flavors": OptionInfo(2, "CLiP: min flavors", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1, "visible": False}),
|
||||
"interrogate_clip_max_flavors": OptionInfo(16, "CLiP: max flavors", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1, "visible": False}),
|
||||
"interrogate_clip_flavor_count": OptionInfo(1024, "CLiP: intermediate flavors", gr.Slider, {"minimum": 256, "maximum": 4096, "step": 64, "visible": False}),
|
||||
"interrogate_clip_chunk_size": OptionInfo(1024, "CLiP: chunk size", gr.Slider, {"minimum": 256, "maximum": 4096, "step": 64, "visible": False}),
|
||||
|
||||
# VLM settings (hidden - controlled via Caption Tab)
|
||||
"interrogate_vlm_sep": OptionInfo("<h2>VLM</h2>", "", gr.HTML, {"visible": False}),
|
||||
"interrogate_vlm_model": OptionInfo(vlm_default, "VLM: default model", gr.Dropdown, {"choices": list(vlm_models), "visible": False}),
|
||||
"interrogate_vlm_prompt": OptionInfo(vlm_prompts[2], "VLM: default prompt", DropdownEditable, {"choices": vlm_prompts, "visible": False}),
|
||||
"interrogate_vlm_system": OptionInfo(vlm_system, "VLM: default prompt", gr.Textbox, {"visible": False}),
|
||||
"interrogate_vlm_num_beams": OptionInfo(1, "VLM: num beams", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1, "visible": False}),
|
||||
"interrogate_vlm_max_length": OptionInfo(512, "VLM: max length", gr.Slider, {"minimum": 1, "maximum": 4096, "step": 1, "visible": False}),
|
||||
"interrogate_vlm_do_sample": OptionInfo(True, "VLM: use sample method", gr.Checkbox, {"visible": False}),
|
||||
"interrogate_vlm_temperature": OptionInfo(0.8, "VLM: temperature", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01, "visible": False}),
|
||||
"interrogate_vlm_top_k": OptionInfo(0, "VLM: top-k", gr.Slider, {"minimum": 0, "maximum": 99, "step": 1, "visible": False}),
|
||||
"interrogate_vlm_top_p": OptionInfo(0, "VLM: top-p", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01, "visible": False}),
|
||||
"interrogate_vlm_keep_prefill": OptionInfo(False, "VLM: keep prefill text in output", gr.Checkbox, {"visible": False}),
|
||||
"interrogate_vlm_keep_thinking": OptionInfo(False, "VLM: keep reasoning trace in output", gr.Checkbox, {"visible": False}),
|
||||
"interrogate_vlm_thinking_mode": OptionInfo(False, "VLM: enable thinking/reasoning mode", gr.Checkbox, {"visible": False}),
|
||||
|
||||
# Common tagger settings (hidden - controlled via Caption Tab)
|
||||
"tagger_sep": OptionInfo("<h2>Tagger Settings</h2>", "", gr.HTML, {"visible": False}),
|
||||
"tagger_max_tags": OptionInfo(74, "Tagger: max tags", gr.Slider, {"minimum": 1, "maximum": 512, "step": 1, "visible": False}),
|
||||
"tagger_sort_alpha": OptionInfo(False, "Tagger: sort alphabetically", gr.Checkbox, {"visible": False}),
|
||||
"tagger_use_spaces": OptionInfo(False, "Tagger: use spaces for tags", gr.Checkbox, {"visible": False}),
|
||||
"tagger_escape": OptionInfo(True, "Tagger: escape brackets", gr.Checkbox, {"visible": False}),
|
||||
"tagger_exclude_tags": OptionInfo("", "Tagger: exclude tags", gr.Textbox, {"visible": False}),
|
||||
|
||||
# DeepBooru-specific settings (hidden - controlled via Caption Tab)
|
||||
"deepbooru_sep": OptionInfo("<h2>DeepBooru</h2>", "", gr.HTML, {"visible": False}),
|
||||
"deepbooru_score_threshold": OptionInfo(0.65, "DeepBooru: score threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": False}),
|
||||
|
||||
# WD14-specific settings (hidden - controlled via Caption Tab)
|
||||
"wd14_sep": OptionInfo("<h2>WD14 Tagger</h2>", "", gr.HTML, {"visible": False}),
|
||||
"wd14_model": OptionInfo("wd-eva02-large-tagger-v3", "WD14: default model", gr.Dropdown, {"choices": [], "visible": False}),
|
||||
"wd14_general_threshold": OptionInfo(0.35, "WD14: general tag threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": False}),
|
||||
"wd14_character_threshold": OptionInfo(0.85, "WD14: character tag threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": False}),
|
||||
"wd14_include_rating": OptionInfo(False, "WD14: include rating tags", gr.Checkbox, {"visible": False}),
|
||||
}))
|
||||
|
||||
options_templates.update(options_section(('huggingface', "Huggingface"), {
|
||||
"huggingface_sep": OptionInfo("<h2>Huggingface</h2>", "", gr.HTML),
|
||||
@@ -793,6 +742,41 @@ options_templates.update(options_section(('hidden_options', "Hidden options"), {
|
||||
"sd_checkpoint_hash": OptionInfo("", "SHA256 hash of the current checkpoint", gr.Textbox, {"visible": False}),
|
||||
"tooltips": OptionInfo("UI Tooltips", "UI tooltips", gr.Radio, {"choices": ["None", "Browser default", "UI tooltips"], "visible": False}),
|
||||
|
||||
# Caption/Interrogate settings (controlled via Caption Tab UI)
|
||||
"interrogate_default_type": OptionInfo("VLM", "Default caption type", gr.Radio, {"choices": ["VLM", "OpenCLiP", "Tagger"], "visible": False}),
|
||||
"tagger_show_scores": OptionInfo(False, "Tagger: show confidence scores in results", gr.Checkbox, {"visible": False}),
|
||||
"interrogate_clip_model": OptionInfo("ViT-L-14/openai", "OpenCLiP: default model", gr.Dropdown, lambda: {"choices": get_clip_models(), "visible": False}, refresh=refresh_clip_models),
|
||||
"interrogate_clip_mode": OptionInfo(caption_types[0], "OpenCLiP: default mode", gr.Dropdown, {"choices": caption_types, "visible": False}),
|
||||
"interrogate_blip_model": OptionInfo(list(caption_models)[0], "OpenCLiP: default captioner", gr.Dropdown, {"choices": list(caption_models), "visible": False}),
|
||||
"interrogate_clip_num_beams": OptionInfo(1, "OpenCLiP: num beams", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1, "visible": False}),
|
||||
"interrogate_clip_min_length": OptionInfo(32, "OpenCLiP: min length", gr.Slider, {"minimum": 1, "maximum": 128, "step": 1, "visible": False}),
|
||||
"interrogate_clip_max_length": OptionInfo(74, "OpenCLiP: max length", gr.Slider, {"minimum": 1, "maximum": 512, "step": 1, "visible": False}),
|
||||
"interrogate_clip_min_flavors": OptionInfo(2, "OpenCLiP: min flavors", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1, "visible": False}),
|
||||
"interrogate_clip_max_flavors": OptionInfo(16, "OpenCLiP: max flavors", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1, "visible": False}),
|
||||
"interrogate_clip_flavor_count": OptionInfo(1024, "OpenCLiP: intermediate flavors", gr.Slider, {"minimum": 256, "maximum": 4096, "step": 64, "visible": False}),
|
||||
"interrogate_clip_chunk_size": OptionInfo(1024, "OpenCLiP: chunk size", gr.Slider, {"minimum": 256, "maximum": 4096, "step": 64, "visible": False}),
|
||||
"interrogate_vlm_model": OptionInfo(vlm_default, "VLM: default model", gr.Dropdown, {"choices": list(vlm_models), "visible": False}),
|
||||
"interrogate_vlm_prompt": OptionInfo(vlm_prompts[2], "VLM: default prompt", DropdownEditable, {"choices": vlm_prompts, "visible": False}),
|
||||
"interrogate_vlm_system": OptionInfo(vlm_system, "VLM: system prompt", gr.Textbox, {"visible": False}),
|
||||
"interrogate_vlm_num_beams": OptionInfo(1, "VLM: num beams", gr.Slider, {"minimum": 1, "maximum": 16, "step": 1, "visible": False}),
|
||||
"interrogate_vlm_max_length": OptionInfo(512, "VLM: max length", gr.Slider, {"minimum": 1, "maximum": 4096, "step": 1, "visible": False}),
|
||||
"interrogate_vlm_do_sample": OptionInfo(True, "VLM: use sample method", gr.Checkbox, {"visible": False}),
|
||||
"interrogate_vlm_temperature": OptionInfo(0.8, "VLM: temperature", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01, "visible": False}),
|
||||
"interrogate_vlm_top_k": OptionInfo(0, "VLM: top-k", gr.Slider, {"minimum": 0, "maximum": 99, "step": 1, "visible": False}),
|
||||
"interrogate_vlm_top_p": OptionInfo(0, "VLM: top-p", gr.Slider, {"minimum": 0, "maximum": 1.0, "step": 0.01, "visible": False}),
|
||||
"interrogate_vlm_keep_prefill": OptionInfo(False, "VLM: keep prefill text in output", gr.Checkbox, {"visible": False}),
|
||||
"interrogate_vlm_keep_thinking": OptionInfo(False, "VLM: keep reasoning trace in output", gr.Checkbox, {"visible": False}),
|
||||
"interrogate_vlm_thinking_mode": OptionInfo(False, "VLM: enable thinking/reasoning mode", gr.Checkbox, {"visible": False}),
|
||||
"tagger_threshold": OptionInfo(0.50, "Tagger: general tag threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": False}),
|
||||
"tagger_include_rating": OptionInfo(False, "Tagger: include rating tags", gr.Checkbox, {"visible": False}),
|
||||
"tagger_max_tags": OptionInfo(74, "Tagger: max tags", gr.Slider, {"minimum": 1, "maximum": 512, "step": 1, "visible": False}),
|
||||
"tagger_sort_alpha": OptionInfo(False, "Tagger: sort alphabetically", gr.Checkbox, {"visible": False}),
|
||||
"tagger_use_spaces": OptionInfo(False, "Tagger: use spaces for tags", gr.Checkbox, {"visible": False}),
|
||||
"tagger_escape_brackets": OptionInfo(True, "Tagger: escape brackets", gr.Checkbox, {"visible": False}),
|
||||
"tagger_exclude_tags": OptionInfo("", "Tagger: exclude tags", gr.Textbox, {"visible": False}),
|
||||
"wd14_model": OptionInfo("wd-eva02-large-tagger-v3", "WD14: default model", gr.Dropdown, {"choices": [], "visible": False}),
|
||||
"wd14_character_threshold": OptionInfo(0.85, "WD14: character tag threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01, "visible": False}),
|
||||
|
||||
# control settings are handled separately
|
||||
"control_hires": OptionInfo(False, "Hires use Control", gr.Checkbox, {"visible": False}),
|
||||
"control_aspect_ratio": OptionInfo(False, "Aspect ratio resize", gr.Checkbox, {"visible": False}),
|
||||
|
||||
+33
-12
@@ -85,28 +85,29 @@ def tagger_batch_wrapper(model_name, batch_files, batch_folder, batch_str, save_
|
||||
def update_tagger_ui(model_name):
|
||||
"""Update UI controls based on selected tagger model.
|
||||
|
||||
When DeepBooru is selected, character_threshold and include_rating are disabled
|
||||
since DeepBooru doesn't support separate character threshold or rating tags.
|
||||
When DeepBooru is selected, character_threshold is disabled since DeepBooru
|
||||
doesn't support separate character threshold.
|
||||
"""
|
||||
from modules.interrogate import tagger
|
||||
is_db = tagger.is_deepbooru(model_name)
|
||||
return [
|
||||
gr.update(interactive=not is_db), # character_threshold
|
||||
gr.update(interactive=not is_db, value=False if is_db else None), # include_rating
|
||||
gr.update(), # include_rating - now supported by both taggers
|
||||
]
|
||||
|
||||
|
||||
def update_tagger_params(model_name, general_threshold, character_threshold, include_rating, max_tags, sort_alpha, use_spaces, escape_brackets, exclude_tags):
|
||||
def update_tagger_params(model_name, general_threshold, character_threshold, include_rating, max_tags, sort_alpha, use_spaces, escape_brackets, exclude_tags, show_scores):
|
||||
"""Save all tagger parameters to shared.opts when UI controls change."""
|
||||
shared.opts.wd14_model = model_name
|
||||
shared.opts.wd14_general_threshold = float(general_threshold)
|
||||
shared.opts.tagger_threshold = float(general_threshold)
|
||||
shared.opts.wd14_character_threshold = float(character_threshold)
|
||||
shared.opts.wd14_include_rating = bool(include_rating)
|
||||
shared.opts.tagger_include_rating = bool(include_rating)
|
||||
shared.opts.tagger_max_tags = int(max_tags)
|
||||
shared.opts.tagger_sort_alpha = bool(sort_alpha)
|
||||
shared.opts.tagger_use_spaces = bool(use_spaces)
|
||||
shared.opts.tagger_escape = bool(escape_brackets)
|
||||
shared.opts.tagger_escape_brackets = bool(escape_brackets)
|
||||
shared.opts.tagger_exclude_tags = str(exclude_tags)
|
||||
shared.opts.tagger_show_scores = bool(show_scores)
|
||||
shared.opts.save()
|
||||
|
||||
|
||||
@@ -138,6 +139,12 @@ def update_vlm_model_params(vlm_model, vlm_system):
|
||||
shared.opts.save()
|
||||
|
||||
|
||||
def update_default_caption_type(caption_type):
|
||||
"""Save the default caption type to shared.opts."""
|
||||
shared.opts.interrogate_default_type = str(caption_type)
|
||||
shared.opts.save()
|
||||
|
||||
|
||||
def create_ui():
|
||||
shared.log.debug('UI initialize: tab=caption')
|
||||
with gr.Row(equal_height=False, variant='compact', elem_classes="caption", elem_id="caption_tab"):
|
||||
@@ -200,7 +207,7 @@ def create_ui():
|
||||
btn_vlm_caption_batch = gr.Button("Batch Caption", variant='primary', elem_id="btn_vlm_caption_batch")
|
||||
with gr.Row():
|
||||
btn_vlm_caption = gr.Button("Caption", variant='primary', elem_id="btn_vlm_caption")
|
||||
with gr.Tab("CLiP Interrogate", elem_id='tab_clip_interrogate'):
|
||||
with gr.Tab("OpenCLiP", elem_id='tab_clip_interrogate'):
|
||||
with gr.Row():
|
||||
clip_model = gr.Dropdown([], value=shared.opts.interrogate_clip_model, label='CLiP Model', elem_id='clip_clip_model')
|
||||
ui_common.create_refresh_button(clip_model, openclip.refresh_clip_models, lambda: {"choices": openclip.refresh_clip_models()}, 'clip_models_refresh')
|
||||
@@ -250,17 +257,19 @@ def create_ui():
|
||||
wd_unload_btn = gr.Button(value='Unload', elem_id='wd_unload', variant='secondary')
|
||||
with gr.Accordion(label='Tagger: Advanced Options', open=True, visible=True):
|
||||
with gr.Row():
|
||||
wd_general_threshold = gr.Slider(label='General threshold', value=shared.opts.wd14_general_threshold, minimum=0.0, maximum=1.0, step=0.01, elem_id='wd_general_threshold')
|
||||
wd_general_threshold = gr.Slider(label='General threshold', value=shared.opts.tagger_threshold, minimum=0.0, maximum=1.0, step=0.01, elem_id='wd_general_threshold')
|
||||
wd_character_threshold = gr.Slider(label='Character threshold', value=shared.opts.wd14_character_threshold, minimum=0.0, maximum=1.0, step=0.01, elem_id='wd_character_threshold')
|
||||
with gr.Row():
|
||||
wd_max_tags = gr.Slider(label='Max tags', value=shared.opts.tagger_max_tags, minimum=1, maximum=512, step=1, elem_id='wd_max_tags')
|
||||
wd_include_rating = gr.Checkbox(label='Include rating', value=shared.opts.wd14_include_rating, elem_id='wd_include_rating')
|
||||
wd_include_rating = gr.Checkbox(label='Include rating', value=shared.opts.tagger_include_rating, elem_id='wd_include_rating')
|
||||
with gr.Row():
|
||||
wd_sort_alpha = gr.Checkbox(label='Sort alphabetically', value=shared.opts.tagger_sort_alpha, elem_id='wd_sort_alpha')
|
||||
wd_use_spaces = gr.Checkbox(label='Use spaces', value=shared.opts.tagger_use_spaces, elem_id='wd_use_spaces')
|
||||
wd_escape = gr.Checkbox(label='Escape brackets', value=shared.opts.tagger_escape, elem_id='wd_escape')
|
||||
wd_escape = gr.Checkbox(label='Escape brackets', value=shared.opts.tagger_escape_brackets, elem_id='wd_escape')
|
||||
with gr.Row():
|
||||
wd_exclude_tags = gr.Textbox(label='Exclude tags', value=shared.opts.tagger_exclude_tags, placeholder='Comma-separated tags to exclude', elem_id='wd_exclude_tags')
|
||||
with gr.Row():
|
||||
wd_show_scores = gr.Checkbox(label='Show confidence scores', value=shared.opts.tagger_show_scores, elem_id='wd_show_scores')
|
||||
gr.HTML('<style>#wd_character_threshold:has(input:disabled), #wd_include_rating:has(input:disabled) { opacity: 0.5; }</style>')
|
||||
with gr.Accordion(label='Tagger: Batch', open=False, visible=True):
|
||||
with gr.Row():
|
||||
@@ -277,6 +286,14 @@ def create_ui():
|
||||
btn_wd_tag_batch = gr.Button("Batch Tag", variant='primary', elem_id="btn_wd_tag_batch")
|
||||
with gr.Row():
|
||||
btn_wd_tag = gr.Button("Tag", variant='primary', elem_id="btn_wd_tag")
|
||||
with gr.Tab("Interrogate", elem_id='tab_interrogate'):
|
||||
with gr.Row():
|
||||
default_caption_type = gr.Radio(
|
||||
choices=["VLM", "OpenCLiP", "Tagger"],
|
||||
value=shared.opts.interrogate_default_type,
|
||||
label="Default Caption Type",
|
||||
elem_id="default_caption_type"
|
||||
)
|
||||
with gr.Column(variant='compact', elem_id='interrogate_output'):
|
||||
with gr.Row(elem_id='interrogate_output_prompt'):
|
||||
prompt = gr.Textbox(label="Answer", lines=12, placeholder="ai generated image description")
|
||||
@@ -320,7 +337,7 @@ def create_ui():
|
||||
wd_model.change(fn=update_tagger_ui, inputs=[wd_model], outputs=[wd_character_threshold, wd_include_rating], show_progress=False)
|
||||
|
||||
# Save tagger parameters to shared.opts when UI controls change
|
||||
tagger_inputs = [wd_model, wd_general_threshold, wd_character_threshold, wd_include_rating, wd_max_tags, wd_sort_alpha, wd_use_spaces, wd_escape, wd_exclude_tags]
|
||||
tagger_inputs = [wd_model, wd_general_threshold, wd_character_threshold, wd_include_rating, wd_max_tags, wd_sort_alpha, wd_use_spaces, wd_escape, wd_exclude_tags, wd_show_scores]
|
||||
wd_model.change(fn=update_tagger_params, inputs=tagger_inputs, outputs=[], show_progress=False)
|
||||
wd_general_threshold.change(fn=update_tagger_params, inputs=tagger_inputs, outputs=[], show_progress=False)
|
||||
wd_character_threshold.change(fn=update_tagger_params, inputs=tagger_inputs, outputs=[], show_progress=False)
|
||||
@@ -330,6 +347,7 @@ def create_ui():
|
||||
wd_use_spaces.change(fn=update_tagger_params, inputs=tagger_inputs, outputs=[], show_progress=False)
|
||||
wd_escape.change(fn=update_tagger_params, inputs=tagger_inputs, outputs=[], show_progress=False)
|
||||
wd_exclude_tags.change(fn=update_tagger_params, inputs=tagger_inputs, outputs=[], show_progress=False)
|
||||
wd_show_scores.change(fn=update_tagger_params, inputs=tagger_inputs, outputs=[], show_progress=False)
|
||||
|
||||
# Save CLiP model parameters to shared.opts when UI controls change
|
||||
clip_model_inputs = [clip_model, blip_model, clip_mode]
|
||||
@@ -342,6 +360,9 @@ def create_ui():
|
||||
vlm_model.change(fn=update_vlm_model_params, inputs=vlm_model_inputs, outputs=[], show_progress=False)
|
||||
vlm_system.change(fn=update_vlm_model_params, inputs=vlm_model_inputs, outputs=[], show_progress=False)
|
||||
|
||||
# Save default caption type to shared.opts when UI control changes
|
||||
default_caption_type.change(fn=update_default_caption_type, inputs=[default_caption_type], outputs=[], show_progress=False)
|
||||
|
||||
for tabname, button in copy_interrogate_buttons.items():
|
||||
generation_parameters_copypaste.register_paste_params_button(generation_parameters_copypaste.ParamBinding(paste_button=button, tabname=tabname, source_text_component=prompt, source_image_component=image,))
|
||||
generation_parameters_copypaste.add_paste_fields("caption", image, None)
|
||||
|
||||
Reference in New Issue
Block a user