diff --git a/modules/interrogate/deepseek.py b/modules/interrogate/deepseek.py
index b2d340248..8611efb82 100644
--- a/modules/interrogate/deepseek.py
+++ b/modules/interrogate/deepseek.py
@@ -18,32 +18,31 @@ from modules import shared, devices, paths, sd_models
# model_path = "deepseek-ai/deepseek-vl2-small"
vl_gpt = None
vl_chat_processor = None
+loaded_repo = None
class fake_attrdict():
- class AttrDict(dict): # dot notation access to dictionary attributes
+ class AttrDict(dict): # dot notation access to dictionary attributes
__getattr__ = dict.get
__setattr__ = dict.__setitem__
__delattr__ = dict.__delitem__
-# def fake_is_flash_attn_2_available():
-# return False
-
-def predict(question, image, repo):
- global vl_gpt, vl_chat_processor # pylint: disable=global-statement
+def load(repo: str):
+ """Load DeepSeek VL2 model (experimental)."""
+ global vl_gpt, vl_chat_processor, loaded_repo # pylint: disable=global-statement
if not shared.cmd_opts.experimental:
shared.log.error(f'Interrogate: type=vlm model="DeepSeek VL2" repo="{repo}" is experimental-only')
- return ''
+ return False
folder = os.path.join(paths.script_path, 'repositories', 'deepseek-vl2')
if not os.path.exists(folder):
shared.log.error(f'Interrogate: type=vlm model="DeepSeek VL2" repo="{repo}" deepseek-vl2 repo not found')
- return ''
- if vl_gpt is None:
+ return False
+ if vl_gpt is None or loaded_repo != repo:
sys.modules['attrdict'] = fake_attrdict
from transformers.models.llama import modeling_llama
modeling_llama.LlamaFlashAttention2 = modeling_llama.LlamaAttention
- _deekseek_vl = importlib.import_module('repositories.deepseek-vl2.deepseek_vl2')
+ importlib.import_module('repositories.deepseek-vl2.deepseek_vl2')
deekseek_vl_models = importlib.import_module('repositories.deepseek-vl2.deepseek_vl2.models')
vl_chat_processor = deekseek_vl_models.DeepseekVLV2Processor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
vl_gpt = AutoModelForCausalLM.from_pretrained(
@@ -51,7 +50,31 @@ def predict(question, image, repo):
trust_remote_code=True,
cache_dir=shared.opts.hfcache_dir,
)
- vl_gpt = vl_gpt.to(device=devices.device, dtype=devices.dtype).eval()
+ vl_gpt.to(dtype=devices.dtype)
+ vl_gpt.eval()
+ loaded_repo = repo
+ shared.log.info(f'Interrogate: type=vlm model="DeepSeek VL2" repo="{repo}"')
+ sd_models.move_model(vl_gpt, devices.device)
+ return True
+
+
+def unload():
+ """Release DeepSeek VL2 model from GPU/memory."""
+ global vl_gpt, vl_chat_processor, loaded_repo # pylint: disable=global-statement
+ if vl_gpt is not None:
+ shared.log.debug(f'DeepSeek unload: model="{loaded_repo}"')
+ sd_models.move_model(vl_gpt, devices.cpu, force=True)
+ vl_gpt = None
+ vl_chat_processor = None
+ loaded_repo = None
+ devices.torch_gc(force=True)
+ else:
+ shared.log.debug('DeepSeek unload: no model loaded')
+
+
+def predict(question, image, repo):
+ if not load(repo):
+ return ''
if len(question) < 2:
question = "Describe the image."
diff --git a/modules/interrogate/joycaption.py b/modules/interrogate/joycaption.py
index 114888f4e..c8d445d9e 100644
--- a/modules/interrogate/joycaption.py
+++ b/modules/interrogate/joycaption.py
@@ -57,26 +57,43 @@ llava_model: LlavaForConditionalGeneration = None
opts = JoyOptions()
-@torch.no_grad()
-def predict(question: str, image, vqa_model: str = None) -> str:
- global llava_model, processor # pylint: disable=global-statement
- opts.max_new_tokens = shared.opts.interrogate_vlm_max_length
- if vqa_model is not None and opts.repo != vqa_model:
- opts.repo = vqa_model
+def load(repo: str = None):
+ """Load JoyCaption model."""
+ global llava_model, processor # pylint: disable=global-statement
+ repo = repo or opts.repo
+ if llava_model is None or opts.repo != repo:
+ opts.repo = repo
llava_model = None
- if llava_model is None:
shared.log.info(f'Interrogate: type=vlm model="JoyCaption" {str(opts)}')
-
- processor = AutoProcessor.from_pretrained(opts.repo, max_pixels=1024*1024)
+ processor = AutoProcessor.from_pretrained(repo, max_pixels=1024*1024, cache_dir=shared.opts.hfcache_dir)
quant_args = model_quant.create_config(module='LLM')
llava_model = LlavaForConditionalGeneration.from_pretrained(
- opts.repo,
+ repo,
torch_dtype=devices.dtype,
- device_map="auto",
cache_dir=shared.opts.hfcache_dir,
**quant_args,
)
llava_model.eval()
+ sd_models.move_model(llava_model, devices.device)
+
+
+def unload():
+ """Release JoyCaption model from GPU/memory."""
+ global llava_model, processor # pylint: disable=global-statement
+ if llava_model is not None:
+ shared.log.debug(f'JoyCaption unload: model="{opts.repo}"')
+ sd_models.move_model(llava_model, devices.cpu, force=True)
+ llava_model = None
+ processor = None
+ devices.torch_gc(force=True)
+ else:
+ shared.log.debug('JoyCaption unload: no model loaded')
+
+
+@torch.no_grad()
+def predict(question: str, image, vqa_model: str = None) -> str:
+ opts.max_new_tokens = shared.opts.interrogate_vlm_max_length
+ load(vqa_model)
if len(question) < 2:
question = "Describe the image."
@@ -86,9 +103,8 @@ def predict(question: str, image, vqa_model: str = None) -> str:
{ "role": "user", "content": question },
]
convo_string = processor.apply_chat_template(convo, tokenize=False, add_generation_prompt=True)
- inputs = processor(text=[convo_string], images=[image], return_tensors="pt").to(devices.device) # Process the inputs
+ inputs = processor(text=[convo_string], images=[image], return_tensors="pt").to(devices.device)
inputs['pixel_values'] = inputs['pixel_values'].to(devices.dtype)
- sd_models.move_model(llava_model, devices.device)
with devices.inference_context():
generate_ids = llava_model.generate( # Generate the captions
**inputs,
diff --git a/modules/interrogate/joytag.py b/modules/interrogate/joytag.py
index ae48cb9c6..f78e49ae5 100644
--- a/modules/interrogate/joytag.py
+++ b/modules/interrogate/joytag.py
@@ -16,7 +16,7 @@ import torchvision.transforms.functional as TVF
import einops
from einops.layers.torch import Rearrange
import huggingface_hub
-from modules import shared, devices
+from modules import shared, devices, sd_models
model = None
@@ -1034,22 +1034,39 @@ def prepare_image(image: Image.Image, target_size: int) -> torch.Tensor:
return image_tensor
-def predict(image: Image.Image):
- global model, tags # pylint: disable=global-statement
+def load():
+ """Load JoyTag model."""
+ global model, tags # pylint: disable=global-statement
if model is None:
folder = huggingface_hub.snapshot_download(MODEL_REPO, cache_dir=shared.opts.hfcache_dir)
model = VisionModel.load_model(folder)
- model = model.to(device=devices.device, dtype=devices.dtype)
+ model.to(dtype=devices.dtype)
model.eval()
with open(os.path.join(folder, 'top_tags.txt'), 'r', encoding='utf8') as f:
tags = [line.strip() for line in f.readlines() if line.strip()]
- shared.log.info(f'Interrogate: type=vlm model="JoyCaption" repo="{MODEL_REPO}" tags={len(tags)}')
+ shared.log.info(f'Interrogate: type=vlm model="JoyTag" repo="{MODEL_REPO}" tags={len(tags)}')
+ sd_models.move_model(model, devices.device)
+
+
+def unload():
+ """Release JoyTag model from GPU/memory."""
+ global model, tags # pylint: disable=global-statement
+ if model is not None:
+ shared.log.debug('JoyTag unload')
+ sd_models.move_model(model, devices.cpu, force=True)
+ model = None
+ tags = None
+ devices.torch_gc(force=True)
+ else:
+ shared.log.debug('JoyTag unload: no model loaded')
+
+
+def predict(image: Image.Image):
+ load()
image_tensor = prepare_image(image, model.image_size).unsqueeze(0).to(device=devices.device, dtype=devices.dtype)
- model = model.to(devices.device)
with devices.inference_context():
- preds = model({ 'image': image_tensor })
+ preds = model({'image': image_tensor})
tag_preds = preds['tags'].sigmoid().cpu()
- model = model.to(devices.cpu)
scores = {tags[i]: tag_preds[0][i] for i in range(len(tags))}
if shared.opts.interrogate_score:
predicted_tags = [f'{tag}:{score:.2f}' for tag, score in scores.items() if score > THRESHOLD]
diff --git a/modules/interrogate/moondream3.py b/modules/interrogate/moondream3.py
new file mode 100644
index 000000000..c1e10e4a9
--- /dev/null
+++ b/modules/interrogate/moondream3.py
@@ -0,0 +1,410 @@
+# Moondream 3 Preview VLM Implementation
+# Source: https://huggingface.co/moondream/moondream3-preview
+# Model: 9.3GB, gated (requires HuggingFace authentication)
+# Architecture: Mixture-of-Experts (9B total params, 2B active)
+import os
+import re
+import torch
+import transformers
+from PIL import Image
+from modules import shared, devices, sd_models
+from modules.interrogate import vqa_detection
+
+
+# Debug logging - function-based to avoid circular import
+debug_enabled = os.environ.get('SD_VQA_DEBUG', None) is not None
+
+def debug(*args, **kwargs):
+ if debug_enabled:
+ shared.log.trace(*args, **kwargs)
+
+
+# Global state
+moondream3_model = None
+loaded = None
+image_cache = {} # Cache encoded images for reuse
+
+
+def get_settings():
+ """
+ Build settings dict for Moondream 3 API from global VQA options.
+ Moondream 3 accepts: temperature, top_p, max_tokens
+ """
+ settings = {}
+ if shared.opts.interrogate_vlm_max_length > 0:
+ settings['max_tokens'] = shared.opts.interrogate_vlm_max_length
+ if shared.opts.interrogate_vlm_temperature > 0:
+ settings['temperature'] = shared.opts.interrogate_vlm_temperature
+ if shared.opts.interrogate_vlm_top_p > 0:
+ settings['top_p'] = shared.opts.interrogate_vlm_top_p
+ return settings if settings else None
+
+
+def load_model(repo: str):
+ """Load Moondream 3 model."""
+ global moondream3_model, loaded # pylint: disable=global-statement
+
+ if moondream3_model is None or loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ moondream3_model = None
+
+ moondream3_model = transformers.AutoModelForCausalLM.from_pretrained(
+ repo,
+ trust_remote_code=True,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ )
+
+ moondream3_model.eval()
+
+ # Initialize KV caches before moving to device (they're lazy by default)
+ if hasattr(moondream3_model, '_setup_caches'):
+ moondream3_model._setup_caches()
+
+ # Disable flex_attention decoding (can cause hangs due to torch.compile)
+ if hasattr(moondream3_model, 'model') and hasattr(moondream3_model.model, 'use_flex_decoding'):
+ moondream3_model.model.use_flex_decoding = False
+
+ loaded = repo
+ devices.torch_gc()
+
+ # Move model to active device
+ sd_models.move_model(moondream3_model, devices.device)
+ return moondream3_model
+
+
+def encode_image(image: Image.Image, cache_key: str = None):
+ """
+ Encode image for reuse across multiple queries.
+
+ Args:
+ image: PIL Image
+ cache_key: Optional cache key for storing encoded image
+
+ Returns:
+ Encoded image tensor
+ """
+ if cache_key and cache_key in image_cache:
+ debug(f'VQA interrogate: handler=moondream3 using cached encoding for cache_key="{cache_key}"')
+ return image_cache[cache_key]
+
+ model = load_model(loaded)
+
+ with devices.inference_context():
+ encoded = model.encode_image(image)
+
+ if cache_key:
+ image_cache[cache_key] = encoded
+ debug(f'VQA interrogate: handler=moondream3 cached encoding cache_key="{cache_key}" cache_size={len(image_cache)}')
+
+ return encoded
+
+
+def query(image: Image.Image, question: str, repo: str, stream: bool = False,
+ temperature: float = None, top_p: float = None, max_tokens: int = None,
+ use_cache: bool = False, reasoning: bool = True):
+ """
+ Visual question answering with optional streaming.
+
+ Args:
+ image: PIL Image
+ question: Question about the image
+ repo: Model repository
+ stream: Enable streaming output (generator)
+ temperature: Sampling temperature (overrides global setting)
+ top_p: Nucleus sampling parameter (overrides global setting)
+ max_tokens: Maximum tokens to generate (overrides global setting)
+ use_cache: Use cached image encoding if available
+
+ Returns:
+ Answer dict or string (or generator if stream=True)
+ """
+ model = load_model(repo)
+
+ # Build settings - per-call parameters override global settings
+ settings = get_settings() or {}
+ if temperature is not None:
+ settings['temperature'] = temperature
+ if top_p is not None:
+ settings['top_p'] = top_p
+ if max_tokens is not None:
+ settings['max_tokens'] = max_tokens
+
+ debug(f'VQA interrogate: handler=moondream3 method=query question="{question}" stream={stream} settings={settings}')
+
+ # Use cached encoding if requested
+ if use_cache:
+ cache_key = f"{id(image)}_{question}"
+ image_input = encode_image(image, cache_key)
+ else:
+ image_input = image
+
+ with devices.inference_context():
+ response = model.query(
+ image=image_input,
+ question=question,
+ stream=stream,
+ settings=settings if settings else None,
+ reasoning=reasoning
+ )
+
+ # Log response structure (for non-streaming)
+ if not stream:
+ if isinstance(response, dict):
+ debug(f'VQA interrogate: handler=moondream3 response_type=dict keys={list(response.keys())}')
+ if 'reasoning' in response:
+ reasoning_text = response['reasoning'].get('text', '')[:100] + '...' if len(response['reasoning'].get('text', '')) > 100 else response['reasoning'].get('text', '')
+ debug(f'VQA interrogate: handler=moondream3 reasoning="{reasoning_text}"')
+ if 'answer' in response:
+ debug(f'VQA interrogate: handler=moondream3 answer="{response["answer"]}"')
+
+ return response
+
+
+def caption(image: Image.Image, repo: str, length: str = 'normal', stream: bool = False,
+ temperature: float = None, top_p: float = None, max_tokens: int = None):
+ """
+ Generate image captions at different lengths.
+
+ Args:
+ image: PIL Image
+ repo: Model repository
+ length: Caption length - 'short', 'normal', or 'long'
+ stream: Enable streaming output (generator)
+ temperature: Sampling temperature (overrides global setting)
+ top_p: Nucleus sampling parameter (overrides global setting)
+ max_tokens: Maximum tokens to generate (overrides global setting)
+
+ Returns:
+ Caption dict or string (or generator if stream=True)
+ """
+ model = load_model(repo)
+
+ # Build settings - per-call parameters override global settings
+ settings = get_settings() or {}
+ if temperature is not None:
+ settings['temperature'] = temperature
+ if top_p is not None:
+ settings['top_p'] = top_p
+ if max_tokens is not None:
+ settings['max_tokens'] = max_tokens
+
+ debug(f'VQA interrogate: handler=moondream3 method=caption length={length} stream={stream} settings={settings}')
+
+ with devices.inference_context():
+ response = model.caption(
+ image,
+ length=length,
+ stream=stream,
+ settings=settings if settings else None
+ )
+
+ # Log response structure (for non-streaming)
+ if not stream and isinstance(response, dict):
+ debug(f'VQA interrogate: handler=moondream3 response_type=dict keys={list(response.keys())}')
+
+ return response
+
+
+def point(image: Image.Image, object_name: str, repo: str):
+ """
+ Identify coordinates of all instances of a specific object in the image.
+
+ Args:
+ image: PIL Image
+ object_name: Name of object to locate
+ repo: Model repository
+
+ Returns:
+ List of (x, y) tuples with coordinates normalized to 0-1 range, or None if not found
+ Example: [(0.733, 0.442), (0.5, 0.6)] for 2 instances
+ """
+ model = load_model(repo)
+
+ debug(f'VQA interrogate: handler=moondream3 method=point object_name="{object_name}"')
+
+ with devices.inference_context():
+ result = model.point(image, object_name)
+
+ debug(f'VQA interrogate: handler=moondream3 point_raw_result="{result}" type={type(result)}')
+ if isinstance(result, dict):
+ debug(f'VQA interrogate: handler=moondream3 point_raw_result_keys={list(result.keys())}')
+
+ points = vqa_detection.parse_points(result)
+ if points:
+ debug(f'VQA interrogate: handler=moondream3 point_result={len(points)} points found')
+ return points
+
+ debug('VQA interrogate: handler=moondream3 point_result=not found')
+ return None
+
+
+def detect(image: Image.Image, object_name: str, repo: str, max_objects: int = 10):
+ """
+ Detect all instances of a specific object with bounding boxes.
+
+ Args:
+ image: PIL Image
+ object_name: Name of object to detect
+ repo: Model repository
+ max_objects: Maximum number of objects to return
+
+ Returns:
+ List of detection dicts with keys:
+ - 'bbox': [x1, y1, x2, y2] normalized to 0-1
+ - 'label': Object label
+ - 'confidence': Detection confidence (0-1)
+ Returns empty list if no objects found.
+ """
+ model = load_model(repo)
+
+ debug(f'VQA interrogate: handler=moondream3 method=detect object_name="{object_name}" max_objects={max_objects}')
+
+ with devices.inference_context():
+ result = model.detect(image, object_name)
+
+ debug(f'VQA interrogate: handler=moondream3 detect_raw_result="{result}" type={type(result)}')
+ if isinstance(result, dict):
+ debug(f'VQA interrogate: handler=moondream3 detect_raw_result_keys={list(result.keys())}')
+
+ detections = vqa_detection.parse_detections(result, object_name, max_objects)
+ debug(f'VQA interrogate: handler=moondream3 detect_result={len(detections)} objects found')
+ return detections
+
+
+def predict(question: str, image: Image.Image, repo: str, model_name: str = None, thinking_mode: bool = False,
+ mode: str = None, stream: bool = False, use_cache: bool = False, **kwargs):
+ """
+ Main entry point for Moondream 3 VQA - auto-detects mode from question.
+
+ Args:
+ question: The question/prompt (e.g., "caption", "where is the cat?", "describe this")
+ image: PIL Image
+ repo: Model repository
+ model_name: Display name for logging
+ thinking_mode: Enable reasoning mode for query
+ mode: Force specific mode ('query', 'caption', 'caption_short', 'caption_long', 'point', 'detect')
+ stream: Enable streaming output (for query/caption)
+ use_cache: Use cached image encoding (for query)
+ **kwargs: Additional parameters (max_objects for detect, etc.)
+
+ Returns:
+ Response string (detection data stored on VQA singleton instance.last_detection_data)
+ (or generator if stream=True for query/caption modes)
+ """
+ debug(f'VQA interrogate: handler=moondream3 model_name="{model_name}" repo="{repo}" question="{question}" image_size={image.size if image else None} mode={mode} stream={stream}')
+
+ # Clean question
+ question = question.replace('<', '').replace('>', '').replace('_', ' ') if question else ''
+
+ # Auto-detect mode from question if not specified
+ if mode is None:
+ question_lower = question.lower()
+
+ # Caption detection
+ if question in ['CAPTION', 'caption'] or 'caption' in question_lower:
+ if 'more detailed' in question_lower or 'very long' in question_lower:
+ mode = 'caption_long'
+ elif 'detailed' in question_lower or 'long' in question_lower:
+ mode = 'caption_normal'
+ elif 'short' in question_lower or 'brief' in question_lower:
+ mode = 'caption_short'
+ else:
+ # Default caption mode (matches vqa.py legacy behavior)
+ if question == 'CAPTION':
+ mode = 'caption_short'
+ elif question == 'DETAILED CAPTION':
+ mode = 'caption_normal'
+ elif question == 'MORE DETAILED CAPTION':
+ mode = 'caption_long'
+ else:
+ mode = 'caption_normal'
+
+ # Point detection
+ elif 'where is' in question_lower or 'locate' in question_lower or 'find' in question_lower or 'point' in question_lower:
+ mode = 'point'
+
+ # Object detection
+ elif 'detect' in question_lower or 'bounding box' in question_lower or 'bbox' in question_lower:
+ mode = 'detect'
+
+ # Default to query
+ else:
+ mode = 'query'
+
+ debug(f'VQA interrogate: handler=moondream3 mode_selected={mode}')
+
+ # Dispatch to appropriate method
+ try:
+ if mode == 'caption_short':
+ response = caption(image, repo, length='short', stream=stream)
+ elif mode == 'caption_long':
+ response = caption(image, repo, length='long', stream=stream)
+ elif mode in ['caption', 'caption_normal']:
+ response = caption(image, repo, length='normal', stream=stream)
+ elif mode == 'point':
+ # Extract object name from question - case insensitive, preserve object names
+ object_name = question
+ for phrase in ['point at', 'where is', 'locate', 'find']:
+ object_name = re.sub(rf'\b{phrase}\b', '', object_name, flags=re.IGNORECASE)
+ object_name = re.sub(r'[?.!,]', '', object_name).strip()
+ object_name = re.sub(r'^\s*the\s+', '', object_name, flags=re.IGNORECASE)
+ debug(f'VQA interrogate: handler=moondream3 point_extracted_object="{object_name}"')
+ result = point(image, object_name, repo)
+ if result:
+ from modules.interrogate import vqa
+ vqa.get_instance().last_detection_data = {'points': result}
+ return vqa_detection.format_points_text(result)
+ return "Object not found"
+ elif mode == 'detect':
+ # Extract object name from question - case insensitive
+ object_name = question
+ for phrase in ['detect', 'find all', 'bounding box', 'bbox', 'find']:
+ object_name = re.sub(rf'\b{phrase}\b', '', object_name, flags=re.IGNORECASE)
+ object_name = re.sub(r'[?.!,]', '', object_name).strip()
+ object_name = re.sub(r'^\s*the\s+', '', object_name, flags=re.IGNORECASE)
+ if ' and ' in object_name.lower():
+ object_name = re.split(r'\s+and\s+', object_name, flags=re.IGNORECASE)[0].strip()
+ debug(f'VQA interrogate: handler=moondream3 detect_extracted_object="{object_name}"')
+
+ results = detect(image, object_name, repo, max_objects=kwargs.get('max_objects', 10))
+ if results:
+ from modules.interrogate import vqa
+ vqa.get_instance().last_detection_data = {'detections': results}
+ return vqa_detection.format_detections_text(results)
+ return "No objects detected"
+ else: # mode == 'query'
+ if len(question) < 2:
+ question = "Describe this image."
+ response = query(image, question, repo, stream=stream, use_cache=use_cache, reasoning=thinking_mode)
+
+ debug(f'VQA interrogate: handler=moondream3 response_before_clean="{response}"')
+ return response
+
+ except Exception as e:
+ from modules import errors
+ errors.display(e, 'Moondream3')
+ return f"Error: {str(e)}"
+
+
+def clear_cache():
+ """Clear image encoding cache."""
+ global image_cache # pylint: disable=global-statement
+ cache_size = len(image_cache)
+ image_cache.clear()
+ debug(f'VQA interrogate: handler=moondream3 cleared image cache cache_size_was={cache_size}')
+ shared.log.debug(f'Moondream3: Cleared image cache ({cache_size} entries)')
+
+
+def unload():
+ """Release Moondream 3 model from GPU/memory."""
+ global moondream3_model, loaded # pylint: disable=global-statement
+ if moondream3_model is not None:
+ shared.log.debug(f'Moondream3 unload: model="{loaded}"')
+ sd_models.move_model(moondream3_model, devices.cpu, force=True)
+ moondream3_model = None
+ loaded = None
+ clear_cache()
+ devices.torch_gc(force=True)
+ else:
+ shared.log.debug('Moondream3 unload: no model loaded')
diff --git a/modules/interrogate/vqa.py b/modules/interrogate/vqa.py
index 7c93c56f3..02c132a33 100644
--- a/modules/interrogate/vqa.py
+++ b/modules/interrogate/vqa.py
@@ -3,31 +3,39 @@ import os
import time
import json
import base64
+import copy
import torch
import transformers
import transformers.dynamic_module_utils
from PIL import Image
-from modules import shared, devices, errors, model_quant, sd_models, sd_models_compile
+from modules import shared, devices, errors, model_quant, sd_models, sd_models_compile, ui_symbols
+from modules.interrogate import vqa_detection
-processor = None
-model = None
-loaded: str = None
-quant_args = None
+# Debug logging - function-based to avoid circular import
+debug_enabled = os.environ.get('SD_VQA_DEBUG', None) is not None
+
+def debug(*args, **kwargs):
+ if debug_enabled:
+ shared.log.trace(*args, **kwargs)
+
vlm_default = "Alibaba Qwen 2.5 VL 3B"
vlm_models = {
"Google Gemma 3 4B": "google/gemma-3-4b-it",
"Google Gemma 3n E2B": "google/gemma-3n-E2B-it", # 1.5GB
"Google Gemma 3n E4B": "google/gemma-3n-E4B-it", # 1.5GB
+ "Nidum Gemma 3 4B Uncensored": "nidum/Nidum-Gemma-3-4B-it-Uncensored",
+ "Allura Gemma 3 Glitter 4B": "allura-org/Gemma-3-Glitter-4B",
"Alibaba Qwen 2.0 VL 2B": "Qwen/Qwen2-VL-2B-Instruct",
"Alibaba Qwen 2.5 Omni 3B": "Qwen/Qwen2.5-Omni-3B",
"Alibaba Qwen 2.5 VL 3B": "Qwen/Qwen2.5-VL-3B-Instruct",
"Alibaba Qwen 3 VL 2B": "Qwen/Qwen3-VL-2B-Instruct",
- "Alibaba Qwen 3 VL 2B Thinking": "Qwen/Qwen3-VL-2B-Thinking",
+ f"Alibaba Qwen 3 VL 2B Thinking {ui_symbols.reasoning}": "Qwen/Qwen3-VL-2B-Thinking",
"Alibaba Qwen 3 VL 4B": "Qwen/Qwen3-VL-4B-Instruct",
- "Alibaba Qwen 3 VL 4B Thinking": "Qwen/Qwen3-VL-4B-Thinking",
+ f"Alibaba Qwen 3 VL 4B Thinking {ui_symbols.reasoning}": "Qwen/Qwen3-VL-4B-Thinking",
"Alibaba Qwen 3 VL 8B": "Qwen/Qwen3-VL-8B-Instruct",
- "Alibaba Qwen 3 VL 8B Thinking": "Qwen/Qwen3-VL-8B-Thinking",
+ f"Alibaba Qwen 3 VL 8B Thinking {ui_symbols.reasoning}": "Qwen/Qwen3-VL-8B-Thinking",
+ "XiaomiMiMo MiMo VL 7B RL": "XiaomiMiMo/MiMo-VL-7B-RL-2508", # 8.3GB
"Huggingface Smol VL2 0.5B": "HuggingFaceTB/SmolVLM-500M-Instruct",
"Huggingface Smol VL2 2B": "HuggingFaceTB/SmolVLM-Instruct",
"Apple FastVLM 0.5B": "apple/FastVLM-0.5B",
@@ -41,7 +49,8 @@ vlm_models = {
"MiaoshouAI PromptGen 2.0 Large": "Disty0/Florence-2-large-PromptGen-v2.0", # 1.5GB
"CogFlorence 2.0 Large": "thwri/CogFlorence-2-Large-Freeze", # 1.6GB
"CogFlorence 2.2 Large": "thwri/CogFlorence-2.2-Large", # 1.6GB
- "Moondream 2": "vikhyatk/moondream2", # 3.7GB
+ f"Moondream 2 {ui_symbols.reasoning}": "vikhyatk/moondream2", # 3.7GB
+ f"Moondream 3 Preview {ui_symbols.reasoning}": "moondream/moondream3-preview", # 9.3GB (gated)
"Google Pix Textcaps": "google/pix2struct-textcaps-base", # 1.1GB
"Google PaliGemma 2 3B": "google/paligemma2-3b-pt-224",
"Salesforce BLIP Base": "Salesforce/blip-vqa-base", # 1.5GB
@@ -60,27 +69,205 @@ vlm_models = {
"AIDC Ovis2 4B": "AIDC-AI/Ovis2-4B",
"ByteDance Sa2VA 1B": "ByteDance/Sa2VA-1B",
"ByteDance Sa2VA 4B": "ByteDance/Sa2VA-4B",
- # "OpenGVLab InternVL 2.5 1B": "OpenGVLab/InternVL2_5-1B"
- # "DeepSeek VL2 Tiny": "deepseek-ai/deepseek-vl2-tiny", # broken
- # "nVidia Eagle 2 1B": "nvidia/Eagle2-1B", # not compatible with latest transformers
}
vlm_system = 'You are image captioning expert, creative, unbiased and uncensored.'
-vlm_prompts = [
- '
',
- '',
- '',
- '',
- '',
- '',
- '',
- '',
- '',
- '',
- '',
- '',
- '',
+
+# Common prompts (work with all VLM models)
+vlm_prompts_common = [
+ "Use Prompt",
+ "Short Caption",
+ "Normal Caption",
+ "Long Caption",
]
+# Florence-2 specific prompts (only shown for Florence/PromptGen models)
+vlm_prompts_florence = [
+ "Phrase Grounding",
+ "Object Detection",
+ "Dense Region Caption",
+ "Region Proposal",
+ "OCR (Read Text)",
+ "OCR with Regions",
+ "Analyze",
+ "Generate Tags",
+ "Mixed Caption",
+ "Mixed Caption+",
+]
+
+# Moondream specific prompts (shared by Moondream 2 and 3)
+vlm_prompts_moondream = [
+ "Point at...",
+ "Detect all...",
+]
+
+# Moondream 2 only prompts (gaze detection not available in Moondream 3)
+vlm_prompts_moondream2 = [
+ "Detect Gaze",
+]
+
+# Mapping from friendly names to internal tokens/commands
+vlm_prompt_mapping = {
+ "Use Prompt": "Use Prompt",
+ "Short Caption": "",
+ "Normal Caption": "",
+ "Long Caption": "",
+ "Phrase Grounding": "",
+ "Object Detection": "",
+ "Dense Region Caption": "",
+ "Region Proposal": "",
+ "OCR (Read Text)": "",
+ "OCR with Regions": "",
+ "Analyze": "",
+ "Generate Tags": "",
+ "Mixed Caption": "",
+ "Mixed Caption+": "",
+ "Point at...": "POINT_MODE",
+ "Detect all...": "DETECT_MODE",
+ "Detect Gaze": "DETECT_GAZE",
+}
+
+# Placeholder hints for prompt field based on selected question
+vlm_prompt_placeholders = {
+ "Use Prompt": "Enter your question or instruction for the model",
+ "Short Caption": "Optional: add specific focus or style instructions",
+ "Normal Caption": "Optional: add specific focus or style instructions",
+ "Long Caption": "Optional: add specific focus or style instructions",
+ "Phrase Grounding": "Optional: specify phrases to ground in the image",
+ "Object Detection": "Optional: specify object types to detect",
+ "Dense Region Caption": "Optional: add specific instructions",
+ "Region Proposal": "Optional: add specific instructions",
+ "OCR (Read Text)": "Optional: add specific instructions",
+ "OCR with Regions": "Optional: add specific instructions",
+ "Analyze": "Optional: add specific analysis instructions",
+ "Generate Tags": "Optional: add specific tagging instructions",
+ "Mixed Caption": "Optional: add specific instructions",
+ "Mixed Caption+": "Optional: add specific instructions",
+ "Point at...": "Enter objects to locate, e.g., 'the red car' or 'all the eyes'",
+ "Detect all...": "Enter object type to detect, e.g., 'cars' or 'faces'",
+ "Detect Gaze": "No input needed - auto-detects face and gaze direction",
+}
+
+# Legacy list for backwards compatibility
+vlm_prompts = vlm_prompts_common + vlm_prompts_florence + vlm_prompts_moondream + vlm_prompts_moondream2
+
+vlm_prefill = 'Answer: the image shows'
+
+
+def get_prompts_for_model(model_name: str) -> list:
+ """Get available prompts based on selected model."""
+ if model_name is None:
+ return vlm_prompts_common
+
+ model_lower = model_name.lower()
+
+ # Check for Florence-2 / PromptGen models
+ if 'florence' in model_lower or 'promptgen' in model_lower:
+ return vlm_prompts_common + vlm_prompts_florence
+
+ # Check for Moondream models (Moondream 2 has gaze detection, Moondream 3 does not)
+ if 'moondream' in model_lower:
+ if 'moondream3' in model_lower or 'moondream 3' in model_lower:
+ return vlm_prompts_common + vlm_prompts_moondream
+ else: # Moondream 2 includes gaze detection
+ return vlm_prompts_common + vlm_prompts_moondream + vlm_prompts_moondream2
+
+ # Default: common prompts only
+ return vlm_prompts_common
+
+
+def get_internal_prompt(friendly_name: str, user_prompt: str = None) -> str:
+ """Convert friendly prompt name to internal token/command."""
+ internal = vlm_prompt_mapping.get(friendly_name, friendly_name)
+
+ # Handle Moondream point/detect modes - prepend trigger phrase
+ if internal == "POINT_MODE" and user_prompt:
+ return f"Point at {user_prompt}"
+ elif internal == "DETECT_MODE" and user_prompt:
+ return f"Detect {user_prompt}"
+
+ return internal
+
+
+def get_prompt_placeholder(friendly_name: str) -> str:
+ """Get placeholder text for the prompt field based on selected question."""
+ return vlm_prompt_placeholders.get(friendly_name, "Enter your question or instruction")
+
+
+def is_florence_task(question: str) -> bool:
+ """Check if the question is a Florence-2 task token (either friendly name or internal token)."""
+ if not question:
+ return False
+ # Check if it's a Florence-specific friendly name
+ if question in vlm_prompts_florence:
+ return True
+ # Check if it's an internal Florence-2 task token (for backwards compatibility)
+ florence_tokens = ['', '', '', '',
+ '', '', '', '', '',
+ '', '', '', '']
+ return question in florence_tokens
+
+
+def is_thinking_model(model_name: str) -> bool:
+ """Check if the model supports thinking mode based on its name."""
+ if not model_name:
+ return False
+ model_lower = model_name.lower()
+ # Check for known thinking models
+ thinking_indicators = [
+ 'thinking', # Qwen3-VL-*-Thinking models
+ 'moondream3', # Moondream 3 supports thinking
+ 'moondream 3',
+ 'moondream2', # Moondream 2 supports reasoning mode
+ 'moondream 2',
+ 'mimo',
+ ]
+ return any(indicator in model_lower for indicator in thinking_indicators)
+
+
+def truncate_b64_in_conversation(conversation, front_chars=50, tail_chars=50, threshold=200):
+ """
+ Deep copy a conversation structure and truncate long base64 image strings for logging.
+ Preserves front and tail of base64 strings with truncation indicator.
+ """
+ conv_copy = copy.deepcopy(conversation)
+
+ def truncate_recursive(obj):
+ if isinstance(obj, dict):
+ for key, value in obj.items():
+ if key == "image" and isinstance(value, str) and len(value) > threshold:
+ # Truncate the base64 image string
+ truncated_count = len(value) - front_chars - tail_chars
+ obj[key] = f"{value[:front_chars]}...[{truncated_count} chars truncated]...{value[-tail_chars:]}"
+ elif isinstance(value, (dict, list)):
+ truncate_recursive(value)
+ elif isinstance(obj, list):
+ for item in obj:
+ truncate_recursive(item)
+
+ truncate_recursive(conv_copy)
+ return conv_copy
+
+
+def keep_think_block_open(text_prompt: str) -> str:
+ """Remove the closing of the final assistant message so the model can continue reasoning."""
+ think_open = ""
+ think_close = ""
+ last_open = text_prompt.rfind(think_open)
+ if last_open == -1:
+ return text_prompt
+ close_index = text_prompt.find(think_close, last_open)
+ if close_index == -1:
+ return text_prompt
+ # Skip any whitespace immediately following the closing tag
+ end_close = close_index + len(think_close)
+ while end_close < len(text_prompt) and text_prompt[end_close] in (' ', '\t'):
+ end_close += 1
+ while end_close < len(text_prompt) and text_prompt[end_close] in ('\r', '\n'):
+ end_close += 1
+ trimmed_prompt = text_prompt[:close_index] + text_prompt[end_close:]
+ debug('VQA interrogate: keep_think_block_open applied to prompt segment near assistant reply')
+ return trimmed_prompt
+
def b64(image):
if image is None:
@@ -92,28 +279,60 @@ def b64(image):
return encoded
-def clean(response, question):
- strip = ['---', '\r', '\t', '**', '"', '“', '”', 'Assistant:', 'Caption:', '<|im_end|>', '']
+def clean(response, question, prefill=None):
+ strip = ['---', '\r', '\t', '**', '"', '"', '"', 'Assistant:', 'Caption:', '<|im_end|>', '']
if isinstance(response, str):
response = response.strip()
elif isinstance(response, dict):
+ text_response = ""
+ if 'reasoning' in response and shared.opts.interrogate_vlm_keep_thinking:
+ r_text = response['reasoning']
+ if isinstance(r_text, dict) and 'text' in r_text:
+ r_text = r_text['text']
+ text_response += f"Reasoning:\n{r_text}\n\nAnswer:\n"
+
if 'answer' in response:
- response = response['answer']
+ text_response += response['answer']
+ elif 'caption' in response:
+ text_response += response['caption']
elif 'task' in response:
- response = response['task']
+ text_response += response['task']
else:
- response = json.dumps(response)
+ if not text_response:
+ text_response = json.dumps(response)
+ response = text_response
elif isinstance(response, list):
response = response[0]
else:
response = str(response)
+
+ # Determine prefill text
+ prefill_text = vlm_prefill if prefill is None else prefill
+ if prefill_text is None:
+ prefill_text = ""
+ prefill_text = prefill_text.strip()
+
question = question.replace('<', '').replace('>', '').replace('_', ' ')
if question in response:
response = response.split(question, 1)[1]
while any(s in response for s in strip):
for s in strip:
response = response.replace(s, '')
- response = response.replace('\n\n', '\n').replace(' ', ' ').replace('* ', '- ').strip()
+ response = response.replace(' ', ' ').replace('* ', '- ').strip()
+
+ # Handle prefill retention/removal
+ if shared.opts.interrogate_vlm_keep_prefill:
+ # Add prefill if it's missing from the cleaned response
+ if len(prefill_text) > 0 and not response.startswith(prefill_text):
+ sep = " "
+ if not response or response[0] in ".,!?;:":
+ sep = ""
+ response = f"{prefill_text}{sep}{response}"
+ else:
+ # Remove prefill if it's present in the cleaned response
+ if len(prefill_text) > 0 and response.startswith(prefill_text):
+ response = response[len(prefill_text):].strip()
+
return response
@@ -133,650 +352,1112 @@ def get_kwargs():
return kwargs
-def fastvlm(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- processor = transformers.AutoTokenizer.from_pretrained(repo, trust_remote_code=True)
- model = transformers.AutoModelForCausalLM.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- # device_map="auto",
- trust_remote_code=True,
- cache_dir=shared.opts.hfcache_dir,
- **quant_args,
- )
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- if len(question) < 2:
- question = "Describe the image."
- question = question.replace('<', '').replace('>', '')
- IMAGE_TOKEN_INDEX = -200 # what the model code looks for
- messages = [{"role": "user", "content": f"\n{question}"}]
- rendered = processor.apply_chat_template(messages, add_generation_prompt=True, tokenize=False)
- pre, post = rendered.split("", 1)
- pre_ids = processor(pre, return_tensors="pt", add_special_tokens=False).input_ids
- post_ids = processor(post, return_tensors="pt", add_special_tokens=False).input_ids
- img_tok = torch.tensor([[IMAGE_TOKEN_INDEX]], dtype=pre_ids.dtype)
- input_ids = torch.cat([pre_ids, img_tok, post_ids], dim=1)
- input_ids = input_ids.to(devices.device)
- attention_mask = torch.ones_like(input_ids, device=devices.device)
- px = model.get_vision_tower().image_processor(images=image, return_tensors="pt")
- px = px["pixel_values"].to(model.device, dtype=model.dtype)
- with devices.inference_context():
- outputs = model.generate(
- inputs=input_ids,
- attention_mask=attention_mask,
- images=px,
- max_new_tokens=128,
- )
- answer = processor.decode(outputs[0], skip_special_tokens=True)
- return answer
+class VQA:
+ """Vision-Language Model interrogation class with per-model self-contained loading."""
+ def __init__(self):
+ self.processor = None
+ self.model = None
+ self.loaded: str = None
+ self.quant_args = None
+ self.last_annotated_image = None
+ self.last_detection_data = None
-def qwen(question: str, image: Image.Image, repo: str = None, system_prompt: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if (model is None) or (loaded != repo):
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- if 'Qwen3-VL' in repo or 'Qwen3VL' in repo:
- cls_name = transformers.Qwen3VLForConditionalGeneration
- elif 'Qwen2.5-VL' in repo or 'Qwen2_5_VL' in repo:
- cls_name = transformers.Qwen2_5_VLForConditionalGeneration
- elif 'Qwen2-VL' in repo or 'Qwen2VL' in repo:
- cls_name = transformers.Qwen2VLForConditionalGeneration
+ def unload(self):
+ """Release VLM model from GPU/memory."""
+ if self.model is not None:
+ shared.log.debug(f'VQA unload: model="{self.loaded}"')
+ sd_models.move_model(self.model, devices.cpu, force=True)
+ self.model = None
+ self.processor = None
+ self.loaded = None
+ devices.torch_gc(force=True, reason='vqa unload')
else:
- cls_name = transformers.AutoModelForCausalLM
- model = cls_name.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- **quant_args,
- )
- processor = transformers.AutoProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- if 'LLM' in shared.opts.cuda_compile:
- model = sd_models_compile.compile_torch(model)
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- question = question.replace('<', '').replace('>', '').replace('_', ' ')
- system_prompt = system_prompt or shared.opts.interrogate_vlm_system
- conversation = [
- {
- "role": "system",
- "content": [{"type": "text", "text": system_prompt}],
- },
- {
- "role": "user",
- "content": [
- {"type": "image", "image": b64(image)},
- {"type": "text", "text": question},
- ],
- }
- ]
- text_prompt = processor.apply_chat_template(conversation, add_generation_prompt=True)
- inputs = processor(text=[text_prompt], images=[image], padding=True, return_tensors="pt")
- inputs = inputs.to(devices.device, devices.dtype)
- output_ids = model.generate(
- **inputs,
- **get_kwargs(),
- )
- generated_ids = [
- output_ids[len(input_ids) :]
- for input_ids, output_ids in zip(inputs.input_ids, output_ids)
- ]
- response = processor.batch_decode(generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=True)
- return response
+ shared.log.debug('VQA unload: no model loaded')
+ def load(self, model_name: str = None):
+ """Load VLM model into memory for the specified model name."""
+ model_name = model_name or shared.opts.interrogate_vlm_model
+ if not model_name:
+ shared.log.warning('VQA load: no model specified')
+ return
+ repo = vlm_models.get(model_name)
+ if repo is None:
+ shared.log.error(f'VQA load: unknown model="{model_name}"')
+ return
-def gemma(question: str, image: Image.Image, repo: str = None, system_prompt: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- if '3n' in repo:
- cls = transformers.Gemma3nForConditionalGeneration # pylint: disable=no-member
- else:
- cls = transformers.Gemma3ForConditionalGeneration
- model = cls.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- **quant_args,
- )
- if 'LLM' in shared.opts.cuda_compile:
- model = sd_models_compile.compile_torch(model)
- processor = transformers.AutoProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- question = question.replace('<', '').replace('>', '').replace('_', ' ')
- system_prompt = system_prompt or shared.opts.interrogate_vlm_system
+ self.quant_args = model_quant.create_config(module='LLM')
+ shared.log.debug(f'VQA load: pre-loading model="{model_name}" repo="{repo}"')
- system_content = []
- if system_prompt is not None and len(system_prompt) > 4:
- system_content.append({"type": "text", "text": system_prompt})
-
- user_content = []
- if question is not None and len(question) > 4:
- user_content.append({"type": "text", "text": question})
- if image is not None:
- user_content.append({"type": "image", "image": b64(image)})
- conversation = [
- { "role": "system", "content": system_content},
- { "role": "user", "content": user_content },
- ]
- inputs = processor.apply_chat_template(
- conversation,
- add_generation_prompt=True,
- tokenize=True,
- return_dict=True,
- return_tensors="pt",
- ).to(device=devices.device, dtype=devices.dtype)
- input_len = inputs["input_ids"].shape[-1]
- with devices.inference_context():
- generation = model.generate(
- **inputs,
- **get_kwargs(),
- )
- generation = generation[0][input_len:]
- response = processor.decode(generation, skip_special_tokens=True)
- return response
-
-
-def paligemma(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- processor = transformers.PaliGemmaProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- model = None
- model = transformers.PaliGemmaForConditionalGeneration.from_pretrained(
- repo,
- cache_dir=shared.opts.hfcache_dir,
- torch_dtype=devices.dtype,
- )
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- question = question.replace('<', '').replace('>', '').replace('_', ' ')
- model_inputs = processor(text=question, images=image, return_tensors="pt").to(devices.device, devices.dtype)
- input_len = model_inputs["input_ids"].shape[-1]
- with devices.inference_context():
- generation = model.generate(
- **model_inputs,
- **get_kwargs(),
- )
- generation = generation[0][input_len:]
- response = processor.decode(generation, skip_special_tokens=True)
- return response
-
-
-def ovis(question: str, image: Image.Image, repo: str = None):
- try:
- import flash_attn # pylint: disable=unused-import
- except Exception:
- shared.log.error(f'Interrogate: vlm="{repo}" flash-attn is not available')
- return ''
- global model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- model = transformers.AutoModelForCausalLM.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- multimodal_max_length=32768,
- trust_remote_code=True,
- cache_dir=shared.opts.hfcache_dir,
- )
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- text_tokenizer = model.get_text_tokenizer()
- visual_tokenizer = model.get_visual_tokenizer()
- max_partition = 9
- question = question.replace('<', '').replace('>', '').replace('_', ' ')
- question = f'\n{question}'
- _prompt, input_ids, pixel_values = model.preprocess_inputs(question, [image], max_partition=max_partition)
- attention_mask = torch.ne(input_ids, text_tokenizer.pad_token_id)
- input_ids = input_ids.unsqueeze(0).to(device=model.device)
- attention_mask = attention_mask.unsqueeze(0).to(device=model.device)
- if pixel_values is not None:
- pixel_values = pixel_values.to(dtype=visual_tokenizer.dtype, device=visual_tokenizer.device)
- pixel_values = [pixel_values]
- with devices.inference_context():
- output_ids = model.generate(
- input_ids,
- pixel_values=pixel_values,
- attention_mask=attention_mask,
- repetition_penalty=None,
- eos_token_id=model.generation_config.eos_token_id,
- pad_token_id=text_tokenizer.pad_token_id,
- use_cache=True,
- **get_kwargs())
- response = text_tokenizer.decode(output_ids[0], skip_special_tokens=True)
- print(f'Output:\n{response}')
- return response
-
-
-def smol(question: str, image: Image.Image, repo: str = None, system_prompt: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- model = transformers.AutoModelForVision2Seq.from_pretrained(
- repo,
- cache_dir=shared.opts.hfcache_dir,
- torch_dtype=devices.dtype,
- **quant_args,
- )
- processor = transformers.AutoProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- if 'LLM' in shared.opts.cuda_compile:
- model = sd_models_compile.compile_torch(model)
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- question = question.replace('<', '').replace('>', '').replace('_', ' ')
- system_prompt = system_prompt or shared.opts.interrogate_vlm_system
- conversation = [
- {
- "role": "system",
- "content": [{"type": "text", "text": system_prompt}],
- },
- {
- "role": "user",
- "content": [
- {"type": "image", "image": b64(image)},
- {"type": "text", "text": question},
- ],
- }
- ]
- text_prompt = processor.apply_chat_template(conversation, add_generation_prompt=True)
- inputs = processor(text=text_prompt, images=[image], padding=True, return_tensors="pt")
- inputs = inputs.to(devices.device, devices.dtype)
- output_ids = model.generate(
- **inputs,
- **get_kwargs(),
- )
- response = processor.batch_decode(output_ids,skip_special_tokens=True)
- return response
-
-
-def git(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- model = transformers.GitForCausalLM.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- )
- processor = transformers.GitProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- pixel_values = processor(images=image, return_tensors="pt").pixel_values
- git_dict = {}
- git_dict['pixel_values'] = pixel_values.to(devices.device, devices.dtype)
- if len(question) > 0:
- input_ids = processor(text=question, add_special_tokens=False).input_ids
- input_ids = [processor.tokenizer.cls_token_id] + input_ids
- input_ids = torch.tensor(input_ids).unsqueeze(0)
- git_dict['input_ids'] = input_ids.to(devices.device)
- with devices.inference_context():
- generated_ids = model.generate(**git_dict)
- response = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
- return response
-
-
-def blip(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- model = transformers.BlipForQuestionAnswering.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- )
- processor = transformers.BlipProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- inputs = processor(image, question, return_tensors="pt")
- inputs = inputs.to(devices.device, devices.dtype)
- with devices.inference_context():
- outputs = model.generate(**inputs)
- response = processor.decode(outputs[0], skip_special_tokens=True)
- return response
-
-
-def vilt(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- model = transformers.ViltForQuestionAnswering.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- )
- processor = transformers.ViltProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- inputs = processor(image, question, return_tensors="pt")
- inputs = inputs.to(devices.device)
- with devices.inference_context():
- outputs = model(**inputs)
- logits = outputs.logits
- idx = logits.argmax(-1).item()
- response = model.config.id2label[idx]
- return response
-
-
-def pix(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- model = transformers.Pix2StructForConditionalGeneration.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- )
- processor = transformers.Pix2StructProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- if len(question) > 0:
- inputs = processor(images=image, text=question, return_tensors="pt").to(devices.device)
- else:
- inputs = processor(images=image, return_tensors="pt").to(devices.device)
- with devices.inference_context():
- outputs = model.generate(**inputs)
- response = processor.decode(outputs[0], skip_special_tokens=True)
- return response
-
-
-def moondream(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}"')
- model = None
- model = transformers.AutoModelForCausalLM.from_pretrained(
- repo,
- revision="2025-06-21",
- trust_remote_code=True,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- )
- processor = transformers.AutoTokenizer.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
- loaded = repo
- model.eval()
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- question = question.replace('<', '').replace('>', '').replace('_', ' ')
- encoded = model.encode_image(image)
- with devices.inference_context():
- if question == 'CAPTION':
- response = model.caption(image, length="short")['caption']
- elif question == 'DETAILED CAPTION':
- response = model.caption(image, length="normal")['caption']
- elif question == 'MORE DETAILED CAPTION':
- response = model.caption(image, length="long")['caption']
- else:
- response = model.answer_question(encoded, question, processor)['answer']
- # model.detect(image, "face")
- # model.point(image, "person")
- # model.detect_gaze(image)
- return response
-
-
-def florence(question: str, image: Image.Image, repo: str = None, revision: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- _get_imports = transformers.dynamic_module_utils.get_imports
-
- def get_imports(f):
- R = _get_imports(f)
- if "flash_attn" in R:
- R.remove("flash_attn") # flash_attn is optional
- return R
-
- revision = None
- if '@' in repo:
- repo, revision = repo.split('@')
- if model is None or loaded != repo:
- shared.log.debug(f'Interrogate load: vlm="{repo}" path="{shared.opts.hfcache_dir}"')
- transformers.dynamic_module_utils.get_imports = get_imports
- model = None
- """
- model = transformers.AutoModelForCausalLM.from_pretrained(
- repo,
- trust_remote_code=True,
- revision=revision,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- **quant_args,
- )
- """
- model = transformers.Florence2ForConditionalGeneration.from_pretrained(
- repo,
- dtype=torch.bfloat16,
- revision=revision,
- torch_dtype=devices.dtype,
- cache_dir=shared.opts.hfcache_dir,
- **quant_args,
- )
- processor = transformers.AutoProcessor.from_pretrained(repo, trust_remote_code=True, revision=revision, cache_dir=shared.opts.hfcache_dir)
- transformers.dynamic_module_utils.get_imports = _get_imports
- loaded = repo
- model.eval()
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- if question.startswith('<'):
- task = question.split('>', 1)[0] + '>'
- else:
- task = ''
- inputs = processor(text=task, images=image, return_tensors="pt")
- input_ids = inputs['input_ids'].to(devices.device)
- pixel_values = inputs['pixel_values'].to(devices.device, devices.dtype)
- with devices.inference_context():
- generated_ids = model.generate(
- input_ids=input_ids,
- pixel_values=pixel_values,
- **get_kwargs()
- )
- generated_text = processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
- response = processor.post_process_generation(generated_text, task="task", image_size=(image.width, image.height))
- return response
-
-
-def sa2(question: str, image: Image.Image, repo: str = None):
- global processor, model, loaded # pylint: disable=global-statement
- if model is None or loaded != repo:
- model = None
- model = transformers.AutoModel.from_pretrained(
- repo,
- torch_dtype=devices.dtype,
- low_cpu_mem_usage=True,
- use_flash_attn=False,
- trust_remote_code=True)
- model = model.eval()
- processor = transformers.AutoTokenizer.from_pretrained(
- repo,
- trust_remote_code=True,
- use_fast=False,
- )
- loaded = repo
- devices.torch_gc()
- sd_models.move_model(model, devices.device)
- if question.startswith('<'):
- task = question.split('>', 1)[0] + '>'
- else:
- task = ''
- input_dict = {
- 'image': image,
- 'text': f'{task}',
- 'past_text': '',
- 'mask_prompts': None,
- 'tokenizer': processor,
- }
- return_dict = model.predict_forward(**input_dict)
- response = return_dict["prediction"] # the text format answer
- return response
-
-
-def interrogate(question:str='', system_prompt:str=None, prompt:str=None, image:Image.Image=None, model_name:str=None, quiet:bool=False):
- global quant_args # pylint: disable=global-statement
- jobid = shared.state.begin('Interrogate LLM')
- t0 = time.time()
- quant_args = model_quant.create_config(module='LLM')
- model_name = model_name or shared.opts.interrogate_vlm_model
- if isinstance(image, list):
- image = image[0] if len(image) > 0 else None
- if isinstance(image, dict) and 'name' in image:
- image = Image.open(image['name'])
- if isinstance(image, Image.Image):
- if image.width > 768 or image.height > 768:
- image.thumbnail((768, 768), Image.Resampling.LANCZOS)
- if image.mode != 'RGB':
- image = image.convert('RGB')
- if prompt is not None and len(prompt) > 0:
- question = prompt
- if len(question) < 2:
- question = "Describe the image."
-
- """
- if shared.sd_loaded:
- from modules.sd_models import apply_balanced_offload # prevent circular import
- apply_balanced_offload(shared.sd_model)
- """
-
- from modules import modelloader
- modelloader.hf_login()
-
- try:
- if model_name is None:
- shared.log.error(f'Interrogate: type=vlm model="{model_name}" no model selected')
- return ''
- vqa_model = vlm_models.get(model_name, None)
- if vqa_model is None:
- shared.log.error(f'Interrogate: type=vlm model="{model_name}" unknown')
- return ''
- # if image is None:
- # shared.log.error(f'Interrogate: type=vlm model="{model_name}" no input image')
- # return ''
-
- if 'git' in vqa_model.lower():
- answer = git(question, image, vqa_model)
- elif 'vilt' in vqa_model.lower():
- answer = vilt(question, image, vqa_model)
- elif 'blip' in vqa_model.lower():
- answer = blip(question, image, vqa_model)
- elif 'pix' in vqa_model.lower():
- answer = pix(question, image, vqa_model)
- elif 'moondream2' in vqa_model.lower():
- answer = moondream(question, image, vqa_model)
- elif 'florence' in vqa_model.lower():
- answer = florence(question, image, vqa_model)
- elif 'qwen' in vqa_model.lower() or 'torii' in vqa_model.lower():
- answer = qwen(question, image, vqa_model, system_prompt)
- elif 'smol' in vqa_model.lower():
- answer = smol(question, image, vqa_model, system_prompt)
- elif 'joytag' in vqa_model.lower():
+ # Dispatch to appropriate loader (same logic as interrogate)
+ repo_lower = repo.lower()
+ if 'qwen' in repo_lower or 'torii' in repo_lower or 'mimo' in repo_lower:
+ self._load_qwen(repo)
+ elif 'gemma' in repo_lower and 'pali' not in repo_lower:
+ self._load_gemma(repo)
+ elif 'smol' in repo_lower:
+ self._load_smol(repo)
+ elif 'florence' in repo_lower:
+ self._load_florence(repo)
+ elif 'moondream2' in repo_lower:
+ self._load_moondream(repo)
+ elif 'git' in repo_lower:
+ self._load_git(repo)
+ elif 'blip' in repo_lower:
+ self._load_blip(repo)
+ elif 'vilt' in repo_lower:
+ self._load_vilt(repo)
+ elif 'pix' in repo_lower:
+ self._load_pix(repo)
+ elif 'paligemma' in repo_lower:
+ self._load_paligemma(repo)
+ elif 'ovis' in repo_lower:
+ self._load_ovis(repo)
+ elif 'sa2' in repo_lower:
+ self._load_sa2(repo)
+ elif 'fastvlm' in repo_lower:
+ self._load_fastvlm(repo)
+ elif 'moondream3' in repo_lower:
+ from modules.interrogate import moondream3
+ moondream3.load_model(repo)
+ shared.log.info(f'VQA load: model="{model_name}" loaded (external handler)')
+ return
+ elif 'joytag' in repo_lower:
from modules.interrogate import joytag
- answer = joytag.predict(image)
- elif 'joycaption' in vqa_model.lower():
+ joytag.load()
+ shared.log.info(f'VQA load: model="{model_name}" loaded (external handler)')
+ return
+ elif 'joycaption' in repo_lower:
from modules.interrogate import joycaption
- answer = joycaption.predict(question, image, vqa_model)
- elif 'deepseek' in vqa_model.lower():
+ joycaption.load(repo)
+ shared.log.info(f'VQA load: model="{model_name}" loaded (external handler)')
+ return
+ elif 'deepseek' in repo_lower:
from modules.interrogate import deepseek
- answer = deepseek.predict(question, image, vqa_model)
- elif 'paligemma' in vqa_model.lower():
- answer = paligemma(question, image, vqa_model)
- elif 'gemma' in vqa_model.lower():
- answer = gemma(question, image, vqa_model, system_prompt)
- elif 'ovis' in vqa_model.lower():
- answer = ovis(question, image, vqa_model)
- elif 'sa2' in vqa_model.lower():
- answer = sa2(question, image, vqa_model)
- elif 'fastvlm' in vqa_model.lower():
- answer = fastvlm(question, image, vqa_model)
+ deepseek.load(repo)
+ shared.log.info(f'VQA load: model="{model_name}" loaded (external handler)')
+ return
else:
- answer = 'unknown model'
- except Exception as e:
- errors.display(e, 'VQA')
- answer = 'error'
+ shared.log.warning(f'VQA load: no pre-loader for model="{model_name}"')
+ return
- if shared.opts.interrogate_offload and model is not None:
- sd_models.move_model(model, devices.cpu, force=True)
- devices.torch_gc(force=True, reason='vqa')
- answer = clean(answer, question)
- t1 = time.time()
- if not quiet:
- shared.log.debug(f'Interrogate: type=vlm model="{model_name}" repo="{vqa_model}" args={get_kwargs()} time={t1-t0:.2f}')
- shared.state.end(jobid)
- return answer
+ sd_models.move_model(self.model, devices.device)
+ shared.log.info(f'VQA load: model="{model_name}" loaded')
+
+ def _load_fastvlm(self, repo: str):
+ """Load FastVLM model and tokenizer."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.processor = transformers.AutoTokenizer.from_pretrained(repo, trust_remote_code=True, cache_dir=shared.opts.hfcache_dir)
+ self.model = transformers.AutoModelForCausalLM.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ trust_remote_code=True,
+ cache_dir=shared.opts.hfcache_dir,
+ **self.quant_args,
+ )
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _fastvlm(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ debug(f'VQA interrogate: handler=fastvlm model_name="{model_name}" repo="{repo}" question="{question}" image_size={image.size if image else None}')
+ self._load_fastvlm(repo)
+ sd_models.move_model(self.model, devices.device)
+ if len(question) < 2:
+ question = "Describe the image."
+ question = question.replace('<', '').replace('>', '')
+ IMAGE_TOKEN_INDEX = -200 # what the model code looks for
+ messages = [{"role": "user", "content": f"\n{question}"}]
+ rendered = self.processor.apply_chat_template(messages, add_generation_prompt=True, tokenize=False)
+ pre, post = rendered.split("", 1)
+ pre_ids = self.processor(pre, return_tensors="pt", add_special_tokens=False).input_ids
+ post_ids = self.processor(post, return_tensors="pt", add_special_tokens=False).input_ids
+ img_tok = torch.tensor([[IMAGE_TOKEN_INDEX]], dtype=pre_ids.dtype)
+ input_ids = torch.cat([pre_ids, img_tok, post_ids], dim=1)
+ input_ids = input_ids.to(devices.device)
+ attention_mask = torch.ones_like(input_ids, device=devices.device)
+ px = self.model.get_vision_tower().image_processor(images=image, return_tensors="pt")
+ px = px["pixel_values"].to(self.model.device, dtype=self.model.dtype)
+ with devices.inference_context():
+ outputs = self.model.generate(
+ inputs=input_ids,
+ attention_mask=attention_mask,
+ images=px,
+ max_new_tokens=128,
+ )
+ answer = self.processor.decode(outputs[0], skip_special_tokens=True)
+ return answer
+
+ def _load_qwen(self, repo: str):
+ """Load Qwen VL model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ if 'Qwen3-VL' in repo or 'Qwen3VL' in repo:
+ cls_name = transformers.Qwen3VLForConditionalGeneration
+ elif 'Qwen2.5-VL' in repo or 'Qwen2_5_VL' in repo or 'MiMo-VL' in repo:
+ cls_name = transformers.Qwen2_5_VLForConditionalGeneration
+ elif 'Qwen2-VL' in repo or 'Qwen2VL' in repo:
+ cls_name = transformers.Qwen2VLForConditionalGeneration
+ else:
+ cls_name = transformers.AutoModelForCausalLM
+ self.model = cls_name.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ **self.quant_args,
+ )
+ self.processor = transformers.AutoProcessor.from_pretrained(repo, max_pixels=1024*1024, cache_dir=shared.opts.hfcache_dir)
+ if 'LLM' in shared.opts.cuda_compile:
+ self.model = sd_models_compile.compile_torch(self.model)
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _qwen(self, question: str, image: Image.Image, repo: str, system_prompt: str = None, model_name: str = None, prefill: str = None, thinking_mode: bool = False):
+ self._load_qwen(repo)
+ sd_models.move_model(self.model, devices.device)
+ # Get model class name for logging
+ cls_name = self.model.__class__.__name__
+ debug(f'VQA interrogate: handler=qwen model_name="{model_name}" model_class="{cls_name}" repo="{repo}" question="{question}" system_prompt="{system_prompt}" image_size={image.size if image else None}')
+
+ # Warn if using Florence-2 task tokens with non-Florence-2 models
+ if is_florence_task(question):
+ shared.log.warning(f'Interrogate: Florence-2 task token "{question}" is designed for Florence-2 models. Using it anyway, but results may vary.')
+ question = question.replace('<', '').replace('>', '').replace('_', ' ')
+ system_prompt = system_prompt or shared.opts.interrogate_vlm_system
+ conversation = [
+ {
+ "role": "system",
+ "content": [{"type": "text", "text": system_prompt}],
+ },
+ {
+ "role": "user",
+ "content": [
+ {"type": "image", "image": b64(image)},
+ {"type": "text", "text": question},
+ ],
+ }
+ ]
+ # Add prefill if provided
+ prefill_value = vlm_prefill if prefill is None else prefill
+ prefill_text = prefill_value.strip()
+
+ # Thinking models emit their own tags via the chat template
+ # Only models with thinking capability can use thinking mode
+ is_thinking = is_thinking_model(model_name)
+
+ # Standardize prefill
+ prefill_value = vlm_prefill if prefill is None else prefill
+ prefill_text = prefill_value.strip()
+ use_prefill = len(prefill_text) > 0
+
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=qwen conversation_roles={[msg["role"] for msg in conversation]}')
+ debug(f'VQA interrogate: handler=qwen full_conversation={truncate_b64_in_conversation(conversation)}')
+ debug(f'VQA interrogate: handler=qwen is_thinking={is_thinking} thinking_mode={thinking_mode} prefill="{prefill_text}"')
+
+ # Generate base prompt using template
+ # Qwen-Thinking template automatically adds "<|im_start|>assistant\n\n" when add_generation_prompt=True
+ try:
+ text_prompt = self.processor.apply_chat_template(
+ conversation,
+ add_generation_prompt=True,
+ )
+ except (TypeError, ValueError) as e:
+ debug(f'VQA interrogate: handler=qwen chat_template fallback add_generation_prompt=True: {e}')
+ text_prompt = self.processor.apply_chat_template(conversation, add_generation_prompt=True)
+
+ # Manually handle thinking tags and prefill
+ if is_thinking:
+ if not thinking_mode:
+ # User wants to SKIP thinking.
+ # Since template opened the block with , we close it immediately.
+ text_prompt += "\n"
+ if use_prefill:
+ text_prompt += prefill_text
+ else:
+ # User wants thinking. Prompt already ends in .
+ # If prefill is provided, it becomes part of the thought process.
+ if use_prefill:
+ text_prompt += prefill_text
+ else:
+ # Standard model (not forcing )
+ if use_prefill:
+ text_prompt += prefill_text
+
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=qwen text_prompt="{text_prompt}"')
+ inputs = self.processor(text=[text_prompt], images=[image], padding=True, return_tensors="pt")
+ inputs = inputs.to(devices.device, devices.dtype)
+ gen_kwargs = get_kwargs()
+ debug(f'VQA interrogate: handler=qwen generation_kwargs={gen_kwargs} input_ids_shape={inputs.input_ids.shape}')
+ output_ids = self.model.generate(
+ **inputs,
+ **gen_kwargs,
+ )
+ debug(f'VQA interrogate: handler=qwen output_ids_shape={output_ids.shape}')
+ generated_ids = [
+ output_ids[len(input_ids):]
+ for input_ids, output_ids in zip(inputs.input_ids, output_ids)
+ ]
+ response = self.processor.batch_decode(generated_ids, skip_special_tokens=True, clean_up_tokenization_spaces=True)
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=qwen response_before_clean="{response}"')
+ # Clean up thinking tags
+ # Note: is in the prompt, not the response - only appears in generated output
+ if len(response) > 0:
+ text = response[0]
+ if shared.opts.interrogate_vlm_keep_thinking:
+ # Handle case where is in prompt (not response) but is in response
+ if '' in text and '' not in text:
+ text = 'Reasoning:\n' + text.replace('', '\n\nAnswer:')
+ else:
+ text = text.replace('', 'Reasoning:\n').replace('', '\n\nAnswer:')
+ else:
+ while '' in text:
+ start = text.find('')
+ end = text.find('')
+
+ if start != -1 and start < end:
+ # Standard ...content... block
+ text = text[:start] + text[end+8:]
+ else:
+ # Missing (implied at start) or malformed
+ # Remove from start up to
+ text = text[end+8:]
+ response[0] = text
+ return response
+
+ def _load_gemma(self, repo: str):
+ """Load Gemma 3 model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ if '3n' in repo:
+ cls = transformers.Gemma3nForConditionalGeneration # pylint: disable=no-member
+ else:
+ cls = transformers.Gemma3ForConditionalGeneration
+ self.model = cls.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ **self.quant_args,
+ )
+ if 'LLM' in shared.opts.cuda_compile:
+ self.model = sd_models_compile.compile_torch(self.model)
+ self.processor = transformers.AutoProcessor.from_pretrained(repo, max_pixels=1024*1024, cache_dir=shared.opts.hfcache_dir)
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _gemma(self, question: str, image: Image.Image, repo: str, system_prompt: str = None, model_name: str = None, prefill: str = None, thinking_mode: bool = False):
+ self._load_gemma(repo)
+ sd_models.move_model(self.model, devices.device)
+ # Get model class name for logging
+ cls_name = self.model.__class__.__name__
+ debug(f'VQA interrogate: handler=gemma model_name="{model_name}" model_class="{cls_name}" repo="{repo}" question="{question}" system_prompt="{system_prompt}" image_size={image.size if image else None}')
+
+ # Warn if using Florence-2 task tokens with non-Florence-2 models
+ if is_florence_task(question):
+ shared.log.warning(f'Interrogate: Florence-2 task token "{question}" is designed for Florence-2 models. Using it anyway, but results may vary.')
+ question = question.replace('<', '').replace('>', '').replace('_', ' ')
+ system_prompt = system_prompt or shared.opts.interrogate_vlm_system
+
+ system_content = []
+ if system_prompt is not None and len(system_prompt) > 4:
+ system_content.append({"type": "text", "text": system_prompt})
+
+ user_content = []
+ if question is not None and len(question) > 4:
+ user_content.append({"type": "text", "text": question})
+ if image is not None:
+ user_content.append({"type": "image", "image": b64(image)})
+ conversation = [
+ {"role": "system", "content": system_content},
+ {"role": "user", "content": user_content},
+ ]
+ # Add prefill if provided
+ prefill_value = vlm_prefill if prefill is None else prefill
+ prefill_text = prefill_value.strip()
+ use_prefill = len(prefill_text) > 0
+ # Thinking models emit their own tags via the chat template
+ # Use manual toggle OR auto-detection based on model name
+ use_thinking = thinking_mode or is_thinking_model(model_name)
+ if use_prefill:
+ conversation.append({
+ "role": "assistant",
+ "content": [{"type": "text", "text": prefill_text}],
+ })
+ debug(f'VQA interrogate: handler=gemma prefill="{prefill_text}"')
+ else:
+ debug('VQA interrogate: handler=gemma prefill disabled (empty), relying on add_generation_prompt')
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=gemma conversation_roles={[msg["role"] for msg in conversation]}')
+ debug(f'VQA interrogate: handler=gemma full_conversation={truncate_b64_in_conversation(conversation)}')
+ debug_prefill_mode = 'add_generation_prompt=False continue_final_message=True' if use_prefill else 'add_generation_prompt=True'
+ debug(f'VQA interrogate: handler=gemma template_mode={debug_prefill_mode}')
+ try:
+ if use_prefill:
+ text_prompt = self.processor.apply_chat_template(
+ conversation,
+ add_generation_prompt=False,
+ continue_final_message=True,
+ tokenize=False,
+ )
+ else:
+ text_prompt = self.processor.apply_chat_template(
+ conversation,
+ add_generation_prompt=True,
+ tokenize=False,
+ )
+ except (TypeError, ValueError) as e:
+ debug(f'VQA interrogate: handler=gemma chat_template fallback add_generation_prompt=True: {e}')
+ text_prompt = self.processor.apply_chat_template(
+ conversation,
+ add_generation_prompt=True,
+ tokenize=False,
+ )
+ if use_prefill and use_thinking:
+ text_prompt = keep_think_block_open(text_prompt)
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=gemma text_prompt="{text_prompt}"')
+ inputs = self.processor(
+ text=[text_prompt],
+ images=[image],
+ padding=True,
+ return_tensors="pt",
+ ).to(device=devices.device, dtype=devices.dtype)
+ input_len = inputs["input_ids"].shape[-1]
+ gen_kwargs = get_kwargs()
+ debug(f'VQA interrogate: handler=gemma generation_kwargs={gen_kwargs} input_len={input_len}')
+ with devices.inference_context():
+ generation = self.model.generate(
+ **inputs,
+ **gen_kwargs,
+ )
+ debug(f'VQA interrogate: handler=gemma output_ids_shape={generation.shape}')
+ generation = generation[0][input_len:]
+ response = self.processor.decode(generation, skip_special_tokens=True)
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=gemma response_before_clean="{response}"')
+
+ # Clean up thinking tags (if any remain)
+ if shared.opts.interrogate_vlm_keep_thinking:
+ response = response.replace('', 'Reasoning:\n').replace('', '\n\nAnswer:')
+ else:
+ text = response
+ while '' in text:
+ start = text.find('')
+ end = text.find('')
+ if start != -1 and start < end:
+ text = text[:start] + text[end+8:]
+ else:
+ text = text[end+8:]
+ response = text
+
+ return response
+
+ def _load_paligemma(self, repo: str):
+ """Load PaliGemma model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.processor = transformers.PaliGemmaProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
+ self.model = None
+ self.model = transformers.PaliGemmaForConditionalGeneration.from_pretrained(
+ repo,
+ cache_dir=shared.opts.hfcache_dir,
+ torch_dtype=devices.dtype,
+ )
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _paligemma(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ self._load_paligemma(repo)
+ sd_models.move_model(self.model, devices.device)
+ question = question.replace('<', '').replace('>', '').replace('_', ' ')
+ model_inputs = self.processor(text=question, images=image, return_tensors="pt").to(devices.device, devices.dtype)
+ input_len = model_inputs["input_ids"].shape[-1]
+ with devices.inference_context():
+ generation = self.model.generate(
+ **model_inputs,
+ **get_kwargs(),
+ )
+ generation = generation[0][input_len:]
+ response = self.processor.decode(generation, skip_special_tokens=True)
+ return response
+
+ def _load_ovis(self, repo: str):
+ """Load Ovis model (requires flash-attn)."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.model = transformers.AutoModelForCausalLM.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ multimodal_max_length=32768,
+ trust_remote_code=True,
+ cache_dir=shared.opts.hfcache_dir,
+ )
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _ovis(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ try:
+ import flash_attn # pylint: disable=unused-import
+ except Exception:
+ shared.log.error(f'Interrogate: vlm="{repo}" flash-attn is not available')
+ return ''
+ self._load_ovis(repo)
+ sd_models.move_model(self.model, devices.device)
+ text_tokenizer = self.model.get_text_tokenizer()
+ visual_tokenizer = self.model.get_visual_tokenizer()
+ max_partition = 9
+ question = question.replace('<', '').replace('>', '').replace('_', ' ')
+ question = f'\n{question}'
+ _prompt, input_ids, pixel_values = self.model.preprocess_inputs(question, [image], max_partition=max_partition)
+ attention_mask = torch.ne(input_ids, text_tokenizer.pad_token_id)
+ input_ids = input_ids.unsqueeze(0).to(device=self.model.device)
+ attention_mask = attention_mask.unsqueeze(0).to(device=self.model.device)
+ if pixel_values is not None:
+ pixel_values = pixel_values.to(dtype=visual_tokenizer.dtype, device=visual_tokenizer.device)
+ pixel_values = [pixel_values]
+ with devices.inference_context():
+ output_ids = self.model.generate(
+ input_ids,
+ pixel_values=pixel_values,
+ attention_mask=attention_mask,
+ repetition_penalty=None,
+ eos_token_id=self.model.generation_config.eos_token_id,
+ pad_token_id=text_tokenizer.pad_token_id,
+ use_cache=True,
+ **get_kwargs())
+ response = text_tokenizer.decode(output_ids[0], skip_special_tokens=True)
+ return response
+
+ def _load_smol(self, repo: str):
+ """Load SmolVLM model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.model = transformers.AutoModelForVision2Seq.from_pretrained(
+ repo,
+ cache_dir=shared.opts.hfcache_dir,
+ torch_dtype=devices.dtype,
+ **self.quant_args,
+ )
+ self.processor = transformers.AutoProcessor.from_pretrained(repo, max_pixels=1024*1024, cache_dir=shared.opts.hfcache_dir)
+ if 'LLM' in shared.opts.cuda_compile:
+ self.model = sd_models_compile.compile_torch(self.model)
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _smol(self, question: str, image: Image.Image, repo: str, system_prompt: str = None, model_name: str = None, prefill: str = None, thinking_mode: bool = False):
+ self._load_smol(repo)
+ sd_models.move_model(self.model, devices.device)
+ # Get model class name for logging
+ cls_name = self.model.__class__.__name__
+ debug(f'VQA interrogate: handler=smol model_name="{model_name}" model_class="{cls_name}" repo="{repo}" question="{question}" system_prompt="{system_prompt}" image_size={image.size if image else None}')
+
+ # Warn if using Florence-2 task tokens with non-Florence-2 models
+ if is_florence_task(question):
+ shared.log.warning(f'Interrogate: Florence-2 task token "{question}" is designed for Florence-2 models. Using it anyway, but results may vary.')
+ question = question.replace('<', '').replace('>', '').replace('_', ' ')
+ system_prompt = system_prompt or shared.opts.interrogate_vlm_system
+ conversation = [
+ {
+ "role": "system",
+ "content": [{"type": "text", "text": system_prompt}],
+ },
+ {
+ "role": "user",
+ "content": [
+ {"type": "image", "image": b64(image)},
+ {"type": "text", "text": question},
+ ],
+ }
+ ]
+ # Add prefill if provided
+ prefill_value = vlm_prefill if prefill is None else prefill
+ prefill_text = prefill_value.strip()
+ use_prefill = len(prefill_text) > 0
+ # Thinking models emit their own tags via the chat template
+ # Use manual toggle OR auto-detection based on model name
+ use_thinking = thinking_mode or is_thinking_model(model_name)
+ if use_prefill:
+ conversation.append({
+ "role": "assistant",
+ "content": [{"type": "text", "text": prefill_text}],
+ })
+ debug(f'VQA interrogate: handler=smol prefill="{prefill_text}"')
+ else:
+ debug('VQA interrogate: handler=smol prefill disabled (empty), relying on add_generation_prompt')
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=smol conversation_roles={[msg["role"] for msg in conversation]}')
+ debug(f'VQA interrogate: handler=smol full_conversation={truncate_b64_in_conversation(conversation)}')
+ debug_prefill_mode = 'add_generation_prompt=False continue_final_message=True' if use_prefill else 'add_generation_prompt=True'
+ debug(f'VQA interrogate: handler=smol template_mode={debug_prefill_mode}')
+ try:
+ if use_prefill:
+ text_prompt = self.processor.apply_chat_template(
+ conversation,
+ add_generation_prompt=False,
+ continue_final_message=True,
+ )
+ else:
+ text_prompt = self.processor.apply_chat_template(conversation, add_generation_prompt=True)
+ except (TypeError, ValueError) as e:
+ debug(f'VQA interrogate: handler=smol chat_template fallback add_generation_prompt=True: {e}')
+ text_prompt = self.processor.apply_chat_template(conversation, add_generation_prompt=True)
+ if use_prefill and use_thinking:
+ text_prompt = keep_think_block_open(text_prompt)
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=smol text_prompt="{text_prompt}"')
+ inputs = self.processor(text=text_prompt, images=[image], padding=True, return_tensors="pt")
+ inputs = inputs.to(devices.device, devices.dtype)
+ gen_kwargs = get_kwargs()
+ debug(f'VQA interrogate: handler=smol generation_kwargs={gen_kwargs}')
+ output_ids = self.model.generate(
+ **inputs,
+ **gen_kwargs,
+ )
+ debug(f'VQA interrogate: handler=smol output_ids_shape={output_ids.shape}')
+ response = self.processor.batch_decode(output_ids, skip_special_tokens=True)
+ if debug_enabled:
+ debug(f'VQA interrogate: handler=smol response_before_clean="{response}"')
+
+ # Clean up thinking tags
+ if len(response) > 0:
+ text = response[0]
+ if shared.opts.interrogate_vlm_keep_thinking:
+ text = text.replace('', 'Reasoning:\n').replace('', '\n\nAnswer:')
+ else:
+ while '' in text:
+ start = text.find('')
+ end = text.find('')
+ if start != -1 and start < end:
+ text = text[:start] + text[end+8:]
+ else:
+ text = text[end+8:]
+ response[0] = text
+
+ return response
+
+ def _load_git(self, repo: str):
+ """Load Microsoft GIT model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.model = transformers.GitForCausalLM.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ )
+ self.processor = transformers.GitProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _git(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ self._load_git(repo)
+ sd_models.move_model(self.model, devices.device)
+ pixel_values = self.processor(images=image, return_tensors="pt").pixel_values
+ git_dict = {}
+ git_dict['pixel_values'] = pixel_values.to(devices.device, devices.dtype)
+ if len(question) > 0:
+ input_ids = self.processor(text=question, add_special_tokens=False).input_ids
+ input_ids = [self.processor.tokenizer.cls_token_id] + input_ids
+ input_ids = torch.tensor(input_ids).unsqueeze(0)
+ git_dict['input_ids'] = input_ids.to(devices.device)
+ with devices.inference_context():
+ generated_ids = self.model.generate(**git_dict)
+ response = self.processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
+ return response
+
+ def _load_blip(self, repo: str):
+ """Load Salesforce BLIP model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.model = transformers.BlipForQuestionAnswering.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ )
+ self.processor = transformers.BlipProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _blip(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ self._load_blip(repo)
+ sd_models.move_model(self.model, devices.device)
+ inputs = self.processor(image, question, return_tensors="pt")
+ inputs = inputs.to(devices.device, devices.dtype)
+ with devices.inference_context():
+ outputs = self.model.generate(**inputs)
+ response = self.processor.decode(outputs[0], skip_special_tokens=True)
+ return response
+
+ def _load_vilt(self, repo: str):
+ """Load ViLT model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.model = transformers.ViltForQuestionAnswering.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ )
+ self.processor = transformers.ViltProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _vilt(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ self._load_vilt(repo)
+ sd_models.move_model(self.model, devices.device)
+ inputs = self.processor(image, question, return_tensors="pt")
+ inputs = inputs.to(devices.device)
+ with devices.inference_context():
+ outputs = self.model(**inputs)
+ logits = outputs.logits
+ idx = logits.argmax(-1).item()
+ response = self.model.config.id2label[idx]
+ return response
+
+ def _load_pix(self, repo: str):
+ """Load Pix2Struct model and processor."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.model = transformers.Pix2StructForConditionalGeneration.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ )
+ self.processor = transformers.Pix2StructProcessor.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _pix(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ self._load_pix(repo)
+ sd_models.move_model(self.model, devices.device)
+ if len(question) > 0:
+ inputs = self.processor(images=image, text=question, return_tensors="pt").to(devices.device)
+ else:
+ inputs = self.processor(images=image, return_tensors="pt").to(devices.device)
+ with devices.inference_context():
+ outputs = self.model.generate(**inputs)
+ response = self.processor.decode(outputs[0], skip_special_tokens=True)
+ return response
+
+ def _load_moondream(self, repo: str):
+ """Load Moondream 2 model and tokenizer."""
+ if self.model is None or self.loaded != repo:
+ shared.log.debug(f'Interrogate load: vlm="{repo}"')
+ self.model = None
+ self.model = transformers.AutoModelForCausalLM.from_pretrained(
+ repo,
+ revision="2025-06-21",
+ trust_remote_code=True,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ )
+ self.processor = transformers.AutoTokenizer.from_pretrained(repo, cache_dir=shared.opts.hfcache_dir)
+ self.loaded = repo
+ self.model.eval()
+ devices.torch_gc()
+
+ def _moondream(self, question: str, image: Image.Image, repo: str, model_name: str = None, thinking_mode: bool = False):
+ debug(f'VQA interrogate: handler=moondream model_name="{model_name}" repo="{repo}" question="{question}" thinking_mode={thinking_mode}')
+ self._load_moondream(repo)
+ sd_models.move_model(self.model, devices.device)
+ question = question.replace('<', '').replace('>', '').replace('_', ' ')
+ with devices.inference_context():
+ if question == 'CAPTION':
+ response = self.model.caption(image, length="short")['caption']
+ elif question == 'DETAILED CAPTION':
+ response = self.model.caption(image, length="normal")['caption']
+ elif question == 'MORE DETAILED CAPTION':
+ response = self.model.caption(image, length="long")['caption']
+ elif question.lower().startswith('point at ') or question == 'POINT_MODE':
+ target = question[9:].strip() if question.lower().startswith('point at ') else ''
+ if not target:
+ return "Please specify an object to locate"
+ debug(f'VQA interrogate: handler=moondream method=point target="{target}"')
+ result = self.model.point(image, target)
+ debug(f'VQA interrogate: handler=moondream point_raw_result={result}')
+ points = vqa_detection.parse_points(result)
+ if points:
+ self.last_detection_data = {'points': points}
+ return vqa_detection.format_points_text(points)
+ return "Object not found"
+ elif question.lower().startswith('detect ') or question == 'DETECT_MODE':
+ target = question[7:].strip() if question.lower().startswith('detect ') else ''
+ if not target:
+ return "Please specify an object to detect"
+ debug(f'VQA interrogate: handler=moondream method=detect target="{target}"')
+ result = self.model.detect(image, target)
+ debug(f'VQA interrogate: handler=moondream detect_raw_result={result}')
+ detections = vqa_detection.parse_detections(result, target)
+ if detections:
+ self.last_detection_data = {'detections': detections}
+ return vqa_detection.format_detections_text(detections, include_confidence=False)
+ return "No objects detected"
+ elif question == 'DETECT_GAZE' or question.lower() == 'detect gaze':
+ debug('VQA interrogate: handler=moondream method=detect_gaze')
+ faces = self.model.detect(image, "face")
+ debug(f'VQA interrogate: handler=moondream detect_gaze faces={faces}')
+ if faces.get('objects'):
+ eye_x, eye_y = vqa_detection.calculate_eye_position(faces['objects'][0])
+ result = self.model.detect_gaze(image, eye=(eye_x, eye_y))
+ debug(f'VQA interrogate: handler=moondream detect_gaze result={result}')
+ if result.get('gaze'):
+ gaze = result['gaze']
+ self.last_detection_data = {'points': [(gaze['x'], gaze['y'])]}
+ return f"Gaze direction: ({gaze['x']:.3f}, {gaze['y']:.3f})"
+ return "No face/gaze detected"
+ else:
+ debug(f'VQA interrogate: handler=moondream method=query question="{question}" reasoning={thinking_mode}')
+ result = self.model.query(image, question, reasoning=thinking_mode)
+ response = result['answer']
+ debug(f'VQA interrogate: handler=moondream query_result keys={list(result.keys()) if isinstance(result, dict) else "not dict"}')
+ if thinking_mode and 'reasoning' in result:
+ reasoning_text = result['reasoning'].get('text', '') if isinstance(result['reasoning'], dict) else str(result['reasoning'])
+ debug(f'VQA interrogate: handler=moondream reasoning_text="{reasoning_text[:100]}..."')
+ if shared.opts.interrogate_vlm_keep_thinking:
+ response = f"Reasoning:\n{reasoning_text}\n\nAnswer:\n{response}"
+ # When keep_thinking is False, just use the answer (reasoning is discarded)
+ return response
+
+ def _load_florence(self, repo: str, revision: str = None):
+ """Load Florence-2 model and processor."""
+ _get_imports = transformers.dynamic_module_utils.get_imports
+
+ def get_imports(f):
+ R = _get_imports(f)
+ if "flash_attn" in R:
+ R.remove("flash_attn") # flash_attn is optional
+ return R
+
+ # Handle revision splitting and caching
+ cache_key = repo
+ effective_revision = revision
+ repo_name = repo
+
+ if repo and '@' in repo:
+ repo_name, revision_from_repo = repo.split('@')
+ effective_revision = revision_from_repo
+
+ if self.model is None or self.loaded != cache_key:
+ shared.log.debug(f'Interrogate load: vlm="{repo_name}" revision="{effective_revision}" path="{shared.opts.hfcache_dir}"')
+ transformers.dynamic_module_utils.get_imports = get_imports
+ self.model = None
+ self.model = transformers.Florence2ForConditionalGeneration.from_pretrained(
+ repo_name,
+ dtype=torch.bfloat16,
+ revision=effective_revision,
+ torch_dtype=devices.dtype,
+ cache_dir=shared.opts.hfcache_dir,
+ **self.quant_args,
+ )
+ self.processor = transformers.AutoProcessor.from_pretrained(repo_name, max_pixels=1024*1024, trust_remote_code=True, revision=effective_revision, cache_dir=shared.opts.hfcache_dir)
+ transformers.dynamic_module_utils.get_imports = _get_imports
+ self.loaded = cache_key
+ self.model.eval()
+ devices.torch_gc()
+
+ def _florence(self, question: str, image: Image.Image, repo: str, revision: str = None, model_name: str = None):
+ self._load_florence(repo, revision)
+ sd_models.move_model(self.model, devices.device)
+ if question.startswith('<'):
+ task = question.split('>', 1)[0] + '>'
+ else:
+ task = ''
+ inputs = self.processor(text=task, images=image, return_tensors="pt")
+ input_ids = inputs['input_ids'].to(devices.device)
+ pixel_values = inputs['pixel_values'].to(devices.device, devices.dtype)
+ with devices.inference_context():
+ generated_ids = self.model.generate(
+ input_ids=input_ids,
+ pixel_values=pixel_values,
+ **get_kwargs()
+ )
+ generated_text = self.processor.batch_decode(generated_ids, skip_special_tokens=False)[0]
+ response = self.processor.post_process_generation(generated_text, task="task", image_size=(image.width, image.height))
+ return response
+
+ def _load_sa2(self, repo: str):
+ """Load SA2VA model and tokenizer."""
+ if self.model is None or self.loaded != repo:
+ self.model = None
+ self.model = transformers.AutoModel.from_pretrained(
+ repo,
+ torch_dtype=devices.dtype,
+ low_cpu_mem_usage=True,
+ use_flash_attn=False,
+ trust_remote_code=True)
+ self.model = self.model.eval()
+ self.processor = transformers.AutoTokenizer.from_pretrained(
+ repo,
+ trust_remote_code=True,
+ use_fast=False,
+ )
+ self.loaded = repo
+ devices.torch_gc()
+
+ def _sa2(self, question: str, image: Image.Image, repo: str, model_name: str = None):
+ self._load_sa2(repo)
+ sd_models.move_model(self.model, devices.device)
+ if question.startswith('<'):
+ task = question.split('>', 1)[0] + '>'
+ else:
+ task = ''
+ input_dict = {
+ 'image': image,
+ 'text': f'{task}',
+ 'past_text': '',
+ 'mask_prompts': None,
+ 'tokenizer': self.processor,
+ }
+ return_dict = self.model.predict_forward(**input_dict)
+ response = return_dict["prediction"] # the text format answer
+ return response
+
+ def interrogate(self, question: str = '', system_prompt: str = None, prompt: str = None, image: Image.Image = None, model_name: str = None, prefill: str = None, thinking_mode: bool = False, quiet: bool = False) -> str:
+ """
+ Main entry point for VQA interrogation. Returns string answer.
+ Detection data stored in self.last_detection_data for annotated image creation.
+ """
+ self.last_annotated_image = None
+ self.last_detection_data = None
+ jobid = shared.state.begin('Interrogate LLM')
+ t0 = time.time()
+ self.quant_args = model_quant.create_config(module='LLM')
+ model_name = model_name or shared.opts.interrogate_vlm_model
+ prefill = vlm_prefill if prefill is None else prefill # Use provided prefill when specified
+ if isinstance(image, list):
+ image = image[0] if len(image) > 0 else None
+ if isinstance(image, dict) and 'name' in image:
+ image = Image.open(image['name'])
+ if isinstance(image, Image.Image):
+ if image.width > 768 or image.height > 768:
+ image.thumbnail((768, 768), Image.Resampling.LANCZOS)
+ if image.mode != 'RGB':
+ image = image.convert('RGB')
+ if image is None:
+ shared.log.error(f'VQA interrogate: model="{model_name}" error="No input image provided"')
+ shared.state.end(jobid)
+ return 'Error: No input image provided. Please upload or select an image.'
+
+ # Convert friendly prompt names to internal tokens/commands
+ if question == "Use Prompt":
+ # Use content from Prompt field directly - requires user input
+ if not prompt or len(prompt.strip()) < 2:
+ shared.log.error(f'VQA interrogate: model="{model_name}" error="Please enter a prompt"')
+ shared.state.end(jobid)
+ return 'Error: Please enter a question or instruction in the Prompt field.'
+ question = prompt
+ elif question in vlm_prompt_mapping:
+ # Check if this is a mode that requires user input (Point/Detect)
+ raw_mapping = vlm_prompt_mapping.get(question)
+ if raw_mapping in ("POINT_MODE", "DETECT_MODE"):
+ # These modes require user input in the prompt field
+ if not prompt or len(prompt.strip()) < 2:
+ shared.log.error(f'VQA interrogate: model="{model_name}" error="Please specify what to find in the prompt field"')
+ shared.state.end(jobid)
+ return 'Error: Please specify what to find in the prompt field (e.g., "the red car" or "faces").'
+ # Convert friendly name to internal token (handles Point/Detect prefix)
+ question = get_internal_prompt(question, prompt)
+ # else: question is already an internal token or custom text
+
+ from modules import modelloader
+ modelloader.hf_login()
+
+ try:
+ if model_name is None:
+ shared.log.error(f'Interrogate: type=vlm model="{model_name}" no model selected')
+ shared.state.end(jobid)
+ return ''
+ vqa_model = vlm_models.get(model_name, None)
+ if vqa_model is None:
+ shared.log.error(f'Interrogate: type=vlm model="{model_name}" unknown')
+ shared.state.end(jobid)
+ return ''
+
+ handler = 'unknown'
+ if 'git' in vqa_model.lower():
+ handler = 'git'
+ answer = self._git(question, image, vqa_model, model_name)
+ elif 'vilt' in vqa_model.lower():
+ handler = 'vilt'
+ answer = self._vilt(question, image, vqa_model, model_name)
+ elif 'blip' in vqa_model.lower():
+ handler = 'blip'
+ answer = self._blip(question, image, vqa_model, model_name)
+ elif 'pix' in vqa_model.lower():
+ handler = 'pix'
+ answer = self._pix(question, image, vqa_model, model_name)
+ elif 'moondream3' in vqa_model.lower():
+ handler = 'moondream3'
+ from modules.interrogate import moondream3
+ answer = moondream3.predict(question, image, vqa_model, model_name, thinking_mode=thinking_mode)
+ elif 'moondream2' in vqa_model.lower():
+ handler = 'moondream'
+ answer = self._moondream(question, image, vqa_model, model_name, thinking_mode)
+ elif 'florence' in vqa_model.lower():
+ handler = 'florence'
+ answer = self._florence(question, image, vqa_model, None, model_name)
+ elif 'qwen' in vqa_model.lower() or 'torii' in vqa_model.lower() or 'mimo' in vqa_model.lower():
+ handler = 'qwen'
+ answer = self._qwen(question, image, vqa_model, system_prompt, model_name, prefill, thinking_mode)
+ elif 'smol' in vqa_model.lower():
+ handler = 'smol'
+ answer = self._smol(question, image, vqa_model, system_prompt, model_name, prefill, thinking_mode)
+ elif 'joytag' in vqa_model.lower():
+ handler = 'joytag'
+ from modules.interrogate import joytag
+ answer = joytag.predict(image)
+ elif 'joycaption' in vqa_model.lower():
+ handler = 'joycaption'
+ from modules.interrogate import joycaption
+ answer = joycaption.predict(question, image, vqa_model)
+ elif 'deepseek' in vqa_model.lower():
+ handler = 'deepseek'
+ from modules.interrogate import deepseek
+ answer = deepseek.predict(question, image, vqa_model)
+ elif 'paligemma' in vqa_model.lower():
+ handler = 'paligemma'
+ answer = self._paligemma(question, image, vqa_model, model_name)
+ elif 'gemma' in vqa_model.lower():
+ handler = 'gemma'
+ answer = self._gemma(question, image, vqa_model, system_prompt, model_name, prefill, thinking_mode)
+ elif 'ovis' in vqa_model.lower():
+ handler = 'ovis'
+ answer = self._ovis(question, image, vqa_model, model_name)
+ elif 'sa2' in vqa_model.lower():
+ handler = 'sa2'
+ answer = self._sa2(question, image, vqa_model, model_name)
+ elif 'fastvlm' in vqa_model.lower():
+ handler = 'fastvlm'
+ answer = self._fastvlm(question, image, vqa_model, model_name)
+ else:
+ answer = 'unknown model'
+ except Exception as e:
+ errors.display(e, 'VQA')
+ answer = 'error'
+
+ if shared.opts.interrogate_offload and self.model is not None:
+ sd_models.move_model(self.model, devices.cpu, force=True)
+ devices.torch_gc(force=True, reason='vqa')
+
+ # Clean the answer
+ answer = clean(answer, question, prefill)
+
+ # Create annotated image if detection data is available
+ if self.last_detection_data and isinstance(self.last_detection_data, dict) and image:
+ detections = self.last_detection_data.get('detections', None)
+ points = self.last_detection_data.get('points', None)
+ if detections or points:
+ self.last_annotated_image = vqa_detection.draw_bounding_boxes(image, detections or [], points)
+ debug(f'VQA interrogate: handler={handler} created annotated image detections={len(detections) if detections else 0} points={len(points) if points else 0}')
+
+ debug(f'VQA interrogate: handler={handler} response_after_clean="{answer}" has_annotation={self.last_annotated_image is not None}')
+ t1 = time.time()
+ if not quiet:
+ shared.log.debug(f'Interrogate: type=vlm model="{model_name}" repo="{vqa_model}" args={get_kwargs()} time={t1-t0:.2f}')
+ shared.state.end(jobid)
+ return answer
+
+ def batch(self, model_name, system_prompt, batch_files, batch_folder, batch_str, question, prompt, write, append, recursive, prefill=None, thinking_mode=False):
+ class BatchWriter:
+ def __init__(self, folder, mode='w'):
+ self.folder = folder
+ self.csv = None
+ self.file = None
+ self.mode = mode
+
+ def add(self, file, prompt_text):
+ txt_file = os.path.splitext(file)[0] + ".txt"
+ if self.mode == 'a':
+ prompt_text = '\n' + prompt_text
+ with open(os.path.join(self.folder, txt_file), self.mode, encoding='utf-8') as f:
+ f.write(prompt_text)
+
+ def close(self):
+ if self.file is not None:
+ self.file.close()
+
+ files = []
+ if batch_files is not None:
+ files += [f.name for f in batch_files]
+ if batch_folder is not None:
+ files += [f.name for f in batch_folder]
+ if batch_str is not None and len(batch_str) > 0 and os.path.exists(batch_str) and os.path.isdir(batch_str):
+ from modules.files_cache import list_files
+ files += list(list_files(batch_str, ext_filter=['.png', '.jpg', '.jpeg', '.webp', '.jxl'], recursive=recursive))
+ if len(files) == 0:
+ shared.log.warning('Interrogate batch: type=vlm no images')
+ return ''
+ jobid = shared.state.begin('Interrogate batch')
+ prompts = []
+ if write:
+ mode = 'w' if not append else 'a'
+ writer = BatchWriter(os.path.dirname(files[0]), mode=mode)
+ orig_offload = shared.opts.interrogate_offload
+ shared.opts.interrogate_offload = False
+ import rich.progress as rp
+ pbar = rp.Progress(rp.TextColumn('[cyan]Caption:'), rp.BarColumn(), rp.MofNCompleteColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console)
+ with pbar:
+ task = pbar.add_task(total=len(files), description='starting...')
+ for file in files:
+ pbar.update(task, advance=1, description=file)
+ try:
+ if shared.state.interrupted:
+ break
+ img = Image.open(file)
+ caption = self.interrogate(question, system_prompt, prompt, img, model_name, prefill, thinking_mode, quiet=True)
+ # Save annotated image if available
+ if self.last_annotated_image and write:
+ annotated_path = os.path.splitext(file)[0] + "_annotated.png"
+ self.last_annotated_image.save(annotated_path)
+ prompts.append(caption)
+ if write:
+ writer.add(file, caption)
+ except Exception as e:
+ shared.log.error(f'Interrogate batch: {e}')
+ if write:
+ writer.close()
+ shared.opts.interrogate_offload = orig_offload
+ shared.state.end(jobid)
+ return '\n\n'.join(prompts)
-def batch(model_name, system_prompt, batch_files, batch_folder, batch_str, question, prompt, write, append, recursive):
- class BatchWriter:
- def __init__(self, folder, mode='w'):
- self.folder = folder
- self.csv = None
- self.file = None
- self.mode = mode
+# Module-level singleton instance
+_instance = None
- def add(self, file, prompt):
- txt_file = os.path.splitext(file)[0] + ".txt"
- if self.mode == 'a':
- prompt = '\n' + prompt
- with open(os.path.join(self.folder, txt_file), self.mode, encoding='utf-8') as f:
- f.write(prompt)
- def close(self):
- if self.file is not None:
- self.file.close()
+def get_instance() -> VQA:
+ """Get or create the singleton VQA instance."""
+ global _instance # pylint: disable=global-statement
+ if _instance is None:
+ _instance = VQA()
+ return _instance
- files = []
- if batch_files is not None:
- files += [f.name for f in batch_files]
- if batch_folder is not None:
- files += [f.name for f in batch_folder]
- if batch_str is not None and len(batch_str) > 0 and os.path.exists(batch_str) and os.path.isdir(batch_str):
- from modules.files_cache import list_files
- files += list(list_files(batch_str, ext_filter=['.png', '.jpg', '.jpeg', '.webp', '.jxl'], recursive=recursive))
- if len(files) == 0:
- shared.log.warning('Interrogate batch: type=vlm no images')
- return ''
- jobid = shared.state.begin('Interrogate batch')
- prompts = []
- if write:
- mode = 'w' if not append else 'a'
- writer = BatchWriter(os.path.dirname(files[0]), mode=mode)
- orig_offload = shared.opts.interrogate_offload
- shared.opts.interrogate_offload = False
- import rich.progress as rp
- pbar = rp.Progress(rp.TextColumn('[cyan]Caption:'), rp.BarColumn(), rp.MofNCompleteColumn(), rp.TaskProgressColumn(), rp.TimeRemainingColumn(), rp.TimeElapsedColumn(), rp.TextColumn('[cyan]{task.description}'), console=shared.console)
- with pbar:
- task = pbar.add_task(total=len(files), description='starting...')
- for file in files:
- pbar.update(task, advance=1, description=file)
- try:
- if shared.state.interrupted:
- break
- image = Image.open(file)
- prompt = interrogate(question, system_prompt, prompt, image, model_name, quiet=True)
- prompts.append(prompt)
- if write:
- writer.add(file, prompt)
- except Exception as e:
- shared.log.error(f'Interrogate batch: {e}')
- if write:
- writer.close()
- shared.opts.interrogate_offload = orig_offload
- shared.state.end(jobid)
- return '\n\n'.join(prompts)
+
+# Backwards-compatible module-level functions
+def interrogate(*args, **kwargs):
+ return get_instance().interrogate(*args, **kwargs)
+
+
+def unload_model():
+ return get_instance().unload()
+
+
+def load_model(model_name: str = None):
+ return get_instance().load(model_name)
+
+
+def get_last_annotated_image():
+ return get_instance().last_annotated_image
+
+
+def batch(*args, **kwargs):
+ return get_instance().batch(*args, **kwargs)
diff --git a/modules/interrogate/vqa_detection.py b/modules/interrogate/vqa_detection.py
new file mode 100644
index 000000000..192d669bc
--- /dev/null
+++ b/modules/interrogate/vqa_detection.py
@@ -0,0 +1,207 @@
+# VQA Detection Utilities
+# Parsing, formatting, and drawing functions for detection results (points, bboxes, gaze)
+
+from PIL import Image, ImageDraw, ImageFont
+from modules import shared
+
+
+def parse_points(result) -> list:
+ """Parse and validate point coordinates from model result.
+
+ Args:
+ result: Model output, typically dict with 'points' key or list of coordinates
+
+ Returns:
+ List of (x, y) tuples with coordinates clamped to 0-1 range.
+ """
+ points = []
+
+ # Dict format: {'points': [{'x': 0.5, 'y': 0.5}, ...]}
+ if isinstance(result, dict) and 'points' in result:
+ points_list = result['points']
+ if points_list and len(points_list) > 0:
+ for point_data in points_list:
+ if isinstance(point_data, dict) and 'x' in point_data and 'y' in point_data:
+ x = max(0.0, min(1.0, float(point_data['x'])))
+ y = max(0.0, min(1.0, float(point_data['y'])))
+ points.append((x, y))
+
+ # Fallback for simple [x, y] format
+ elif isinstance(result, (list, tuple)) and len(result) == 2:
+ try:
+ x = max(0.0, min(1.0, float(result[0])))
+ y = max(0.0, min(1.0, float(result[1])))
+ points.append((x, y))
+ except (ValueError, TypeError):
+ pass
+
+ return points
+
+
+def parse_detections(result, label: str, max_objects: int = None) -> list:
+ """Parse and validate detection bboxes from model result.
+
+ Args:
+ result: Model output, typically dict with 'objects' key
+ label: Label to assign to detected objects
+ max_objects: Maximum number of objects to return (None for all)
+
+ Returns:
+ List of {'bbox': [x1,y1,x2,y2], 'label': str, 'confidence': float}
+ with coordinates clamped to 0-1 range.
+ """
+ detections = []
+
+ if isinstance(result, dict) and 'objects' in result:
+ objects = result['objects']
+ if max_objects is not None:
+ objects = objects[:max_objects]
+
+ for obj in objects:
+ if all(k in obj for k in ['x_min', 'y_min', 'x_max', 'y_max']):
+ bbox = [
+ max(0.0, min(1.0, float(obj['x_min']))),
+ max(0.0, min(1.0, float(obj['y_min']))),
+ max(0.0, min(1.0, float(obj['x_max']))),
+ max(0.0, min(1.0, float(obj['y_max'])))
+ ]
+ detections.append({
+ 'bbox': bbox,
+ 'label': label,
+ 'confidence': obj.get('confidence', 1.0)
+ })
+
+ return detections
+
+
+def format_points_text(points: list) -> str:
+ """Format point coordinates as human-readable text.
+
+ Args:
+ points: List of (x, y) tuples with normalized coordinates
+
+ Returns:
+ Formatted text string describing the points.
+ """
+ if not points:
+ return "Object not found"
+
+ if len(points) == 1:
+ return f"Found at: ({points[0][0]:.3f}, {points[0][1]:.3f})"
+
+ lines = [f"Found {len(points)} instances:"]
+ for i, (x, y) in enumerate(points, 1):
+ lines.append(f" {i}. ({x:.3f}, {y:.3f})")
+ return '\n'.join(lines)
+
+
+def format_detections_text(detections: list, include_confidence: bool = True) -> str:
+ """Format detections with bboxes as human-readable text.
+
+ Args:
+ detections: List of detection dicts with 'bbox', 'label', 'confidence'
+ include_confidence: Whether to include confidence scores in output
+
+ Returns:
+ Formatted text string describing the detections.
+ """
+ if not detections:
+ return "No objects detected"
+
+ lines = []
+ for det in detections:
+ bbox = det['bbox']
+ label = det.get('label', 'object')
+ confidence = det.get('confidence', 1.0)
+
+ if include_confidence and confidence < 1.0:
+ lines.append(f"{label}: [{bbox[0]:.3f}, {bbox[1]:.3f}, {bbox[2]:.3f}, {bbox[3]:.3f}] (confidence: {confidence:.2f})")
+ else:
+ lines.append(f"{label}: [{bbox[0]:.3f}, {bbox[1]:.3f}, {bbox[2]:.3f}, {bbox[3]:.3f}]")
+
+ return '\n'.join(lines)
+
+
+def calculate_eye_position(face_bbox: dict) -> tuple:
+ """Calculate approximate eye position from face bounding box.
+
+ Args:
+ face_bbox: Dict with 'x_min', 'y_min', 'x_max', 'y_max' keys
+
+ Returns:
+ (eye_x, eye_y) tuple with normalized coordinates.
+ """
+ eye_x = (face_bbox['x_min'] + face_bbox['x_max']) / 2
+ eye_y = face_bbox['y_min'] + (face_bbox['y_max'] - face_bbox['y_min']) * 0.3 # Approximate eye level
+ return (eye_x, eye_y)
+
+
+def draw_bounding_boxes(image: Image.Image, detections: list, points: list = None) -> Image.Image:
+ """
+ Draw bounding boxes and/or points on an image.
+
+ Args:
+ image: PIL Image to annotate
+ detections: List of detection dicts with format:
+ [{'label': str, 'bbox': [x1, y1, x2, y2], 'confidence': float}, ...]
+ where coordinates are normalized 0-1
+ points: Optional list of (x, y) tuples with normalized 0-1 coordinates
+
+ Returns:
+ Annotated PIL Image with boxes and labels drawn, or None if no annotations
+ """
+ if not detections and not points:
+ return None
+
+ # Create a copy to avoid modifying original
+ annotated = image.copy()
+ draw = ImageDraw.Draw(annotated)
+ width, height = image.size
+
+ # Try to load a font, fall back to default if unavailable
+ try:
+ font_size = max(12, int(min(width, height) * 0.02))
+ font_path = shared.opts.font or "javascript/notosans-nerdfont-regular.ttf"
+ font = ImageFont.truetype(font_path, size=font_size)
+ except Exception:
+ font = ImageFont.load_default()
+
+ # Draw bounding boxes
+ if detections:
+ colors = ['#FF0000', '#00FF00', '#0000FF', '#FFFF00', '#FF00FF', '#00FFFF', '#FFA500', '#800080']
+ for idx, det in enumerate(detections):
+ bbox = det['bbox']
+ label = det.get('label', 'object')
+ confidence = det.get('confidence', 1.0)
+
+ # Convert normalized coordinates to pixel coordinates
+ x1 = int(bbox[0] * width)
+ y1 = int(bbox[1] * height)
+ x2 = int(bbox[2] * width)
+ y2 = int(bbox[3] * height)
+
+ # Choose color
+ color = colors[idx % len(colors)]
+
+ # Draw box
+ draw.rectangle([x1, y1, x2, y2], outline=color, width=max(2, int(min(width, height) * 0.003)))
+
+ # Draw label with background
+ label_text = f"{label} {confidence:.2f}" if confidence < 1.0 else label
+ bbox_font = draw.textbbox((x1, y1), label_text, font=font)
+ text_width = bbox_font[2] - bbox_font[0]
+ text_height = bbox_font[3] - bbox_font[1]
+ draw.rectangle([x1, y1 - text_height - 4, x1 + text_width + 4, y1], fill=color)
+ draw.text((x1 + 2, y1 - text_height - 2), label_text, fill='white', font=font)
+
+ # Draw points
+ if points:
+ point_radius = max(3, int(min(width, height) * 0.01))
+ for px, py in points:
+ x = int(px * width)
+ y = int(py * height)
+ # Draw point as a circle
+ draw.ellipse([x - point_radius, y - point_radius, x + point_radius, y + point_radius],
+ fill='#FF0000', outline='#FFFFFF', width=2)
+
+ return annotated
diff --git a/modules/shared.py b/modules/shared.py
index fc21147f1..7ec809d78 100644
--- a/modules/shared.py
+++ b/modules/shared.py
@@ -656,12 +656,12 @@ options_templates.update(options_section(('interrogate', "Interrogate"), {
"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(16, "CLiP: intermediate flavors", gr.Slider, {"minimum": 0, "maximum": 32, "step": 1, "visible": False}),
- "interrogate_clip_chunk_size": OptionInfo(1024, "CLiP: chunk size", gr.Slider, {"minimum": 256, "maximum": 4096, "step": 8, "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}),
"interrogate_vlm_sep": OptionInfo("VLM
", "", gr.HTML),
"interrogate_vlm_model": OptionInfo(vlm_default, "VLM: default model", gr.Dropdown, {"choices": list(vlm_models)}),
- "interrogate_vlm_prompt": OptionInfo(vlm_prompts[2], "VLM: default prompt", DropdownEditable, {"choices": vlm_prompts }),
+ "interrogate_vlm_prompt": OptionInfo(vlm_prompts[0], "VLM: default prompt", DropdownEditable, {"choices": vlm_prompts }),
"interrogate_vlm_system": OptionInfo(vlm_system, "VLM: default prompt"),
"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}),
@@ -669,6 +669,9 @@ options_templates.update(options_section(('interrogate', "Interrogate"), {
"interrogate_vlm_temperature": OptionInfo(0, "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}),
"deepbooru_sep": OptionInfo("DeepBooru
", "", gr.HTML),
"deepbooru_score_threshold": OptionInfo(0.65, "DeepBooru: score threshold", gr.Slider, {"minimum": 0, "maximum": 1, "step": 0.01}),
diff --git a/modules/ui_caption.py b/modules/ui_caption.py
index e0ef85933..36468b3e4 100644
--- a/modules/ui_caption.py
+++ b/modules/ui_caption.py
@@ -3,14 +3,41 @@ from modules import shared, ui_common, generation_parameters_copypaste
from modules.interrogate import openclip
+def vlm_caption_wrapper(question, system_prompt, prompt, image, model_name, prefill, thinking_mode):
+ """Wrapper for vqa.interrogate that handles annotated image display."""
+ from modules.interrogate import vqa
+ answer = vqa.interrogate(question, system_prompt, prompt, image, model_name, prefill, thinking_mode)
+ annotated_image = vqa.get_last_annotated_image()
+ if annotated_image is not None:
+ return answer, gr.update(value=annotated_image, visible=True)
+ return answer, gr.update(visible=False)
+
+
+def update_vlm_prompts_for_model(model_name):
+ """Update the task dropdown choices based on selected model."""
+ from modules.interrogate import vqa
+ prompts = vqa.get_prompts_for_model(model_name)
+ return gr.update(choices=prompts, value=prompts[0] if prompts else "Use Prompt")
+
+
+def update_vlm_prompt_placeholder(question):
+ """Update the prompt field placeholder based on selected task."""
+ from modules.interrogate import vqa
+ placeholder = vqa.get_prompt_placeholder(question)
+ return gr.update(placeholder=placeholder)
+
+
def update_vlm_params(*args):
- vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p = args
+ vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode = args
shared.opts.interrogate_vlm_max_length = int(vlm_max_tokens)
shared.opts.interrogate_vlm_num_beams = int(vlm_num_beams)
shared.opts.interrogate_vlm_temperature = float(vlm_temperature)
shared.opts.interrogate_vlm_do_sample = bool(vlm_do_sample)
shared.opts.interrogate_vlm_top_k = int(vlm_top_k)
shared.opts.interrogate_vlm_top_p = float(vlm_top_p)
+ shared.opts.interrogate_vlm_keep_prefill = bool(vlm_keep_prefill)
+ shared.opts.interrogate_vlm_keep_thinking = bool(vlm_keep_thinking)
+ shared.opts.interrogate_vlm_thinking_mode = bool(vlm_thinking_mode)
shared.opts.save(shared.config_filename)
@@ -36,14 +63,19 @@ def create_ui():
with gr.Tabs(elem_id="mode_caption"):
with gr.Tab("VLM Caption", elem_id="tab_vlm_caption"):
from modules.interrogate import vqa
+ current_vlm_model = shared.opts.interrogate_vlm_model or vqa.vlm_default
+ initial_prompts = vqa.get_prompts_for_model(current_vlm_model)
with gr.Row():
vlm_system = gr.Textbox(label="System prompt", value=vqa.vlm_system, lines=1, elem_id='vlm_system')
with gr.Row():
- vlm_question = gr.Dropdown(label="Predefined question", allow_custom_value=False, choices=vqa.vlm_prompts, value=vqa.vlm_prompts[2], elem_id='vlm_question')
+ vlm_question = gr.Dropdown(label="Task", allow_custom_value=False, choices=initial_prompts, value=initial_prompts[0] if initial_prompts else "Use Prompt", elem_id='vlm_question')
with gr.Row():
- vlm_prompt = gr.Textbox(label="Prompt", placeholder="optionally enter custom prompt", lines=2, elem_id='vlm_prompt')
+ vlm_prompt = gr.Textbox(label="Prompt", placeholder=vqa.get_prompt_placeholder(initial_prompts[0] if initial_prompts else "Use Prompt"), lines=2, elem_id='vlm_prompt')
with gr.Row(elem_id='interrogate_buttons_query'):
- vlm_model = gr.Dropdown(list(vqa.vlm_models), value=vqa.vlm_default, label='VLM Model', elem_id='vlm_model')
+ vlm_model = gr.Dropdown(list(vqa.vlm_models), value=current_vlm_model, label='VLM Model', elem_id='vlm_model')
+ with gr.Row():
+ vlm_load_btn = gr.Button(value='Load', elem_id='vlm_load', variant='secondary')
+ vlm_unload_btn = gr.Button(value='Unload', elem_id='vlm_unload', variant='secondary')
with gr.Accordion(label='Advanced options', open=False, visible=True):
with gr.Row():
vlm_max_tokens = gr.Slider(label='VLM max tokens', value=shared.opts.interrogate_vlm_max_length, minimum=16, maximum=4096, step=1, elem_id='vlm_max_tokens')
@@ -54,12 +86,21 @@ def create_ui():
vlm_top_p = gr.Slider(label='Top-P', value=shared.opts.interrogate_vlm_top_p, minimum=0.0, maximum=1.0, step=0.01, elem_id='vlm_top_p')
with gr.Row():
vlm_do_sample = gr.Checkbox(label='Use sample', value=shared.opts.interrogate_vlm_do_sample, elem_id='vlm_do_sample')
- vlm_max_tokens.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p], outputs=[])
- vlm_num_beams.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p], outputs=[])
- vlm_temperature.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p], outputs=[])
- vlm_do_sample.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p], outputs=[])
- vlm_top_k.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p], outputs=[])
- vlm_top_p.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p], outputs=[])
+ vlm_thinking_mode = gr.Checkbox(label='Thinking Mode', value=shared.opts.interrogate_vlm_thinking_mode, elem_id='vlm_thinking_mode')
+ with gr.Row():
+ vlm_keep_thinking = gr.Checkbox(label='Keep Thinking Trace', value=shared.opts.interrogate_vlm_keep_thinking, elem_id='vlm_keep_thinking')
+ vlm_keep_prefill = gr.Checkbox(label='Keep Prefill', value=shared.opts.interrogate_vlm_keep_prefill, elem_id='vlm_keep_prefill')
+ with gr.Row():
+ vlm_prefill = gr.Textbox(label='Prefill Text', value='', lines=1, elem_id='vlm_prefill', placeholder='Optional prefill text for model to continue from')
+ vlm_max_tokens.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_num_beams.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_temperature.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_do_sample.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_top_k.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_top_p.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_keep_prefill.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_keep_thinking.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
+ vlm_thinking_mode.change(fn=update_vlm_params, inputs=[vlm_max_tokens, vlm_num_beams, vlm_temperature, vlm_do_sample, vlm_top_k, vlm_top_p, vlm_keep_prefill, vlm_keep_thinking, vlm_thinking_mode], outputs=[])
with gr.Accordion(label='Batch caption', open=False, visible=True):
with gr.Row():
vlm_batch_files = gr.File(label="Files", show_label=True, file_count='multiple', file_types=['image'], interactive=True, height=100, elem_id='vlm_batch_files')
@@ -118,6 +159,8 @@ def create_ui():
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")
+ with gr.Row(elem_id='interrogate_output_image'):
+ output_image = gr.Image(type='pil', label="Annotated Image", interactive=False, visible=False, elem_id='interrogate_output_image_display')
with gr.Row(elem_id='interrogate_output_classes'):
medium = gr.Label(elem_id="interrogate_label_medium", label="Medium", num_top_classes=5, visible=False)
artist = gr.Label(elem_id="interrogate_label_artist", label="Artist", num_top_classes=5, visible=False)
@@ -127,11 +170,19 @@ def create_ui():
with gr.Row(elem_id='copy_buttons_interrogate'):
copy_interrogate_buttons = generation_parameters_copypaste.create_buttons(["txt2img", "img2img", "control", "extras"])
- btn_clip_interrogate_img.click(openclip.interrogate_image, inputs=[image, clip_model, blip_model, clip_mode], outputs=[prompt])
- btn_clip_analyze_img.click(openclip.analyze_image, inputs=[image, clip_model, blip_model], outputs=[medium, artist, movement, trending, flavor])
- btn_clip_interrogate_batch.click(fn=openclip.interrogate_batch, inputs=[clip_batch_files, clip_batch_folder, clip_batch_str, clip_model, blip_model, clip_mode, clip_save_output, clip_save_append, clip_folder_recursive], outputs=[prompt])
- btn_vlm_caption.click(fn=vqa.interrogate, inputs=[vlm_question, vlm_system, vlm_prompt, image, vlm_model], outputs=[prompt])
- btn_vlm_caption_batch.click(fn=vqa.batch, inputs=[vlm_model, vlm_system, vlm_batch_files, vlm_batch_folder, vlm_batch_str, vlm_question, vlm_prompt, vlm_save_output, vlm_save_append, vlm_folder_recursive], outputs=[prompt])
+ btn_clip_interrogate_img.click(openclip.interrogate_image, inputs=[image, clip_model, blip_model, clip_mode], outputs=[prompt]).then(fn=lambda: gr.update(visible=False), inputs=[], outputs=[output_image])
+ btn_clip_analyze_img.click(openclip.analyze_image, inputs=[image, clip_model, blip_model], outputs=[medium, artist, movement, trending, flavor]).then(fn=lambda: gr.update(visible=False), inputs=[], outputs=[output_image])
+ btn_clip_interrogate_batch.click(fn=openclip.interrogate_batch, inputs=[clip_batch_files, clip_batch_folder, clip_batch_str, clip_model, blip_model, clip_mode, clip_save_output, clip_save_append, clip_folder_recursive], outputs=[prompt]).then(fn=lambda: gr.update(visible=False), inputs=[], outputs=[output_image])
+ btn_vlm_caption.click(fn=vlm_caption_wrapper, inputs=[vlm_question, vlm_system, vlm_prompt, image, vlm_model, vlm_prefill, vlm_thinking_mode], outputs=[prompt, output_image])
+ btn_vlm_caption_batch.click(fn=vqa.batch, inputs=[vlm_model, vlm_system, vlm_batch_files, vlm_batch_folder, vlm_batch_str, vlm_question, vlm_prompt, vlm_save_output, vlm_save_append, vlm_folder_recursive, vlm_prefill, vlm_thinking_mode], outputs=[prompt]).then(fn=lambda: gr.update(visible=False), inputs=[], outputs=[output_image])
+
+ # Dynamic UI updates based on selected model and task
+ vlm_model.change(fn=update_vlm_prompts_for_model, inputs=[vlm_model], outputs=[vlm_question])
+ vlm_question.change(fn=update_vlm_prompt_placeholder, inputs=[vlm_question], outputs=[vlm_prompt])
+
+ # Load/Unload model buttons
+ vlm_load_btn.click(fn=vqa.load_model, inputs=[vlm_model], outputs=[])
+ vlm_unload_btn.click(fn=vqa.unload_model, inputs=[], outputs=[])
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,))