From 62880fe8274266705d0c4be8655bfb2d32f77ca9 Mon Sep 17 00:00:00 2001 From: Jordan Thompson Date: Mon, 17 Nov 2025 14:13:23 -0800 Subject: [PATCH] Add manual dataset caption node --- nodes.py | 634 +++++++++++++++++++++++++++++++++++++++++++++---- pyproject.toml | 4 +- 2 files changed, 589 insertions(+), 49 deletions(-) diff --git a/nodes.py b/nodes.py index f832678..e281f7c 100644 --- a/nodes.py +++ b/nodes.py @@ -2,17 +2,18 @@ import os import json import io import uuid +import shutil import numpy as np -import lmstudio as lms from typing import Any, Dict, List, Tuple, Optional from PIL import Image -import shutil + import folder_paths from comfy.utils import ProgressBar + def read_config() -> Dict[str, Any]: """Read configuration from lmstudio_config.json.""" config_path = os.path.join(os.path.dirname(__file__), "lmstudio_config.json") @@ -21,6 +22,7 @@ def read_config() -> Dict[str, Any]: return json.load(f) return {} + def read_tasks() -> Dict[str, str]: """Read task files from the tasks directory.""" tasks_dir = os.path.join(os.path.dirname(__file__), "tasks") @@ -34,7 +36,9 @@ def read_tasks() -> Dict[str, str]: tasks[task_name] = f.read().strip() return tasks + def fetch_models() -> Optional[List[str]]: + import lmstudio as lms """Fetch available models from LM Studio using the official SDK.""" try: downloaded = lms.list_downloaded_models() @@ -47,6 +51,7 @@ def fetch_models() -> Optional[List[str]]: print(f"Error fetching models: {e}") return None + def image_tensor_to_png_bytes(image_tensor, max_edge: int = 1024) -> bytes: """Convert an image tensor to PNG bytes for lmstudio.prepare_image.""" t = image_tensor @@ -70,22 +75,16 @@ def image_tensor_to_png_bytes(image_tensor, max_edge: int = 1024) -> bytes: img.save(buf, format="PNG") return buf.getvalue() -def _temp_images_dir() -> str: - try: - d = folder_paths.get_temp_directory() - except Exception: - d = os.path.join(os.path.dirname(__file__), "temp_images") - os.makedirs(d, exist_ok=True) - return d def image_tensor_to_temp_png_path(image_tensor, max_edge: int = 1024) -> str: data = image_tensor_to_png_bytes(image_tensor, max_edge=max_edge) - d = _temp_images_dir() + d = folder_paths.get_temp_directory() p = os.path.join(d, f"{uuid.uuid4().hex}.png") with open(p, "wb") as f: f.write(data) return p + def listify(x): """Convert input to list if it's not already.""" if x is None: @@ -94,11 +93,24 @@ def listify(x): return list(x) return [x] + def string_list(strings: List[str]) -> Tuple[List[str]]: """Return strings as a tuple containing a list (for ComfyUI output).""" return (strings,) -def _map_params(params: Dict[str, Any]) -> Dict[str, Any]: + +def strip_thinking_tags(text: str) -> str: + """Remove ... and ... tags and their content from text.""" + import re + # Remove thinking tags and everything between them (including newlines) + result = re.sub(r'.*?', '', text, flags=re.DOTALL) + result = re.sub(r'.*?', '', result, flags=re.DOTALL) + # Clean up any extra whitespace left behind + result = re.sub(r'\n\s*\n\s*\n', '\n\n', result) # Replace 3+ newlines with 2 + return result.strip() + + +def map_lms_params(params: Dict[str, Any]) -> Dict[str, Any]: out: Dict[str, Any] = {} for k, v in params.items(): if v is None: @@ -119,6 +131,206 @@ def _map_params(params: Dict[str, Any]) -> Dict[str, Any]: out[k] = v return out + +def check_chat_fits_in_context(model_handle: Any, chat: Any) -> Tuple[bool, int, int]: + """ + Check if a chat conversation fits within the model's context window. + + Args: + model_handle: The loaded LM Studio model instance + chat: The lms.Chat object containing the conversation + + Returns: + Tuple of (fits: bool, token_count: int, context_length: int) + """ + try: + # Convert the conversation to a string using the prompt template + formatted = model_handle.apply_prompt_template(chat) + # Count the number of tokens in the string + token_count = len(model_handle.tokenize(formatted)) + # Get the current loaded context length of the model + context_length = model_handle.get_context_length() + # Check if it fits (leaving some room for the response) + fits = token_count < context_length + return (fits, token_count, context_length) + except Exception as e: + print(f"Warning: Could not check context length: {e}") + # Return conservative defaults if check fails + return (True, 0, 0) + + +def truncate_conversation_to_fit(model_handle: Any, messages: List[Dict[str, Any]], max_tokens_reserve: int = 512) -> List[Dict[str, Any]]: + """ + Truncate conversation history from the beginning until it fits in the model's context. + Always preserves the system message (if present) and removes oldest user/assistant pairs. + + Args: + model_handle: The loaded LM Studio model instance + messages: List of message dictionaries with 'role' and 'content' keys + max_tokens_reserve: Number of tokens to reserve for the response + + Returns: + Truncated list of messages that fits in context + """ + if not messages: + return messages + + try: + context_length = model_handle.get_context_length() + + system_msg = None + conversation = [] + for msg in messages: + if msg.get("role") == "system": + system_msg = msg + else: + conversation.append(msg) + + if not conversation: + return messages + + test_messages = [system_msg] if system_msg else [] + test_messages.extend(conversation) + chat = lms.Chat.from_history({"messages": test_messages}) + fits, token_count, _ = check_chat_fits_in_context(model_handle, chat) + + if fits and (token_count + max_tokens_reserve) < context_length: + return messages + + print(f"Conversation exceeds context length ({token_count} + {max_tokens_reserve} >= {context_length}). Truncating from beginning...") + + while len(conversation) > 2: + conversation.pop(0) + if conversation and conversation[0].get("role") == "assistant": + conversation.pop(0) + + test_messages = [system_msg] if system_msg else [] + test_messages.extend(conversation) + chat = lms.Chat.from_history({"messages": test_messages}) + fits, token_count, _ = check_chat_fits_in_context(model_handle, chat) + + if fits and (token_count + max_tokens_reserve) < context_length: + print(f"Truncated to {len(conversation)} messages ({token_count} tokens)") + return test_messages + + test_messages = [system_msg] if system_msg else [] + test_messages.extend(conversation) + return test_messages + + except Exception as e: + print(f"Warning: Could not truncate conversation: {e}") + return messages + + +def _is_under_allowed_root(path: str, allowed_roots: List[str]) -> bool: + try: + if not allowed_roots: + return False + p = os.path.realpath(path) + for root in allowed_roots: + if not root: + continue + r = os.path.realpath(root) + try: + if os.path.commonpath([p, r]) == r: + return True + except Exception: + continue + return False + except Exception: + return False + + +def classify_orientation(width: int, height: int, square_tolerance: float = 0.05) -> str: + """Classify image orientation as landscape, portrait, or square.""" + aspect = width / height + if abs(aspect - 1.0) <= square_tolerance: + return "square" + if aspect > 1.0: + return "landscape" + return "portrait" + + +def compute_median(values: List[float]) -> float: + """Compute median of a list of values.""" + if not values: + raise ValueError("Cannot compute median of empty list.") + sorted_values = sorted(values) + n = len(sorted_values) + mid = n // 2 + if n % 2 == 1: + return sorted_values[mid] + return (sorted_values[mid - 1] + sorted_values[mid]) / 2.0 + + +def analyze_aspect_ratios(image_paths: List[str]) -> Dict[str, float]: + """Analyze images and compute target aspect ratios for each orientation bucket.""" + orientation_ratios: Dict[str, List[float]] = { + "landscape": [], + "portrait": [], + "square": [], + } + + for image_path in image_paths: + try: + with Image.open(image_path) as image: + width, height = image.size + except Exception: + continue + + orientation = classify_orientation(width, height) + aspect = width / height + orientation_ratios[orientation].append(aspect) + + target_ratios: Dict[str, float] = {} + + for orientation in ["landscape", "portrait", "square"]: + ratios = orientation_ratios[orientation] + if ratios: + median_ratio = compute_median(ratios) + target_ratios[orientation] = median_ratio + else: + if orientation == "landscape": + target_ratios[orientation] = 16.0 / 9.0 + elif orientation == "portrait": + target_ratios[orientation] = 9.0 / 16.0 + else: + target_ratios[orientation] = 1.0 + + print("Aspect ratio analysis:") + for orientation in ["landscape", "portrait", "square"]: + ratios = orientation_ratios[orientation] + count = len(ratios) + ratio_value = target_ratios[orientation] + print(f" {orientation.capitalize()}: {count} image(s), target ratio ≈ {ratio_value:.4f}") + + return target_ratios + + +def crop_image_to_ratio(image: Image.Image, target_ratio: float) -> Image.Image: + """Crop image to match target aspect ratio using center crop.""" + width, height = image.size + current_ratio = width / height + + if abs(current_ratio - target_ratio) < 1e-4: + return image + + if current_ratio > target_ratio: + new_width = int(round(height * target_ratio)) + left = (width - new_width) // 2 + upper = 0 + right = left + new_width + lower = height + else: + new_height = int(round(width / target_ratio)) + left = 0 + upper = (height - new_height) // 2 + right = width + lower = upper + new_height + + return image.crop((left, upper, right, lower)) + + class WASLMStudioModel: @classmethod def INPUT_TYPES(cls): @@ -215,6 +427,8 @@ class WASLMStudioModel: seed: int, image_max_size: str, ): + import lmstudio as lms + cfg = read_config() chosen_id = manual_model_id.strip() if manual_model_id.strip() else model @@ -306,6 +520,8 @@ class WASLMStudioQuery: images=None, options: Optional[Dict[str, Any]] = None, ): + import lmstudio as lms + model_id = model.get("model_id", "") temperature = float(model.get("temperature", 0.2)) max_tokens = int(model.get("max_tokens", 512)) @@ -342,7 +558,7 @@ class WASLMStudioQuery: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) out.append(str(pred)) pbar.update_absolute(idx + 1) responses_out = out @@ -366,7 +582,7 @@ class WASLMStudioQuery: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) responses_out = [str(pred)] else: chat = lms.Chat(system_prompt) if system_prompt else lms.Chat() @@ -383,7 +599,7 @@ class WASLMStudioQuery: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) responses_out = [str(pred)] except Exception as e: @@ -408,6 +624,10 @@ class WASLMStudioQuery: except Exception: pass + # Apply thinking tag filtering if requested + if options and options.get("strip_thinking_tags", False): + responses_out = [strip_thinking_tags(r) for r in responses_out] + return string_list(responses_out) @@ -448,6 +668,10 @@ class WASLMStudioOptions: "FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01, "tooltip": "Generic repetition penalty (model/engine specific). 1.0 means no penalty."}, ), + "strip_thinking_tags": ( + "BOOLEAN", + {"default": False, "tooltip": "Remove ... tags and their content from the output. Useful for getting clean captions without reasoning text."}, + ), }, "optional": { "stop": ( @@ -472,6 +696,7 @@ class WASLMStudioOptions: frequency_penalty: float, presence_penalty: float, repeat_penalty: float, + strip_thinking_tags: bool, stop: str = "", ): params: Dict[str, Any] = { @@ -483,6 +708,7 @@ class WASLMStudioOptions: "frequency_penalty": float(frequency_penalty), "presence_penalty": float(presence_penalty), "repeat_penalty": float(repeat_penalty), + "strip_thinking_tags": bool(strip_thinking_tags), } if stop.strip(): params["stop"] = [s for s in stop.split(",") if s] @@ -634,6 +860,8 @@ class WASLMStudioChat: options: Optional[Dict[str, Any]] = None, temp_convo: bool = False, ): + import lmstudio as lms + model_id = model.get("model_id", "") temperature = float(model.get("temperature", 0.2)) max_tokens = int(model.get("max_tokens", 512)) @@ -656,9 +884,14 @@ class WASLMStudioChat: responses_out: List[str] = ["Error: Failed to process request"] tmp_img_paths: List[str] = [] + model_handle = None try: imgs = listify(images) model_handle = lms.llm(model_id) if model_id else lms.llm() + + # Truncate conversation history if it exceeds context length + # Reserve tokens for max_tokens response + msgs = truncate_conversation_to_fit(model_handle, msgs, max_tokens_reserve=max_tokens) if imgs: if mode == "one-by-one": out: List[str] = [] @@ -681,7 +914,7 @@ class WASLMStudioChat: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) out.append(str(pred)) msgs.append({"role": "user", "content": user_prompt or ""}) msgs.append({"role": "assistant", "content": str(pred)}) @@ -707,7 +940,7 @@ class WASLMStudioChat: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) responses_out = [str(pred)] msgs.append({"role": "user", "content": user_prompt or ""}) msgs.append({"role": "assistant", "content": str(pred)}) @@ -726,7 +959,7 @@ class WASLMStudioChat: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) responses_out = [str(pred)] msgs.append({"role": "user", "content": user_prompt or ""}) msgs.append({"role": "assistant", "content": str(pred)}) @@ -759,6 +992,11 @@ class WASLMStudioChat: responses_out = [m.get("content", "") for m in msgs if m.get("role") == "assistant"] queries_out = [m.get("content", "") for m in msgs if m.get("role") == "user"] + + # Apply thinking tag filtering if requested + if options and options.get("strip_thinking_tags", False): + responses_out = [strip_thinking_tags(r) for r in responses_out] + return (responses_out, queries_out, name) @@ -839,6 +1077,8 @@ class WASLMStudioCaption: user_prompt: str, options: Optional[Dict[str, Any]] = None, ): + import lmstudio as lms + model_id = model.get("model_id", "") temperature = float(model.get("temperature", 0.2)) max_tokens = int(model.get("max_tokens", 512)) @@ -876,7 +1116,7 @@ class WASLMStudioCaption: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) out.append(str(pred)) pbar.update_absolute(idx + 1) result = string_list(out) @@ -901,7 +1141,7 @@ class WASLMStudioCaption: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) + pred = model_handle.respond(chat, config=map_lms_params(params)) result = string_list([str(pred)]) except Exception as e: @@ -920,28 +1160,13 @@ class WASLMStudioCaption: except Exception: pass + # Apply thinking tag filtering if requested + if options and options.get("strip_thinking_tags", False): + result = string_list([strip_thinking_tags(r) for r in result[0]]) + return result -def _is_under_allowed_root(path: str, allowed_roots: List[str]) -> bool: - try: - if not allowed_roots: - return False - p = os.path.realpath(path) - for root in allowed_roots: - if not root: - continue - r = os.path.realpath(root) - try: - if os.path.commonpath([p, r]) == r: - return True - except Exception: - continue - return False - except Exception: - return False - - class WASLoadImageDirectory: @classmethod def INPUT_TYPES(cls): @@ -1002,6 +1227,13 @@ class WASLoadImageDirectory: "tooltip": "Resize/crop method: 'none' (no resize), 'crop_center' (crop to aspect), 'stretch' (distort to fit), 'pad' (letterbox), 'fit' (scale to fit).", }, ), + "normalize_aspect_ratios": ( + "BOOLEAN", + { + "default": False, + "tooltip": "Automatically normalize aspect ratios by analyzing the dataset and grouping images into landscape/portrait/square buckets with median aspect ratios.", + }, + ), }, "optional": { }, @@ -1012,7 +1244,7 @@ class WASLoadImageDirectory: CATEGORY = "LM Studio" FUNCTION = "load_dir" - def load_dir(self, directory_path: str, recursive: bool, extensions: str, dataset_output_path: str, copy_images: bool, force_aspect: str, max_size: str, resize_mode: str): + def load_dir(self, directory_path: str, recursive: bool, extensions: str, dataset_output_path: str, copy_images: bool, force_aspect: str, max_size: str, resize_mode: str, normalize_aspect_ratios: bool): cfg = read_config() allowed_roots = cfg.get("allowed_root_directories", []) if not _is_under_allowed_root(directory_path, allowed_roots): @@ -1034,7 +1266,10 @@ class WASLoadImageDirectory: "__force_aspect": force_aspect, "__max_size": int(max_size) if max_size.isdigit() else 1024, "__resize_mode": resize_mode, + "__normalize_aspect_ratios": bool(normalize_aspect_ratios), } + + image_paths: List[str] = [] if recursive: walker = os.walk(directory_path) for root, _, files in walker: @@ -1042,11 +1277,17 @@ class WASLoadImageDirectory: if os.path.splitext(f)[1].lower() in exts: p = os.path.join(root, f) dataset[p] = {"path": p, "filename": f} + image_paths.append(p) else: for f in os.listdir(directory_path): p = os.path.join(directory_path, f) if os.path.isfile(p) and os.path.splitext(f)[1].lower() in exts: dataset[p] = {"path": p, "filename": f} + image_paths.append(p) + + if normalize_aspect_ratios and image_paths: + target_ratios = analyze_aspect_ratios(image_paths) + dataset["__target_aspect_ratios"] = target_ratios return (dataset,) @@ -1074,6 +1315,18 @@ class WASLMStudioCaptionDataset: "STRING", {"default": "", "multiline": True, "tooltip": "Optional extra instructions."}, ), + "trigger_word_or_phrase": ( + "STRING", + {"default": "", "tooltip": "Optional trigger word or phrase to add to each caption."}, + ), + "trigger_concat_mode": ( + ["prepend", "append"], + {"default": "prepend", "tooltip": "Whether to prepend or append the trigger word/phrase to the caption."}, + ), + "caption_behavior": ( + ["overwrite", "prepend", "append"], + {"default": "overwrite", "tooltip": "Behavior when caption file already exists: overwrite (replace), prepend (add before existing), append (add after existing)."}, + ), }, "optional": { "options": ( @@ -1089,12 +1342,21 @@ class WASLMStudioCaptionDataset: FUNCTION = "caption_dataset" OUTPUT_IS_LIST = (True, True, False) - def caption_dataset(self, model: Dict[str, Any], dataset: Dict[str, Any], task_name: str, user_prompt: str, options: Optional[Dict[str, Any]] = None): + def caption_dataset(self, model: Dict[str, Any], dataset: Dict[str, Any], task_name: str, user_prompt: str, trigger_word_or_phrase: str, trigger_concat_mode: str, caption_behavior: str, options: Optional[Dict[str, Any]] = None): + import lmstudio as lms + model_id = model.get("model_id", "") temperature = float(model.get("temperature", 0.2)) max_tokens = int(model.get("max_tokens", 512)) seed = model.get("seed", None) unload = model.get("unload", False) + + # Apply options overrides if provided + if options: + temperature = float(options.get("temperature", temperature)) + max_tokens = int(options.get("max_tokens", max_tokens)) + if "seed" in options: + seed = options.get("seed", seed) system_text = WASLMStudioCaption.cached_tasks.get(task_name, task_name) meta_output = "" @@ -1102,18 +1364,24 @@ class WASLMStudioCaptionDataset: meta_force_aspect = "none" meta_max_size = 1024 meta_resize_mode = "none" + meta_normalize_aspect = False + meta_target_ratios = {} try: meta_output = str(dataset.get("__output_path", "") or "").strip() meta_copy = bool(dataset.get("__copy_images", False)) meta_force_aspect = str(dataset.get("__force_aspect", "none")) meta_max_size = int(dataset.get("__max_size", 1024)) meta_resize_mode = str(dataset.get("__resize_mode", "none")) + meta_normalize_aspect = bool(dataset.get("__normalize_aspect_ratios", False)) + meta_target_ratios = dataset.get("__target_aspect_ratios", {}) except Exception: meta_output = "" meta_copy = False meta_force_aspect = "none" meta_max_size = 1024 meta_resize_mode = "none" + meta_normalize_aspect = False + meta_target_ratios = {} written: List[str] = [] captions: List[str] = [] @@ -1145,22 +1413,60 @@ class WASLMStudioCaptionDataset: params[k] = v except Exception: pass - pred = model_handle.respond(chat, config=_map_params(params)) - text = str(pred) - captions.append(text.strip()) - # Determine output directory + pred = model_handle.respond(chat, config=map_lms_params(params)) + text = str(pred).strip() + + # Apply thinking tag filtering if requested + if options and options.get("strip_thinking_tags", False): + text = strip_thinking_tags(text) + + if trigger_word_or_phrase.strip(): + trigger = trigger_word_or_phrase.strip() + if trigger_concat_mode == "prepend": + text = f"{trigger} {text}" + else: # append + text = f"{text} {trigger}" + + captions.append(text) src_dir = os.path.dirname(p) out_dir = meta_output if meta_output else src_dir os.makedirs(out_dir, exist_ok=True) base = os.path.splitext(os.path.basename(p))[0] txt_path = os.path.join(out_dir, base + ".txt") - # Copy image if requested and output differs from source + + # Handle existing caption file based on caption_behavior + if os.path.exists(txt_path) and caption_behavior != "overwrite": + try: + with open(txt_path, "r", encoding="utf-8") as f: + existing_caption = f.read().strip() + if caption_behavior == "prepend": + text = f"{text} {existing_caption}" + elif caption_behavior == "append": + text = f"{existing_caption} {text}" + except Exception as read_err: + print(f"Could not read existing caption at {txt_path}: {read_err}") + if meta_copy and meta_output and os.path.realpath(out_dir) != os.path.realpath(src_dir): dst_img = os.path.join(out_dir, os.path.basename(p)) if not os.path.exists(dst_img): - shutil.copy2(p, dst_img) + if meta_normalize_aspect and meta_target_ratios: + try: + with Image.open(p) as img: + img = img.convert("RGB") + width, height = img.size + orientation = classify_orientation(width, height) + target_ratio = meta_target_ratios.get(orientation, width / height) + normalized_img = crop_image_to_ratio(img, target_ratio) + if meta_max_size > 0: + normalized_img.thumbnail((meta_max_size, meta_max_size), Image.LANCZOS) + normalized_img.save(dst_img, quality=95, optimize=True) + except Exception as norm_err: + print(f"Normalization failed for {p}, copying original: {norm_err}") + shutil.copy2(p, dst_img) + else: + shutil.copy2(p, dst_img) with open(txt_path, "w", encoding="utf-8") as f: - f.write(text.strip()) + f.write(text) written.append(txt_path) pbar.update_absolute(idx + 1) except Exception as ie: @@ -1197,13 +1503,245 @@ class WASLMStudioCaptionDataset: f" Max Size: {meta_max_size}", f" Resize Mode: {meta_resize_mode}", f" Copy Images: {meta_copy}", + f" Normalize Aspect Ratios: {meta_normalize_aspect}", ] if meta_output: result_lines.append(f" Output Path: {meta_output}") + if meta_normalize_aspect and meta_target_ratios: + result_lines.append(f" Target Aspect Ratios:") + for orientation, ratio in meta_target_ratios.items(): + result_lines.append(f" {orientation.capitalize()}: {ratio:.4f}") result_summary = "\n".join(result_lines) return (captions, written, result_summary) + +class WASLMStudioCaptionDatasetCustom: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ( + "LMSTUDIO_MODEL", + {"tooltip": "Model settings from LM Studio Model node."}, + ), + "dataset": ( + "LMSTUDIO_DATASET_IMAGES", + {"tooltip": "Dictionary produced by WASLoadImageDirectory mapping file paths to metadata."}, + ), + "system_prompt": ( + "STRING", + { + "default": "You are a professional image captioner. Describe the image in natural, concise prose with attention to composition, lighting, and details.", + "multiline": True, + "tooltip": "System prompt that defines the assistant's behavior and captioning style.", + }, + ), + "user_prompt": ( + "STRING", + {"default": "", "multiline": True, "tooltip": "Optional extra instructions for each image."}, + ), + "trigger_word_or_phrase": ( + "STRING", + {"default": "", "tooltip": "Optional trigger word or phrase to add to each caption."}, + ), + "trigger_concat_mode": ( + ["prepend", "append"], + {"default": "prepend", "tooltip": "Whether to prepend or append the trigger word/phrase to the caption."}, + ), + "caption_behavior": ( + ["overwrite", "prepend", "append"], + {"default": "overwrite", "tooltip": "Behavior when caption file already exists: overwrite (replace), prepend (add before existing), append (add after existing)."}, + ), + }, + "optional": { + "options": ( + "LMSTUDIO_OPTIONS", + {"tooltip": "Override generation options."}, + ), + }, + } + + RETURN_TYPES = ("STRING", "STRING", "STRING") + RETURN_NAMES = ("captions", "written_caption_paths", "result") + CATEGORY = "LM Studio" + FUNCTION = "caption_dataset" + OUTPUT_IS_LIST = (True, True, False) + + def caption_dataset(self, model: Dict[str, Any], dataset: Dict[str, Any], system_prompt: str, user_prompt: str, trigger_word_or_phrase: str, trigger_concat_mode: str, caption_behavior: str, options: Optional[Dict[str, Any]] = None): + import lmstudio as lms + + model_id = model.get("model_id", "") + temperature = float(model.get("temperature", 0.2)) + max_tokens = int(model.get("max_tokens", 512)) + seed = model.get("seed", None) + unload = model.get("unload", False) + + # Apply options overrides if provided + if options: + temperature = float(options.get("temperature", temperature)) + max_tokens = int(options.get("max_tokens", max_tokens)) + if "seed" in options: + seed = options.get("seed", seed) + + system_text = system_prompt.strip() + meta_output = "" + meta_copy = False + meta_force_aspect = "none" + meta_max_size = 1024 + meta_resize_mode = "none" + meta_normalize_aspect = False + meta_target_ratios = {} + try: + meta_output = str(dataset.get("__output_path", "") or "").strip() + meta_copy = bool(dataset.get("__copy_images", False)) + meta_force_aspect = str(dataset.get("__force_aspect", "none")) + meta_max_size = int(dataset.get("__max_size", 1024)) + meta_resize_mode = str(dataset.get("__resize_mode", "none")) + meta_normalize_aspect = bool(dataset.get("__normalize_aspect_ratios", False)) + meta_target_ratios = dataset.get("__target_aspect_ratios", {}) + except Exception: + meta_output = "" + meta_copy = False + meta_force_aspect = "none" + meta_max_size = 1024 + meta_resize_mode = "none" + meta_normalize_aspect = False + meta_target_ratios = {} + + written: List[str] = [] + captions: List[str] = [] + failed_count = 0 + skipped_count = 0 + model_handle = None + try: + model_handle = lms.llm(model_id) if model_id else lms.llm() + image_paths = [k for k in dataset.keys() if not str(k).startswith("__")] + pbar = ProgressBar(len(image_paths)) + for idx, p in enumerate(image_paths): + try: + if not os.path.isfile(p): + skipped_count += 1 + pbar.update_absolute(idx + 1) + continue + img_handle = lms.prepare_image(p) + chat = lms.Chat(system_text) if system_text else lms.Chat() + chat.add_user_message(user_prompt or "", images=[img_handle]) + params: Dict[str, Any] = { + "temperature": temperature, + "max_tokens": max_tokens, + } + if seed is not None: + params["seed"] = seed + if options: + try: + for k, v in options.items(): + params[k] = v + except Exception: + pass + pred = model_handle.respond(chat, config=map_lms_params(params)) + text = str(pred).strip() + + # Apply thinking tag filtering if requested + if options and options.get("strip_thinking_tags", False): + text = strip_thinking_tags(text) + + if trigger_word_or_phrase.strip(): + trigger = trigger_word_or_phrase.strip() + if trigger_concat_mode == "prepend": + text = f"{trigger} {text}" + else: # append + text = f"{text} {trigger}" + + captions.append(text) + src_dir = os.path.dirname(p) + out_dir = meta_output if meta_output else src_dir + os.makedirs(out_dir, exist_ok=True) + base = os.path.splitext(os.path.basename(p))[0] + txt_path = os.path.join(out_dir, base + ".txt") + + # Handle existing caption file based on caption_behavior + if os.path.exists(txt_path) and caption_behavior != "overwrite": + try: + with open(txt_path, "r", encoding="utf-8") as f: + existing_caption = f.read().strip() + if caption_behavior == "prepend": + text = f"{text} {existing_caption}" + elif caption_behavior == "append": + text = f"{existing_caption} {text}" + except Exception as read_err: + print(f"Could not read existing caption at {txt_path}: {read_err}") + + if meta_copy and meta_output and os.path.realpath(out_dir) != os.path.realpath(src_dir): + dst_img = os.path.join(out_dir, os.path.basename(p)) + if not os.path.exists(dst_img): + if meta_normalize_aspect and meta_target_ratios: + try: + with Image.open(p) as img: + img = img.convert("RGB") + width, height = img.size + orientation = classify_orientation(width, height) + target_ratio = meta_target_ratios.get(orientation, width / height) + normalized_img = crop_image_to_ratio(img, target_ratio) + if meta_max_size > 0: + normalized_img.thumbnail((meta_max_size, meta_max_size), Image.LANCZOS) + normalized_img.save(dst_img, quality=95, optimize=True) + except Exception as norm_err: + print(f"Normalization failed for {p}, copying original: {norm_err}") + shutil.copy2(p, dst_img) + else: + shutil.copy2(p, dst_img) + with open(txt_path, "w", encoding="utf-8") as f: + f.write(text) + written.append(txt_path) + pbar.update_absolute(idx + 1) + except Exception as ie: + print(f"Caption failed for {p}: {ie}") + failed_count += 1 + pbar.update_absolute(idx + 1) + continue + except Exception as e: + print(f"Error in caption_dataset: {e}") + finally: + if unload and model_handle is not None: + try: + model_handle.unload() + except Exception: + pass + + # Build result summary + total_images = len(image_paths) + success_count = len(written) + result_lines = [ + f"Dataset Captioning Complete", + f"="*50, + f"Total Images: {total_images}", + f"Successfully Captioned: {success_count}", + f"Failed: {failed_count}", + f"Skipped: {skipped_count}", + f"", + f"Settings:", + f" System Prompt: {system_text[:50]}..." if len(system_text) > 50 else f" System Prompt: {system_text}", + f" Model: {model_id}", + f" Temperature: {temperature}", + f" Max Tokens: {max_tokens}", + f" Force Aspect: {meta_force_aspect}", + f" Max Size: {meta_max_size}", + f" Resize Mode: {meta_resize_mode}", + f" Copy Images: {meta_copy}", + f" Normalize Aspect Ratios: {meta_normalize_aspect}", + ] + if meta_output: + result_lines.append(f" Output Path: {meta_output}") + if meta_normalize_aspect and meta_target_ratios: + result_lines.append(f" Target Aspect Ratios:") + for orientation, ratio in meta_target_ratios.items(): + result_lines.append(f" {orientation.capitalize()}: {ratio:.4f}") + result_summary = "\n".join(result_lines) + + return (captions, written, result_summary) + + # Node mappings for ComfyUI NODE_CLASS_MAPPINGS = { "WASLMStudioModel": WASLMStudioModel, @@ -1213,6 +1751,7 @@ NODE_CLASS_MAPPINGS = { "WASLMStudioOptions": WASLMStudioOptions, "WASLoadImageDirectory": WASLoadImageDirectory, "WASLMStudioCaptionDataset": WASLMStudioCaptionDataset, + "WASLMStudioCaptionDatasetCustom": WASLMStudioCaptionDatasetCustom, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -1223,4 +1762,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WASLMStudioOptions": "LM Studio Options", "WASLoadImageDirectory": "WAS Load Image Directory", "WASLMStudioCaptionDataset": "LM Studio Easy-Caption Dataset", + "WASLMStudioCaptionDatasetCustom": "LM Studio Caption Dataset", } \ No newline at end of file diff --git a/pyproject.toml b/pyproject.toml index e50b78f..9f51f8f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "was-lmstudio-easyquery" -version = "0.1.0" +version = "0.2.0" description = "LM Studio calls with image support and easy task directions." requires-python = ">=3.10" license = { file = "LICENSE" } @@ -22,4 +22,4 @@ PublisherId = "was" DisplayName = "WAS LMStudio Easy-Query" # Icon = "https://raw.githubusercontent.com/WASasquatch/WAS_LMStudio_EasyQuery/main/icon.png" # Banner = "https://raw.githubusercontent.com/WASasquatch/WAS_LMStudio_EasyQuery/main/banner.png" -requires-comfyui = ">=1.0.0" \ No newline at end of file +requires-comfyui = ">=1.0.0"