Files
ComfyAssets-ComfyUI_PromptM…/py/autotag.py
T

703 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