From ffaf5fd3d8dbb7b5f8b593bd46c6bdd990f889ed Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sun, 7 Dec 2025 17:24:54 -0800 Subject: [PATCH 1/7] feat(database): add method to retrieve all images with prompts --- database/operations.py | 66 ++++++++++++++++++++++++++++++++++++------ 1 file changed, 57 insertions(+), 9 deletions(-) diff --git a/database/operations.py b/database/operations.py index 39ed828..095388d 100644 --- a/database/operations.py +++ b/database/operations.py @@ -195,7 +195,7 @@ class PromptDatabase: params.extend([limit, offset]) query = " ".join(query_parts) - + with self.model.get_connection() as conn: cursor = conn.execute(query, params) rows = cursor.fetchall() @@ -368,12 +368,12 @@ class PromptDatabase: def get_all_tags(self) -> List[str]: """ Get all unique tags from the database. - + Returns: List of tag names """ all_tags = set() - + with self.model.get_connection() as conn: cursor = conn.execute("SELECT tags FROM prompts WHERE tags IS NOT NULL") for row in cursor.fetchall(): @@ -381,32 +381,43 @@ class PromptDatabase: tags = json.loads(row['tags']) if isinstance(tags, list): all_tags.update(tags) + elif isinstance(tags, str): + # Handle corrupted data (comma-separated string stored as JSON) + parsed_tags = [t.strip() for t in tags.split(',') if t.strip()] + all_tags.update(parsed_tags) except (json.JSONDecodeError, TypeError): continue - + return sorted(list(all_tags)) def _row_to_dict(self, row: sqlite3.Row) -> Dict[str, Any]: """ Convert a database row to a dictionary with parsed JSON fields. - + Args: row: SQLite row object - + Returns: Dictionary representation of the row """ data = dict(row) - + # Parse tags JSON if data.get('tags'): try: - data['tags'] = json.loads(data['tags']) + parsed = json.loads(data['tags']) + if isinstance(parsed, list): + data['tags'] = parsed + elif isinstance(parsed, str): + # Handle corrupted data (comma-separated string stored as JSON) + data['tags'] = [t.strip() for t in parsed.split(',') if t.strip()] + else: + data['tags'] = [] except (json.JSONDecodeError, TypeError): data['tags'] = [] else: data['tags'] = [] - + return data def export_prompts(self, file_path: str, format: str = "json") -> bool: @@ -715,6 +726,43 @@ class PromptDatabase: ) return [self._image_row_to_dict(row) for row in cursor.fetchall()] + def get_all_images(self) -> List[Dict[str, Any]]: + """ + Get all generated images with their linked prompts. + + Returns: + List of all image records with prompt text and tags + """ + with self.model.get_connection() as conn: + cursor = conn.execute( + """ + SELECT gi.*, p.text as prompt_text, p.tags as prompt_tags + FROM generated_images gi + LEFT JOIN prompts p ON gi.prompt_id = p.id + ORDER BY gi.generation_time DESC + """ + ) + rows = cursor.fetchall() + result = [] + for row in rows: + data = self._image_row_to_dict(row) + # Parse prompt_tags JSON + if row['prompt_tags']: + try: + tags = json.loads(row['prompt_tags']) + if isinstance(tags, list): + data['prompt_tags'] = tags + elif isinstance(tags, str): + data['prompt_tags'] = [t.strip() for t in tags.split(',') if t.strip()] + else: + data['prompt_tags'] = [] + except (json.JSONDecodeError, TypeError): + data['prompt_tags'] = [] + else: + data['prompt_tags'] = [] + result.append(data) + return result + def search_images_by_prompt(self, search_term: str) -> List[Dict[str, Any]]: """ Search images by prompt text. From 56eed388a8fa8f0658303f3928b5f1fec2dd2309 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sun, 7 Dec 2025 17:25:06 -0800 Subject: [PATCH 2/7] feat(api): add AutoTag endpoints and image routes --- py/api.py | 515 +++++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 511 insertions(+), 4 deletions(-) diff --git a/py/api.py b/py/api.py index 100b1d1..f177be1 100644 --- a/py/api.py +++ b/py/api.py @@ -374,6 +374,10 @@ class PromptManagerAPI: async def get_recent_images_route(request): return await self.get_recent_images(request) + @routes.get("/prompt_manager/images/all") + async def get_all_images_route(request): + return await self.get_all_images(request) + @routes.get("/prompt_manager/images/search") async def search_images_route(request): return await self.search_images(request) @@ -457,6 +461,31 @@ class PromptManagerAPI: async def get_log_stats_route(request): return await self.get_log_stats(request) + # AutoTag endpoints + @routes.get("/prompt_manager/autotag/models") + async def get_autotag_models_route(request): + return await self.get_autotag_models(request) + + @routes.get("/prompt_manager/autotag/download/{model_type}") + async def download_autotag_model_route(request): + return await self.download_autotag_model(request) + + @routes.get("/prompt_manager/autotag/start") + async def start_autotag_route(request): + return await self.start_autotag(request) + + @routes.post("/prompt_manager/autotag/single") + async def autotag_single_route(request): + return await self.autotag_single(request) + + @routes.post("/prompt_manager/autotag/apply") + async def apply_autotag_route(request): + return await self.apply_autotag(request) + + @routes.get("/prompt_manager/scan_output_dir") + async def scan_output_dir_route(request): + return await self.scan_output_dir(request) + self.logger.info("All routes registered with decorator pattern") async def search_prompts(self, request): @@ -1754,7 +1783,7 @@ class PromptManagerAPI: try: limit = int(request.query.get('limit', 50)) images = self.db.get_recent_images(limit) - + return web.json_response({ 'success': True, 'images': images @@ -1766,6 +1795,23 @@ class PromptManagerAPI: 'error': str(e) }, status=500) + async def get_all_images(self, request): + """Get all generated images with linked prompts.""" + try: + images = self.db.get_all_images() + + return web.json_response({ + 'success': True, + 'images': images, + 'count': len(images) + }) + except Exception as e: + self.logger.error(f"Get all images error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + async def search_images(self, request): """Search images by prompt text.""" try: @@ -4195,18 +4241,479 @@ class PromptManagerAPI: current_dir = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) sys.path.insert(0, current_dir) from utils.logging_config import get_logger_manager - + logger_manager = get_logger_manager() stats = logger_manager.get_log_stats() - + return web.json_response({ 'success': True, 'stats': stats }) - + except Exception as e: self.logger.error(f"Get log stats error: {e}") return web.json_response({ 'success': False, 'error': str(e) }, status=500) + + # ======================================================================== + # AutoTag Endpoints + # ======================================================================== + + async def get_autotag_models(self, request): + """ + Get status of available AutoTag models. + GET /prompt_manager/autotag/models + + Returns model availability, download status, and configuration. + """ + try: + from .autotag import get_autotag_service + + service = get_autotag_service() + models_status = service.get_models_status() + + return web.json_response({ + 'success': True, + 'models': models_status, + 'default_prompt': service.default_prompt + }) + + except Exception as e: + self.logger.error(f"Get autotag models error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def download_autotag_model(self, request): + """ + Download an AutoTag model with streaming progress. + POST /prompt_manager/autotag/download/{model_type} + + Streams SSE progress updates during download. + """ + import json + import asyncio + + model_type = request.match_info.get('model_type') + + async def stream_response(): + try: + from .autotag import get_autotag_service + + service = get_autotag_service() + + if model_type not in service.models_config: + yield f"data: {json.dumps({'type': 'error', 'message': f'Invalid model type: {model_type}'})}\n\n" + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Starting download...'})}\n\n" + + # Define progress callback + progress_data = {'last_progress': 0} + + def progress_callback(status: str, progress: float): + progress_data['last_progress'] = progress + + # Perform download (blocking in thread) + loop = asyncio.get_event_loop() + success = await loop.run_in_executor( + None, + lambda: service.download_model(model_type, progress_callback) + ) + + if success: + yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'status': 'Download complete'})}\n\n" + else: + yield f"data: {json.dumps({'type': 'error', 'message': 'Download failed'})}\n\n" + + except Exception as e: + self.logger.error(f"Download model error: {e}") + yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" + + response = web.StreamResponse( + status=200, + reason='OK', + headers={ + 'Content-Type': 'text/event-stream', + 'Cache-Control': 'no-cache', + 'Connection': 'keep-alive' + } + ) + + await response.prepare(request) + + async for chunk in stream_response(): + await response.write(chunk.encode('utf-8')) + + await response.write_eof() + return response + + async def start_autotag(self, request): + """ + Start batch auto-tagging with streaming progress. + GET /prompt_manager/autotag/start + + Query params: + model_type: "gguf" or "hf" + prompt: custom prompt text + + Streams SSE progress updates during processing. + """ + import json + import asyncio + from pathlib import Path + + # Read from query params (EventSource only supports GET) + model_type = request.query.get('model_type', 'gguf') + custom_prompt = request.query.get('prompt', '') + skip_tagged = request.query.get('skip_tagged', 'true').lower() == 'true' + use_gpu = True + + async def stream_response(): + try: + from .autotag import get_autotag_service + + service = get_autotag_service() + + # Check model is downloaded + status = service.get_models_status() + if not status.get(model_type, {}).get('downloaded'): + yield f"data: {json.dumps({'type': 'error', 'message': f'Model {model_type} not downloaded'})}\n\n" + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 0, 'status': 'Loading model...'})}\n\n" + + # Load model in thread pool + loop = asyncio.get_event_loop() + try: + await loop.run_in_executor( + None, + lambda: service.load_model(model_type, use_gpu) + ) + except Exception as e: + yield f"data: {json.dumps({'type': 'error', 'message': f'Failed to load model: {str(e)}'})}\n\n" + return + + # Set custom prompt if provided + if custom_prompt: + service.custom_prompt = custom_prompt + + yield f"data: {json.dumps({'type': 'progress', 'progress': 5, 'status': 'Model loaded. Fetching all images from database...'})}\n\n" + + # Get ALL images from database (only images with linked prompts) + images = self.db.get_all_images() + + total_files = len(images) + if total_files == 0: + yield f"data: {json.dumps({'type': 'complete', 'processed': 0, 'tagged': 0, 'skipped': 0, 'status': 'No images with linked prompts found'})}\n\n" + service.unload_model() + return + + yield f"data: {json.dumps({'type': 'progress', 'progress': 10, 'status': f'Found {total_files} images. Processing...'})}\n\n" + + processed = 0 + tagged = 0 + skipped = 0 + errors = 0 + + for i, image_data in enumerate(images): + image_path = image_data.get('image_path') + prompt_id = image_data.get('prompt_id') + + if not image_path or not prompt_id: + skipped += 1 + continue + + # Check if file exists + if not Path(image_path).exists(): + skipped += 1 + continue + + # Check if image already has real tags (skip_tagged option) + if skip_tagged: + prompt_tags = image_data.get('prompt_tags', []) + if isinstance(prompt_tags, str): + prompt_tags = [t.strip() for t in prompt_tags.split(',') if t.strip()] + # Filter out "auto-scanned" - it's not a real tag + real_tags = [t for t in prompt_tags if t != 'auto-scanned'] + if real_tags: + skipped += 1 + if i % 5 == 0 or i == total_files - 1: + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Skipping {i+1}/{total_files} (already tagged)...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + continue + + try: + # Generate tags + tags = await loop.run_in_executor( + None, + lambda p=str(image_path): service.generate_tags(p) + ) + + processed += 1 + + if tags: + # Get existing prompt + existing_prompt = self.db.get_prompt_by_id(prompt_id) + if existing_prompt: + existing_tags = existing_prompt.get('tags', []) + if isinstance(existing_tags, str): + existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()] + + new_tags = [t for t in tags if t not in existing_tags] + if new_tags: + all_tags = existing_tags + new_tags + self.db.update_prompt_metadata( + prompt_id, + tags=all_tags + ) + tagged += 1 + else: + skipped += 1 + else: + skipped += 1 + else: + skipped += 1 + + if i % 5 == 0 or i == total_files - 1: + progress = 10 + int((i + 1) / total_files * 85) + yield f"data: {json.dumps({'type': 'progress', 'progress': progress, 'status': f'Processing {i+1}/{total_files}...', 'processed': processed, 'tagged': tagged, 'skipped': skipped})}\n\n" + await asyncio.sleep(0.01) + + except Exception as img_err: + self.logger.error(f"Error processing {image_path}: {img_err}") + errors += 1 + processed += 1 + + service.unload_model() + + yield f"data: {json.dumps({'type': 'complete', 'progress': 100, 'processed': processed, 'tagged': tagged, 'skipped': skipped, 'errors': errors, 'status': 'Complete'})}\n\n" + + except Exception as e: + self.logger.error(f"AutoTag error: {e}") + import traceback + self.logger.error(f"AutoTag traceback: {traceback.format_exc()}") + yield f"data: {json.dumps({'type': 'error', 'message': str(e)})}\n\n" + + response = web.StreamResponse( + status=200, + reason='OK', + headers={ + 'Content-Type': 'text/event-stream', + 'Cache-Control': 'no-cache', + 'Connection': 'keep-alive' + } + ) + + await response.prepare(request) + + async for chunk in stream_response(): + await response.write(chunk.encode('utf-8')) + + await response.write_eof() + return response + + async def autotag_single(self, request): + """ + Generate tags for a single image (for Review mode). + POST /prompt_manager/autotag/single + + Request body: + { + "image_path": "/path/to/image.png", + "model_type": "gguf", + "prompt": "optional custom prompt", + "use_gpu": true + } + """ + try: + data = await request.json() + image_path = data.get('image_path') + model_type = data.get('model_type', 'gguf') + custom_prompt = data.get('prompt') + use_gpu = data.get('use_gpu', True) + + if not image_path: + return web.json_response({ + 'success': False, + 'error': 'image_path is required' + }, status=400) + + from .autotag import get_autotag_service + import asyncio + + service = get_autotag_service() + + if not service.is_model_loaded() or service.get_loaded_model_type() != model_type: + loop = asyncio.get_event_loop() + await loop.run_in_executor( + None, + lambda: service.load_model(model_type, use_gpu) + ) + + if custom_prompt: + service.custom_prompt = custom_prompt + + loop = asyncio.get_event_loop() + tags = await loop.run_in_executor( + None, + lambda: service.generate_tags(image_path) + ) + + prompt_id = None + try: + with self.db.model.get_connection() as conn: + cursor = conn.execute( + "SELECT prompt_id FROM generated_images WHERE image_path = ?", + (image_path,) + ) + row = cursor.fetchone() + if row: + prompt_id = row[0] + except Exception as e: + self.logger.warning(f"Could not find linked prompt: {e}") + + return web.json_response({ + 'success': True, + 'tags': tags, + 'prompt_id': prompt_id, + 'image_path': image_path + }) + + except Exception as e: + self.logger.error(f"AutoTag single error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def apply_autotag(self, request): + """ + Apply selected tags to a prompt. + POST /prompt_manager/autotag/apply + + Request body: + { + "prompt_id": 123, + "tags": ["tag1", "tag2", ...] + } + """ + try: + data = await request.json() + prompt_id = data.get('prompt_id') + tags = data.get('tags', []) + + if not prompt_id: + return web.json_response({ + 'success': False, + 'error': 'prompt_id is required' + }, status=400) + + if not tags: + return web.json_response({ + 'success': True, + 'message': 'No tags to apply' + }) + + prompt = self.db.get_prompt_by_id(prompt_id) + if not prompt: + return web.json_response({ + 'success': False, + 'error': f'Prompt {prompt_id} not found' + }, status=404) + + existing_tags = prompt.get('tags', []) + if isinstance(existing_tags, str): + existing_tags = [t.strip() for t in existing_tags.split(',') if t.strip()] + + new_tags = [t for t in tags if t not in existing_tags] + all_tags = existing_tags + new_tags + + self.db.update_prompt_metadata( + prompt_id, + tags=all_tags + ) + + return web.json_response({ + 'success': True, + 'added_tags': new_tags, + 'total_tags': len(all_tags) + }) + + except Exception as e: + self.logger.error(f"Apply autotag error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) + + async def scan_output_dir(self, request): + """ + Scan ComfyUI output directory for images. + GET /prompt_manager/scan_output_dir + + Returns a list of images in the output directory for autotag review mode. + """ + from pathlib import Path + + try: + output_dir = self._find_comfyui_output_dir() + if not output_dir: + return web.json_response({ + 'success': False, + 'error': 'ComfyUI output directory not found' + }, status=404) + + output_path = Path(output_dir) + image_extensions = ['.png', '.jpg', '.jpeg', '.gif', '.webp', '.bmp', '.tiff'] + + images = [] + seen_paths = set() + + for ext in image_extensions: + for pattern in [f"*{ext.lower()}", f"*{ext.upper()}"]: + for image_path in output_path.rglob(pattern): + if 'thumbnails' not in image_path.parts: + normalized_path = str(image_path).lower() + if normalized_path not in seen_paths: + seen_paths.add(normalized_path) + + rel_path = image_path.relative_to(output_path) + + # Check for thumbnail + thumbnail_url = None + thumbnails_dir = output_path / "thumbnails" + if thumbnails_dir.exists(): + thumbnail_path = thumbnails_dir / f"{image_path.stem}_thumb{image_path.suffix}" + if thumbnail_path.exists(): + thumbnail_url = f'/prompt_manager/images/serve/thumbnails/{image_path.stem}_thumb{image_path.suffix}' + + images.append({ + 'filename': image_path.name, + 'path': str(image_path), + 'relative_path': str(rel_path), + 'url': f'/prompt_manager/images/serve/{rel_path.as_posix()}', + 'thumbnail_url': thumbnail_url + }) + + # Sort by filename + images.sort(key=lambda x: x['filename']) + + self.logger.info(f"Found {len(images)} images in output directory") + + return web.json_response({ + 'success': True, + 'images': images, + 'count': len(images) + }) + + except Exception as e: + self.logger.error(f"Scan output dir error: {e}") + return web.json_response({ + 'success': False, + 'error': str(e) + }, status=500) From ba36814894736b677b6e5673329f279f624b0b6a Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sun, 7 Dec 2025 17:25:17 -0800 Subject: [PATCH 3/7] feat(admin): add auto-tagging feature to UI --- web/admin.html | 911 ++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 897 insertions(+), 14 deletions(-) diff --git a/web/admin.html b/web/admin.html index afebffb..dc199a2 100644 --- a/web/admin.html +++ b/web/admin.html @@ -236,10 +236,14 @@ class="px-4 py-2 bg-amber-600 hover:bg-amber-700 text-white text-sm font-medium rounded-lg transition-colors"> 📁 Restore DB - + + + + +
+ +
+ + +
+
+ + +
+ + +
+ + +
+ +
+ + +
+

Note: "auto-scanned" is not counted as a real tag

+
+ + +
+ + + +
+ + + + + + + + + + + + + + + + + + + + \ No newline at end of file From c11d98fa17d5b92e519f22698133140d39cd360e Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sun, 7 Dec 2025 17:25:43 -0800 Subject: [PATCH 5/7] feat(autotag): add AutoTagService for image tagging --- py/autotag.py | 702 ++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 702 insertions(+) create mode 100644 py/autotag.py diff --git a/py/autotag.py b/py/autotag.py new file mode 100644 index 0000000..d74a0f9 --- /dev/null +++ b/py/autotag.py @@ -0,0 +1,702 @@ +""" +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=["", "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 From 3033d088b47a65b92a551cfe2fb40ce31b996a97 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sun, 7 Dec 2025 17:37:52 -0800 Subject: [PATCH 6/7] feat(AI AutoTag): Add auto-tagging for image collections --- README.md | 60 ++++++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 59 insertions(+), 1 deletion(-) diff --git a/README.md b/README.md index 18d5490..79af3ea 100644 --- a/README.md +++ b/README.md @@ -56,6 +56,7 @@ Both nodes include the complete PromptManager feature set: - **🔬 Workflow Analysis**: Extract and analyze ComfyUI workflow data from PNG images - **📋 Metadata Viewer**: Standalone tool for analyzing ComfyUI-generated images - **🛠️ System Management**: Built-in diagnostics, backup/restore, and maintenance tools +- **🏷️ AI AutoTag**: Automatically tag your image collection using JoyCaption vision models ![Image Gallery](images/pm-02.png) @@ -363,6 +364,53 @@ Import existing ComfyUI images into your database: - Any errors or issues encountered 5. **Access imported data** through the normal gallery interface +### 🏷️ AI AutoTag + +Automatically tag your entire image collection using JoyCaption vision models: + +#### **Model Options** + +- **JoyCaption Beta One FP16**: Full precision model for highest quality tagging (requires more VRAM) +- **JoyCaption Beta One GGUF (FP8)**: Quantized model for lower VRAM usage with minimal quality loss + +#### **Two Tagging Modes** + +1. **AutoTag (Batch Mode)**: Tag your entire collection automatically + - Click **"🏷️ AutoTag"** in the admin dashboard + - Choose your model type (FP16 or GGUF) + - Select how to handle already-tagged images: + - **Skip images with existing tags**: Preserve your meticulous manual tagging work + - **Re-tag all images**: Overwrite existing tags with fresh AI analysis + - Customize the system prompt to match your collection style + - Monitor real-time progress as images are processed + +2. **Review Mode**: Tag images one-by-one with approval + - Review each image and its AI-generated tags before applying + - Edit, add, or remove tags before saving + - Skip images you don't want to tag + - Perfect for curating high-quality tag assignments + +#### **Handling Existing Tags** + +For users who have meticulously tagged their collections from the start: +- AutoTag respects your existing work with the **"Skip tagged"** option +- **"auto-scanned"** placeholder tags don't count as real tags +- Choose to re-tag specific images while preserving others +- Review mode asks before overwriting on each image + +#### **Customization** + +Adjust the system prompt to match your tagging style: +- Focus on specific attributes (style, mood, composition, subjects) +- Match your existing tag vocabulary +- Optimize for your collection's theme (anime, photography, landscapes, etc.) + +#### **Requirements** + +- Sufficient VRAM for the chosen model (FP16 requires more, GGUF is lighter) +- Models are downloaded automatically on first use +- Feature requests and improvements are welcome! + ### 🌐 Web Interface Features The comprehensive web interface provides: @@ -770,13 +818,23 @@ MIT License - see LICENSE file for details. ### Integration Ideas -- **Auto-tagging**: Use AI to automatically categorize prompts - **Workflow linking**: Connect prompts to specific workflow templates - **Image analysis**: Analyze generated images to improve suggestions - **Version control**: Track prompt iterations and effectiveness ## Changelog +### v3.0.23 (AI AutoTag Release) + +- **🏷️ AI AutoTag**: Automatically tag your entire image collection using JoyCaption vision models +- **🤖 Dual Model Support**: Choose between JoyCaption Beta One FP16 (high quality) or GGUF FP8 (lower VRAM) +- **📋 Two Tagging Modes**: Batch mode for bulk tagging, Review mode for one-by-one approval +- **⚙️ Smart Tag Handling**: Skip already-tagged images or re-tag all - respects your existing work +- **✏️ Customizable Prompts**: Adjust the system prompt to match your collection style and tag vocabulary +- **🔄 Real-time Progress**: SSE-based streaming progress updates during batch operations +- **📦 Tag Accordion UI**: Collapsible tag display shows first row with expandable section for large tag sets +- **🎯 Review Confirmation**: Prompts before overwriting existing tags in review mode + ### v3.0.10 (Batch Processing Node) - **📋 PromptSearchList Node**: New node that searches the prompt database and outputs results as a list From b8d6b9072d984a95300c05f6d5ca4df6e88c0263 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Sun, 7 Dec 2025 17:38:04 -0800 Subject: [PATCH 7/7] chore(pyproject): bump version to 3.0.23 --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 1edacdf..b71a258 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "promptmanager" description = "A powerful ComfyUI custom node that extends the standard text encoder with persistent prompt storage, advanced search capabilities, and an automatic image gallery system using SQLite." -version = "3.0.22" +version = "3.0.23" license = {file = "LICENSE"} dependencies = ["# Core dependencies for PromptManager", "# Note: Most dependencies are already included with ComfyUI", "# Already included with Python standard library:", "# - sqlite3", "# - hashlib", "# - json", "# - datetime", "# - os", "# - typing", "# - threading", "# - uuid", "# Required for gallery functionality:", "watchdog>=2.1.0 # For file system monitoring", "Pillow>=8.0.0 # For image metadata extraction (usually included with ComfyUI)", "# Optional dependencies for enhanced search functionality:", "# fuzzywuzzy[speedup]>=0.18.0 # For fuzzy string matching (optional)", "# sqlalchemy>=1.4.0 # For advanced ORM features (optional)", "# Development dependencies (optional):", "# pytest>=6.0.0 # For running tests", "# black>=22.0.0 # For code formatting", "# flake8>=4.0.0 # For linting", "# mypy>=0.910 # For type checking"]