Automated formatting pass across 30 files to establish consistent code style enforced by CI. No logic changes.
712 lines
24 KiB
Python
712 lines
24 KiB
Python
"""
|
|
AutoTag Service Module
|
|
|
|
Provides JoyCaption-based automatic tagging for images using LLM models.
|
|
Refactored from standalone_tagger.py for integration with PromptManager API.
|
|
"""
|
|
|
|
import gc
|
|
import os
|
|
from pathlib import Path
|
|
from typing import Any, Callable, Dict, List, Optional, Tuple
|
|
from PIL import Image
|
|
|
|
# Try to import logging from utils, fallback to standard logging
|
|
try:
|
|
from ..utils.logging_config import get_logger
|
|
except ImportError:
|
|
import logging
|
|
|
|
def get_logger(name: str) -> logging.Logger:
|
|
logger = logging.getLogger(name)
|
|
if not logger.handlers:
|
|
handler = logging.StreamHandler()
|
|
handler.setFormatter(
|
|
logging.Formatter(
|
|
"%(asctime)s - %(name)s - %(levelname)s - %(message)s"
|
|
)
|
|
)
|
|
logger.addHandler(handler)
|
|
logger.setLevel(logging.DEBUG)
|
|
return logger
|
|
|
|
|
|
# Model configurations - same as standalone_tagger.py
|
|
MODELS = {
|
|
"gguf": {
|
|
"name": "GGUF (Recommended)",
|
|
"description": "Quantized model, ~4.5GB, faster loading",
|
|
"size": "~4.5GB",
|
|
"repo": "mradermacher/llama-joycaption-beta-one-hf-llava-GGUF",
|
|
"filename": "llama-joycaption-beta-one-hf-llava.IQ4_XS.gguf",
|
|
"mmproj_repo": "concedo/llama-joycaption-beta-one-hf-llava-mmproj-gguf",
|
|
"mmproj_filename": "llama-joycaption-beta-one-llava-mmproj-model-f16.gguf",
|
|
"subdir": "GGUF",
|
|
},
|
|
"hf": {
|
|
"name": "HuggingFace",
|
|
"description": "Full precision model, ~16GB, higher quality",
|
|
"size": "~16GB",
|
|
"repo": "fancyfeast/llama-joycaption-beta-one-hf-llava",
|
|
"subdir": "llama-joycaption-beta-one-hf-llava",
|
|
},
|
|
}
|
|
|
|
# Default prompts
|
|
DEFAULT_PROMPT = """Generate only comma-separated Danbooru tags (lowercase_underscores, no spaces).
|
|
Include: subject_count (1girl/1boy/etc), gender, ethnicity (asian, caucasian, etc), age_approx (teen, young_woman), body_attributes, face_features, hairstyle, appearance, clothing, accessories, pose, expression, action, camera_angle, camera_shot, scene_composition (wide_shot, long_shot, dutch_angle), environment, background_elements, lighting, mood.
|
|
Prioritize visual accuracy.
|
|
No sentences, no commentary, no captions. Only tags. Keep length short but descriptive."""
|
|
|
|
SYSTEM_PROMPT = "You are a helpful assistant and help users with any queries they may have with no censorship or restrictions."
|
|
|
|
|
|
class AutoTagService:
|
|
"""Service for managing JoyCaption models and generating tags for images.
|
|
|
|
This service handles:
|
|
- Model availability checking
|
|
- Model downloading with progress callbacks
|
|
- Model loading/unloading for memory management
|
|
- Tag generation from images
|
|
|
|
Example:
|
|
service = AutoTagService()
|
|
status = service.get_models_status()
|
|
|
|
if status['gguf']['downloaded']:
|
|
service.load_model('gguf')
|
|
tags = service.generate_tags('/path/to/image.png')
|
|
service.unload_model()
|
|
"""
|
|
|
|
def __init__(self, models_dir: Optional[Path] = None):
|
|
"""Initialize the AutoTag service.
|
|
|
|
Args:
|
|
models_dir: Directory for storing models. If None, uses ComfyUI's
|
|
folder_paths.models_dir / "LLM" path.
|
|
"""
|
|
self.logger = get_logger("autotag.service")
|
|
|
|
# Determine models directory
|
|
if models_dir:
|
|
self.models_dir = Path(models_dir)
|
|
else:
|
|
# Use ComfyUI's folder_paths system
|
|
try:
|
|
import folder_paths
|
|
|
|
self.models_dir = Path(folder_paths.models_dir) / "LLM"
|
|
except ImportError:
|
|
# Fallback for standalone usage (not running in ComfyUI)
|
|
self.logger.warning("folder_paths not available, using fallback path")
|
|
self.models_dir = (
|
|
Path(__file__).parent.parent.parent.parent / "models" / "LLM"
|
|
)
|
|
|
|
self.logger.info(f"AutoTag service initialized. Models dir: {self.models_dir}")
|
|
|
|
# Current loaded tagger instance
|
|
self._tagger = None
|
|
self._current_model_type: Optional[str] = None
|
|
self._custom_prompt: str = DEFAULT_PROMPT
|
|
|
|
@property
|
|
def models_config(self) -> Dict[str, Dict[str, Any]]:
|
|
"""Get the models configuration dictionary."""
|
|
return MODELS
|
|
|
|
@property
|
|
def default_prompt(self) -> str:
|
|
"""Get the default tag generation prompt."""
|
|
return DEFAULT_PROMPT
|
|
|
|
@property
|
|
def custom_prompt(self) -> str:
|
|
"""Get the current custom prompt."""
|
|
return self._custom_prompt
|
|
|
|
@custom_prompt.setter
|
|
def custom_prompt(self, value: str):
|
|
"""Set a custom prompt for tag generation."""
|
|
self._custom_prompt = value
|
|
|
|
def get_models_status(self) -> Dict[str, Dict[str, Any]]:
|
|
"""Get availability status for all model types.
|
|
|
|
Returns:
|
|
Dictionary with status for each model type:
|
|
{
|
|
'gguf': {
|
|
'name': 'GGUF (Recommended)',
|
|
'description': '...',
|
|
'size': '~4.5GB',
|
|
'downloaded': True,
|
|
'model_exists': True,
|
|
'mmproj_exists': True, # GGUF only
|
|
'model_path': '/path/to/model'
|
|
},
|
|
'hf': {...}
|
|
}
|
|
"""
|
|
status = {}
|
|
|
|
for model_type, config in MODELS.items():
|
|
model_status = {
|
|
"name": config["name"],
|
|
"description": config["description"],
|
|
"size": config["size"],
|
|
"downloaded": False,
|
|
"model_path": None,
|
|
}
|
|
|
|
if model_type == "gguf":
|
|
model_exists, mmproj_exists = self._check_gguf_models()
|
|
model_status["model_exists"] = model_exists
|
|
model_status["mmproj_exists"] = mmproj_exists
|
|
model_status["downloaded"] = model_exists and mmproj_exists
|
|
if model_status["downloaded"]:
|
|
model_status["model_path"] = str(
|
|
self.models_dir / config["subdir"] / config["filename"]
|
|
)
|
|
else: # hf
|
|
model_status["downloaded"] = self._check_hf_model()
|
|
if model_status["downloaded"]:
|
|
# Get the actual path (local or cache)
|
|
model_status["model_path"] = str(self._get_hf_model_path())
|
|
|
|
status[model_type] = model_status
|
|
|
|
return status
|
|
|
|
def _check_gguf_models(self) -> Tuple[bool, bool]:
|
|
"""Check if GGUF model and mmproj files exist.
|
|
|
|
Returns:
|
|
Tuple of (model_exists, mmproj_exists)
|
|
"""
|
|
config = MODELS["gguf"]
|
|
gguf_dir = self.models_dir / config["subdir"]
|
|
model_path = gguf_dir / config["filename"]
|
|
mmproj_path = gguf_dir / config["mmproj_filename"]
|
|
return model_path.exists(), mmproj_path.exists()
|
|
|
|
def _check_hf_model(self) -> bool:
|
|
"""Check if HuggingFace model exists.
|
|
|
|
Checks both the local models directory and the HuggingFace cache.
|
|
|
|
Returns:
|
|
True if model directory contains config.json
|
|
"""
|
|
config = MODELS["hf"]
|
|
|
|
# Check local directory first
|
|
model_dir = self.models_dir / config["subdir"]
|
|
self.logger.debug(f"Checking local HF model path: {model_dir}")
|
|
if (model_dir / "config.json").exists():
|
|
self.logger.debug("Found HF model in local directory")
|
|
return True
|
|
|
|
# Check HuggingFace cache as fallback
|
|
hf_cache_path = self._get_hf_cache_path(config["repo"])
|
|
if hf_cache_path:
|
|
self.logger.debug(f"Found HF model in cache: {hf_cache_path}")
|
|
return True
|
|
|
|
self.logger.debug("HF model not found in local dir or cache")
|
|
return False
|
|
|
|
def _get_hf_cache_path(self, repo_id: str) -> Optional[Path]:
|
|
"""Get the path to a model in the HuggingFace cache.
|
|
|
|
Args:
|
|
repo_id: The HuggingFace repo ID (e.g., 'fancyfeast/llama-joycaption-beta-one-hf-llava')
|
|
|
|
Returns:
|
|
Path to the cached model directory, or None if not found
|
|
"""
|
|
try:
|
|
from huggingface_hub import scan_cache_dir, HFCacheInfo
|
|
except ImportError:
|
|
self.logger.debug("huggingface_hub not available for cache check")
|
|
return None
|
|
|
|
try:
|
|
cache_info = scan_cache_dir()
|
|
for repo in cache_info.repos:
|
|
if repo.repo_id == repo_id and repo.repo_type == "model":
|
|
# Get the latest revision's snapshot path
|
|
for revision in repo.revisions:
|
|
snapshot_path = revision.snapshot_path
|
|
if (Path(snapshot_path) / "config.json").exists():
|
|
return Path(snapshot_path)
|
|
except Exception as e:
|
|
self.logger.debug(f"Error scanning HF cache: {e}")
|
|
|
|
return None
|
|
|
|
def _get_hf_model_path(self) -> Optional[Path]:
|
|
"""Get the actual path to the HuggingFace model.
|
|
|
|
Checks local directory first, then HuggingFace cache.
|
|
|
|
Returns:
|
|
Path to the model directory, or None if not found
|
|
"""
|
|
config = MODELS["hf"]
|
|
|
|
# Check local directory first
|
|
model_dir = self.models_dir / config["subdir"]
|
|
if (model_dir / "config.json").exists():
|
|
return model_dir
|
|
|
|
# Check HuggingFace cache as fallback
|
|
cache_path = self._get_hf_cache_path(config["repo"])
|
|
if cache_path:
|
|
return cache_path
|
|
|
|
return None
|
|
|
|
def download_model(
|
|
self,
|
|
model_type: str,
|
|
progress_callback: Optional[Callable[[str, float], None]] = None,
|
|
) -> bool:
|
|
"""Download a model with optional progress updates.
|
|
|
|
Args:
|
|
model_type: Either 'gguf' or 'hf'
|
|
progress_callback: Optional callback(status_message, progress_percent)
|
|
|
|
Returns:
|
|
True if download successful, False otherwise
|
|
|
|
Raises:
|
|
ValueError: If model_type is not valid
|
|
"""
|
|
if model_type not in MODELS:
|
|
raise ValueError(
|
|
f"Invalid model type: {model_type}. Must be 'gguf' or 'hf'"
|
|
)
|
|
|
|
try:
|
|
from huggingface_hub import hf_hub_download, snapshot_download
|
|
except ImportError:
|
|
self.logger.error("huggingface_hub not installed")
|
|
if progress_callback:
|
|
progress_callback("Error: huggingface_hub not installed", 0)
|
|
return False
|
|
|
|
try:
|
|
if model_type == "gguf":
|
|
return self._download_gguf_models(progress_callback)
|
|
else:
|
|
return self._download_hf_model(progress_callback)
|
|
except Exception as e:
|
|
self.logger.error(f"Download failed: {e}")
|
|
if progress_callback:
|
|
progress_callback(f"Error: {str(e)}", 0)
|
|
return False
|
|
|
|
def _download_gguf_models(
|
|
self, progress_callback: Optional[Callable[[str, float], None]] = None
|
|
) -> bool:
|
|
"""Download GGUF model and mmproj files."""
|
|
from huggingface_hub import hf_hub_download
|
|
|
|
config = MODELS["gguf"]
|
|
gguf_dir = self.models_dir / config["subdir"]
|
|
gguf_dir.mkdir(parents=True, exist_ok=True)
|
|
|
|
model_path = gguf_dir / config["filename"]
|
|
mmproj_path = gguf_dir / config["mmproj_filename"]
|
|
|
|
# Download main model
|
|
if not model_path.exists():
|
|
if progress_callback:
|
|
progress_callback(f"Downloading {config['filename']}...", 10)
|
|
|
|
self.logger.info(f"Downloading GGUF model: {config['filename']}")
|
|
hf_hub_download(
|
|
repo_id=config["repo"],
|
|
filename=config["filename"],
|
|
local_dir=str(gguf_dir),
|
|
local_dir_use_symlinks=False,
|
|
)
|
|
self.logger.info("GGUF model downloaded")
|
|
|
|
if progress_callback:
|
|
progress_callback("Main model ready", 50)
|
|
|
|
# Download mmproj
|
|
if not mmproj_path.exists():
|
|
if progress_callback:
|
|
progress_callback(f"Downloading {config['mmproj_filename']}...", 60)
|
|
|
|
self.logger.info(f"Downloading mmproj: {config['mmproj_filename']}")
|
|
hf_hub_download(
|
|
repo_id=config["mmproj_repo"],
|
|
filename=config["mmproj_filename"],
|
|
local_dir=str(gguf_dir),
|
|
local_dir_use_symlinks=False,
|
|
)
|
|
self.logger.info("mmproj downloaded")
|
|
|
|
if progress_callback:
|
|
progress_callback("Download complete", 100)
|
|
|
|
return True
|
|
|
|
def _download_hf_model(
|
|
self, progress_callback: Optional[Callable[[str, float], None]] = None
|
|
) -> bool:
|
|
"""Download HuggingFace model."""
|
|
from huggingface_hub import snapshot_download
|
|
|
|
config = MODELS["hf"]
|
|
model_dir = self.models_dir / config["subdir"]
|
|
|
|
if not self._check_hf_model():
|
|
if progress_callback:
|
|
progress_callback(f"Downloading {config['repo']}...", 10)
|
|
|
|
self.logger.info(f"Downloading HF model: {config['repo']}")
|
|
snapshot_download(
|
|
repo_id=config["repo"],
|
|
local_dir=str(model_dir),
|
|
local_dir_use_symlinks=False,
|
|
)
|
|
self.logger.info("HF model downloaded")
|
|
|
|
if progress_callback:
|
|
progress_callback("Download complete", 100)
|
|
|
|
return True
|
|
|
|
def load_model(self, model_type: str, use_gpu: bool = True) -> bool:
|
|
"""Load a model into memory for tag generation.
|
|
|
|
Args:
|
|
model_type: Either 'gguf' or 'hf'
|
|
use_gpu: Whether to use GPU acceleration (default True)
|
|
|
|
Returns:
|
|
True if model loaded successfully
|
|
|
|
Raises:
|
|
ValueError: If model_type is invalid
|
|
RuntimeError: If model not downloaded or loading fails
|
|
"""
|
|
if model_type not in MODELS:
|
|
raise ValueError(f"Invalid model type: {model_type}")
|
|
|
|
# Unload existing model first
|
|
if self._tagger is not None:
|
|
self.unload_model()
|
|
|
|
status = self.get_models_status()
|
|
if not status[model_type]["downloaded"]:
|
|
raise RuntimeError(f"Model {model_type} not downloaded")
|
|
|
|
try:
|
|
if model_type == "gguf":
|
|
self._tagger = self._load_gguf_tagger(use_gpu)
|
|
else:
|
|
self._tagger = self._load_hf_tagger()
|
|
|
|
self._current_model_type = model_type
|
|
self.logger.info(f"Model {model_type} loaded successfully")
|
|
return True
|
|
|
|
except Exception as e:
|
|
self.logger.error(f"Failed to load model {model_type}: {e}")
|
|
self._tagger = None
|
|
self._current_model_type = None
|
|
raise RuntimeError(f"Failed to load model: {e}")
|
|
|
|
def _load_gguf_tagger(self, use_gpu: bool = True):
|
|
"""Load GGUF-based tagger."""
|
|
from llama_cpp import Llama
|
|
from llama_cpp.llama_chat_format import Llava15ChatHandler
|
|
|
|
config = MODELS["gguf"]
|
|
gguf_dir = self.models_dir / config["subdir"]
|
|
model_path = gguf_dir / config["filename"]
|
|
mmproj_path = gguf_dir / config["mmproj_filename"]
|
|
|
|
self.logger.info("Loading GGUF model...")
|
|
n_gpu_layers = -1 if use_gpu else 0
|
|
|
|
tagger = Llama(
|
|
model_path=str(model_path),
|
|
n_ctx=4096,
|
|
n_batch=2048,
|
|
n_threads=4,
|
|
n_gpu_layers=n_gpu_layers,
|
|
verbose=False,
|
|
chat_handler=Llava15ChatHandler(clip_model_path=str(mmproj_path)),
|
|
offload_kqv=True,
|
|
)
|
|
|
|
self.logger.info("GGUF model loaded")
|
|
return ("gguf", tagger)
|
|
|
|
def _load_hf_tagger(self, quantization: str = "8bit"):
|
|
"""Load HuggingFace-based tagger."""
|
|
import torch
|
|
from transformers import (
|
|
AutoProcessor,
|
|
LlavaForConditionalGeneration,
|
|
BitsAndBytesConfig,
|
|
)
|
|
|
|
# Get the actual model path (local or cache)
|
|
model_path = self._get_hf_model_path()
|
|
if model_path is None:
|
|
raise RuntimeError(
|
|
"HuggingFace model not found in local directory or cache"
|
|
)
|
|
|
|
self.logger.info(f"Loading HuggingFace model from {model_path}...")
|
|
device = "cuda" if torch.cuda.is_available() else "cpu"
|
|
|
|
processor = AutoProcessor.from_pretrained(str(model_path))
|
|
model_kwargs = {"device_map": "cuda" if device == "cuda" else "cpu"}
|
|
|
|
if quantization == "8bit":
|
|
qnt_config = BitsAndBytesConfig(
|
|
load_in_8bit=True,
|
|
bnb_8bit_compute_dtype=torch.float16,
|
|
bnb_8bit_use_double_quant=True,
|
|
llm_int8_skip_modules=["vision_tower", "multi_modal_projector"],
|
|
)
|
|
model = LlavaForConditionalGeneration.from_pretrained(
|
|
str(model_path),
|
|
torch_dtype=torch.float16,
|
|
quantization_config=qnt_config,
|
|
**model_kwargs,
|
|
)
|
|
else:
|
|
model = LlavaForConditionalGeneration.from_pretrained(
|
|
str(model_path), torch_dtype=torch.bfloat16, **model_kwargs
|
|
)
|
|
|
|
model.eval()
|
|
self.logger.info("HuggingFace model loaded")
|
|
# Track the compute dtype for pixel_values conversion
|
|
compute_dtype = torch.float16 if quantization == "8bit" else torch.bfloat16
|
|
return ("hf", (model, processor, device, compute_dtype))
|
|
|
|
def unload_model(self):
|
|
"""Unload the current model and free memory."""
|
|
if self._tagger is not None:
|
|
self.logger.info(f"Unloading model: {self._current_model_type}")
|
|
del self._tagger
|
|
self._tagger = None
|
|
self._current_model_type = None
|
|
gc.collect()
|
|
|
|
# Try to clear CUDA cache if available
|
|
try:
|
|
import torch
|
|
|
|
if torch.cuda.is_available():
|
|
torch.cuda.empty_cache()
|
|
except ImportError:
|
|
pass
|
|
|
|
self.logger.info("Model unloaded, memory freed")
|
|
|
|
def is_model_loaded(self) -> bool:
|
|
"""Check if a model is currently loaded."""
|
|
return self._tagger is not None
|
|
|
|
def get_loaded_model_type(self) -> Optional[str]:
|
|
"""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]:
|
|
"""Generate tags for an image.
|
|
|
|
Args:
|
|
image_path: Path to the image file
|
|
prompt: Custom prompt for tag generation. Uses default if None.
|
|
|
|
Returns:
|
|
List of generated tags
|
|
|
|
Raises:
|
|
RuntimeError: If no model is loaded
|
|
FileNotFoundError: If image doesn't exist
|
|
"""
|
|
if self._tagger is None:
|
|
raise RuntimeError("No model loaded. Call load_model() first.")
|
|
|
|
if not os.path.exists(image_path):
|
|
raise FileNotFoundError(f"Image not found: {image_path}")
|
|
|
|
use_prompt = prompt or self._custom_prompt
|
|
|
|
# Load image
|
|
image = Image.open(image_path)
|
|
if image.mode != "RGB":
|
|
image = image.convert("RGB")
|
|
|
|
# Generate based on model type
|
|
model_type, tagger_obj = self._tagger
|
|
|
|
if model_type == "gguf":
|
|
raw_tags = self._generate_gguf(tagger_obj, image, use_prompt)
|
|
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
|
|
|
|
def _generate_gguf(self, model, image: Image.Image, prompt: str) -> str:
|
|
"""Generate tags using GGUF model."""
|
|
import base64
|
|
import io
|
|
|
|
# Resize image
|
|
image = image.resize((336, 336), Image.Resampling.BILINEAR)
|
|
|
|
# Encode to base64
|
|
buffer = io.BytesIO()
|
|
image.save(buffer, format="PNG")
|
|
buffer.seek(0)
|
|
img_base64 = base64.b64encode(buffer.read()).decode("utf-8")
|
|
data_uri = f"data:image/png;base64,{img_base64}"
|
|
|
|
# Create message
|
|
messages = [
|
|
{"role": "system", "content": SYSTEM_PROMPT},
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{"type": "text", "text": prompt},
|
|
{"type": "image_url", "image_url": {"url": data_uri}},
|
|
],
|
|
},
|
|
]
|
|
|
|
# Generate
|
|
response = model.create_chat_completion(
|
|
messages=messages,
|
|
max_tokens=512,
|
|
temperature=0.6,
|
|
top_p=0.9,
|
|
stop=["</s>", "User:", "Assistant:"],
|
|
stream=False,
|
|
)
|
|
|
|
return response["choices"][0]["message"]["content"].strip()
|
|
|
|
def _generate_hf(
|
|
self,
|
|
model,
|
|
processor,
|
|
device: str,
|
|
compute_dtype,
|
|
image: Image.Image,
|
|
prompt: str,
|
|
) -> str:
|
|
"""Generate tags using HuggingFace model."""
|
|
import torch
|
|
|
|
# Resize image
|
|
image = image.resize((336, 336), Image.Resampling.LANCZOS)
|
|
|
|
convo = [
|
|
{"role": "system", "content": SYSTEM_PROMPT},
|
|
{"role": "user", "content": prompt},
|
|
]
|
|
convo_string = processor.apply_chat_template(
|
|
convo, tokenize=False, add_generation_prompt=True
|
|
)
|
|
|
|
inputs = processor(text=[convo_string], images=[image], return_tensors="pt").to(
|
|
device
|
|
)
|
|
|
|
# Convert pixel_values to the model's compute dtype
|
|
# (float16 for 8-bit quantized, bfloat16 for non-quantized)
|
|
if "pixel_values" in inputs and inputs["pixel_values"] is not None:
|
|
inputs["pixel_values"] = inputs["pixel_values"].to(compute_dtype)
|
|
|
|
with torch.inference_mode(), torch.cuda.amp.autocast(enabled=True):
|
|
generate_ids = model.generate(
|
|
**inputs,
|
|
max_new_tokens=512,
|
|
do_sample=True,
|
|
temperature=0.6,
|
|
top_p=0.9,
|
|
use_cache=True,
|
|
)[0]
|
|
|
|
generate_ids = generate_ids[inputs["input_ids"].shape[1] :]
|
|
return processor.tokenizer.decode(
|
|
generate_ids, skip_special_tokens=True
|
|
).strip()
|
|
|
|
def _parse_tags(self, raw_output: str) -> List[str]:
|
|
"""Parse raw model output into a list of clean tags.
|
|
|
|
Args:
|
|
raw_output: Raw text output from the model
|
|
|
|
Returns:
|
|
List of cleaned, deduplicated tags
|
|
"""
|
|
# Patterns to exclude
|
|
exclude_prefixes = (
|
|
"copyright:",
|
|
"meta:",
|
|
"photo_",
|
|
"photo:",
|
|
)
|
|
|
|
# Split by common delimiters
|
|
tags = []
|
|
|
|
# Handle comma-separated tags
|
|
for part in raw_output.split(","):
|
|
tag = part.strip().lower()
|
|
# Remove any quotes or extra characters
|
|
tag = tag.strip("\"'")
|
|
# Replace spaces with underscores (Danbooru style)
|
|
tag = tag.replace(" ", "_")
|
|
# Remove empty tags
|
|
if tag and len(tag) > 1:
|
|
# Filter out unwanted tag patterns
|
|
if not tag.startswith(exclude_prefixes):
|
|
tags.append(tag)
|
|
|
|
# Deduplicate while preserving order
|
|
seen = set()
|
|
unique_tags = []
|
|
for tag in tags:
|
|
if tag not in seen:
|
|
seen.add(tag)
|
|
unique_tags.append(tag)
|
|
|
|
return unique_tags
|
|
|
|
|
|
# Singleton instance for API use
|
|
_service_instance: Optional[AutoTagService] = None
|
|
|
|
|
|
def get_autotag_service() -> AutoTagService:
|
|
"""Get or create the singleton AutoTagService instance."""
|
|
global _service_instance
|
|
if _service_instance is None:
|
|
_service_instance = AutoTagService()
|
|
return _service_instance
|