commit 8588aa5430a4e30decedac25953f00ccd31cd7b2 Author: Obada Date: Mon Jul 28 11:31:30 2025 +0300 initial commit diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..b035a9b --- /dev/null +++ b/__init__.py @@ -0,0 +1,57 @@ +# ComfyUI/custom_nodes/ComfyUI_LocalLLMNodes/__init__.py +# --- Conditional Import Logic --- +LOCAL_LLM_NODES_AVAILABLE = False +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +# --- Check for core dependencies --- +# We need transformers and torch for the LLM nodes to work. +try: + import transformers + import torch + # Optionally check for bitsandbytes here if you plan to use quantization immediately + # import bitsandbytes as bnb + CORE_DEPS_AVAILABLE = True +except ImportError as e: + print(f"[ComfyUI_LocalLLMNodes] Core dependencies (transformers, torch) not found or not importable: {e}") + CORE_DEPS_AVAILABLE = False + +if CORE_DEPS_AVAILABLE: + try: + # --- Import Node Classes --- + from .local_llm_connector import SetLocalLLMServiceConnector + from .local_prompt_generator import ( + LocalKontextPromptGenerator, + AddUserLocalKontextPreset, + RemoveUserLocalKontextPreset + ) + + # --- Define Mappings --- + # Using the class names directly as keys is standard. + NODE_CLASS_MAPPINGS = { + "SetLocalLLMServiceConnector": SetLocalLLMServiceConnector, + "LocalKontextPromptGenerator": LocalKontextPromptGenerator, + "AddUserLocalKontextPreset": AddUserLocalKontextPreset, + "RemoveUserLocalKontextPreset": RemoveUserLocalKontextPreset, + } + + NODE_DISPLAY_NAME_MAPPINGS = { + "SetLocalLLMServiceConnector": "Set Local LLM Service Connector 🐑", + "LocalKontextPromptGenerator": "Local Kontext Prompt Generator 🐑", + "AddUserLocalKontextPreset": "Add User Local Kontext Preset 🐑", + "RemoveUserLocalKontextPreset": "Remove User Local Kontext Preset 🐑", + } + + LOCAL_LLM_NODES_AVAILABLE = True + print("[ComfyUI_LocalLLMNodes] All nodes are available.") + + except Exception as e: + print(f"[ComfyUI_LocalLLMNodes] Error importing node classes: {e}") + # If import fails, NODE_CLASS_MAPPINGS and NODE_DISPLAY_NAME_MAPPINGS remain empty, + # and ComfyUI won't register any nodes from this package. +else: + print("[ComfyUI_LocalLLMNodes] Nodes will NOT be available due to missing core dependencies.") + +# --- Define what ComfyUI sees --- +# ComfyUI looks for these specific dictionaries in __init__.py +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/local_llm_connector.py b/local_llm_connector.py new file mode 100644 index 0000000..4472780 --- /dev/null +++ b/local_llm_connector.py @@ -0,0 +1,268 @@ +# ComfyUI/custom_nodes/ComfyUI_LocalLLMNodes/local_llm_connector.py +import os +import folder_paths + +# --- Simple logging utility for this package --- +def log(message): + print(f"[LocalLLMNodes] {message}") + +# --- Library for local LLM inference --- +# Requires 'transformers' and 'torch'. These are checked in __init__.py. +try: + from transformers import AutoTokenizer, AutoModelForCausalLM + # Optional: For quantization (requires 'bitsandbytes') + # from transformers import BitsAndBytesConfig + import torch + HF_AVAILABLE = True +except ImportError: + log("Warning: transformers or torch library not found. Local LLM node will not work.") + HF_AVAILABLE = False + AutoTokenizer = None + AutoModelForCausalLM = None + torch = None + +# Define category directly for this package +LOCAL_LLM_CATEGORY = "Local LLM Nodes/LLM Connectors" + +def get_local_llm_model_names(): + """Discovers subdirectories within models/LLM that could be local LLM models.""" + llm_models_dir = os.path.join(folder_paths.models_dir, "LLM") + model_names = [] + + if os.path.exists(llm_models_dir): + try: + # List directories inside models/LLM + for item in os.listdir(llm_models_dir): + item_path = os.path.join(llm_models_dir, item) + if os.path.isdir(item_path): + # Basic check: assume any directory is a potential model + model_names.append(item) + except Exception as e: + log(f"Error scanning models/LLM directory: {e}") + else: + log(f"models/LLM directory not found: {llm_models_dir}") + + if not model_names: + model_names = ["No_Local_Models_Found"] # Placeholder if no models found + + return model_names + +class SetLocalLLMServiceConnector: + """ + A node to select and prepare a connection to a local LLM model. + Models should be placed in ComfyUI/models/LLM/your_model_name. + Requires 'transformers' and 'torch': pip install transformers torch + """ + @classmethod + def INPUT_TYPES(cls): + model_names = get_local_llm_model_names() + return { + "required": { + "local_model_name": (model_names, {"default": model_names[0] if model_names else "No_Local_Models_Found"}), + }, + } + + # --- CRITICAL: Match the type identifier expected by consuming nodes --- + # Based on MieNodes structure (TextTranslator, KontextPromptGenerator), this should be "LLMServiceConnector" + # This allows your node to connect to standard MieNodes prompt generators if desired. + RETURN_TYPES = ("LLMServiceConnector",) + FUNCTION = "get_connector" + CATEGORY = LOCAL_LLM_CATEGORY + OUTPUT_NODE = False + + def get_connector(self, local_model_name): + """ + Returns a connector object that can invoke the selected local model. + The actual model loading happens within the connector's invoke method. + """ + if not HF_AVAILABLE: + raise Exception("The 'transformers' and 'torch' libraries are required for the Local LLM node but are not installed.") + + if local_model_name == "No_Local_Models_Found": + raise Exception("No local LLM models found in models/LLM directory. Please place your models there.") + + model_path = os.path.join(folder_paths.models_dir, "LLM", local_model_name) + + if not os.path.exists(model_path): + raise FileNotFoundError(f"Selected local LLM model path not found: {model_path}") + + # Create and return the connector instance. Model loading is deferred. + connector = LocalLLMServiceConnector(model_path) + return (connector,) + + +class LocalLLMServiceConnector: + """ + Represents the connection to a specific local LLM. + Handles loading (on first use) and invocation. + Compatible with TextTranslator and KontextPromptGenerator via the 'invoke' method. + The type identifier used by nodes expecting this connector is "LLMServiceConnector". + """ + def __init__(self, model_path): + self.model_path = model_path + self.model = None + self.tokenizer = None + self.is_loaded = False + + def _load_model(self): + """Loads the model and tokenizer if not already loaded.""" + if self.is_loaded: + return # Already loaded + + try: + log(f"Loading local LLM model from: {self.model_path}") + # --- Model Loading Configuration --- + tokenizer_kwargs = { + "trust_remote_code": True # Needed for some non-standard models + } + + # --- Example Quantization Config (Uncomment and use if needed) --- + # Requires 'bitsandbytes': pip install bitsandbytes + # quantization_config = BitsAndBytesConfig( + # load_in_4bit=True, + # bnb_4bit_compute_dtype=torch.float16, + # bnb_4bit_use_double_quant=True, + # bnb_4bit_quant_type="nf4" + # ) + + model_kwargs = { + "trust_remote_code": True, + # --- Add Quantization Config if using --- + # "quantization_config": quantization_config, # <-- Add for quantization + # --- Device Placement --- + "device_map": "auto", # Automatic device placement (CPU/GPU) - Often crucial with quantization + # --- Precision (if not using quantization config) --- + # "torch_dtype": torch.float16 if torch.cuda.is_available() else torch.float32, + } + + # --- Load Tokenizer --- + self.tokenizer = AutoTokenizer.from_pretrained(self.model_path, **tokenizer_kwargs) + # Handle models without a pad token (common with Llama-based models) + if self.tokenizer.pad_token_id is None: + self.tokenizer.pad_token_id = self.tokenizer.eos_token_id + log(f"Set pad_token_id to eos_token_id ({self.tokenizer.eos_token_id})") + + # --- Load Model --- + # Note: device_map="auto" is often crucial here, especially with quantization + self.model = AutoModelForCausalLM.from_pretrained(self.model_path, **model_kwargs) + + # Manual device placement is usually handled by device_map="auto" + # if "device_map" not in model_kwargs and torch.cuda.is_available(): + # self.model.to("cuda") + # elif "device_map" not in model_kwargs: + # self.model.to("cpu") + + self.is_loaded = True + log(f"Local LLM model loaded successfully from: {self.model_path}") + except Exception as e: + error_msg = f"Failed to load local LLM model from {self.model_path}: {e}" + log(error_msg) + # Re-raise to stop execution if loading fails + raise Exception(error_msg) from e # Chain the exception + + def invoke(self, messages, **generation_kwargs): + """ + Generates text using the local LLM based on the messages list. + Mimics the API expected by TextTranslator/KontextPromptGenerator. + :param messages: List of message dictionaries (like OpenAI format). + Example: [{"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Hello!"}] + :param generation_kwargs: Additional arguments for text generation (e.g., max_new_tokens, temperature, seed). + :return: The generated text string. + """ + try: + if not self.is_loaded: + self._load_model() # Load model on first invocation + + if not self.model or not self.tokenizer: + raise Exception("Local LLM model or tokenizer failed to load.") + + # --- Format messages for the local model --- + # Try using the tokenizer's chat template if available (more robust) + prompt = "" + try: + # Many models come with a chat template + prompt = self.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + if not isinstance(prompt, str): + # Fallback if apply_chat_template behaves unexpectedly + raise ValueError("apply_chat_template did not return a string") + except Exception as e: + # Fallback to simple formatting if chat template fails or isn't available + log(f"Falling back to simple prompt formatting: {e}") + prompt_parts = [] + for message in messages: + role = message.get('role', 'user') + content = message.get('content', '').strip() # Strip content whitespace + if content: # Only add non-empty content lines + prompt_parts.append(f"{role.capitalize()}: {content}") + prompt = "\n".join(prompt_parts) + if prompt: # Add the final prompt indicator only if there's content + prompt += "\nAssistant:" # Encourage assistant response + else: + prompt = "Assistant:" # Fallback if messages were empty + + if not prompt.strip(): + log("Warning: Generated prompt is empty or whitespace.") + return "" # Return empty string if prompt is empty + + # --- Tokenize the prompt --- + try: + inputs = self.tokenizer(prompt, return_tensors="pt") + # Consider moving inputs to model's device if not using device_map="auto" + # if hasattr(self.model, 'device') and self.model.device.type != 'meta': + # inputs = {k: v.to(self.model.device) for k, v in inputs.items()} + except Exception as e: + error_msg = f"Error tokenizing prompt: {e}" + log(error_msg) + raise Exception(error_msg) from e + + # --- Set default generation parameters --- + # These defaults should be reasonable for prompt generation tasks + default_kwargs = { + "max_new_tokens": 250, # Slightly higher default for complex prompts + "temperature": 0.7, # Default creativity + "do_sample": True, # Enable sampling for variety + "pad_token_id": self.tokenizer.pad_token_id if self.tokenizer.pad_token_id is not None else self.tokenizer.eos_token_id, + # "eos_token_id": self.tokenizer.eos_token_id, # Optional: explicitly set EOS + } + # Update defaults with any provided kwargs (e.g., from consuming nodes) + # Filter out kwargs that might cause issues if not explicitly supported + filtered_kwargs = {k: v for k, v in generation_kwargs.items() if k in ['max_new_tokens', 'temperature', 'top_p', 'top_k', 'do_sample', 'num_beams', 'early_stopping', 'pad_token_id', 'eos_token_id', 'seed']} + # Handle 'seed' if passed - PyTorch manual seed (affects stochastic operations) + if 'seed' in generation_kwargs and generation_kwargs['seed'] is not None: + try: + torch.manual_seed(generation_kwargs['seed']) + except Exception as e: + log(f"Warning: Could not set seed: {e}") + + default_kwargs.update(filtered_kwargs) + + # --- Generate text --- + try: + self.model.eval() + with torch.no_grad(): + outputs = self.model.generate(**inputs, **default_kwargs) + except Exception as e: + error_msg = f"Error during model.generate: {e}" + log(error_msg) + raise Exception(error_msg) from e + + # --- Decode the generated tokens --- + try: + # Extract only the newly generated part (skip the input prompt tokens) + generated_tokens = outputs[:, inputs['input_ids'].shape[-1]:] + generated_text = self.tokenizer.decode(generated_tokens[0], skip_special_tokens=True) + # Ensure the output is a clean string + final_text = generated_text.strip() + return final_text # <-- Return ONLY the generated string + + except Exception as e: + error_msg = f"Error decoding generated tokens: {e}" + log(error_msg) + raise Exception(error_msg) from e + + except Exception as e: + # Catch any error that occurred within the try block and log it + error_msg = f"Error in invoke method: {str(e)}" + log(error_msg) + # Re-raise the exception so the calling node knows it failed + raise e # Or raise Exception(error_msg) from e \ No newline at end of file diff --git a/local_prompt_generator.py b/local_prompt_generator.py new file mode 100644 index 0000000..b1255eb --- /dev/null +++ b/local_prompt_generator.py @@ -0,0 +1,255 @@ +# ComfyUI/custom_nodes/ComfyUI_LocalLLMNodes/local_prompt_generator.py +import hashlib +import os +import json + +# --- Simple logging utility for this package --- +def log(message): + print(f"[LocalPromptGen] {message}") + +# --- Replicate necessary preset logic locally --- +# We need to replicate the core preset loading/management logic here. + +# --- Define paths and functions for user presets --- +script_directory = os.path.dirname(os.path.abspath(__file__)) +USER_PRESETS_FILE = os.path.join(script_directory, "user_kontext_presets.json") + +# --- Replicate Built-in Presets (from MieNodes' prompt_generator.py) --- +# Ensure these match the original structure exactly. +KONTEXT_PRESETS = { + "Kontext Standard": { + "system": "You are an expert AI assistant specializing in generating detailed, creative prompts for image generation models. Your task is to take user-provided image descriptions and edit instructions, then synthesize them into a single, highly descriptive prompt optimized for image generation. Focus on integrating key visual elements like character features, clothing details, scene setting, lighting, and artistic style. Ensure the final prompt is concise, avoids redundancy, and clearly conveys the user's intent for the desired image output." + }, + "Kontext Detailed": { + "system": "You are a master prompt engineer for AI image generators. Your role is to meticulously craft prompts by deeply analyzing user inputs. Given descriptions of two images (e.g., a person and clothing) and specific edit instructions, you must seamlessly merge these elements. Prioritize descriptive keywords for physical attributes, textures, colors, background environments, and artistic influences. The resulting prompt should be rich in detail, logically structured, and highly effective at guiding the image generator to produce the envisioned scene with high fidelity." + }, + "Kontext Minimalist": { + "system": "You are an AI assistant focused on creating concise, clear prompts for image generation. Your task is to distill user-provided descriptions and edit instructions into a short, essential prompt. Identify the core subject, key action or interaction, and the most important visual style or setting. Eliminate unnecessary details and focus on the primary elements that define the scene. The output prompt should be direct and easy for the image generator to interpret accurately." + }, + "Kontext Artistic Style Focus": { + "system": "You are an AI prompt specialist with an emphasis on artistic style and rendering techniques. Users will provide descriptions of elements and specific edits. Your goal is to construct a prompt that heavily emphasizes the desired artistic style (e.g., 'oil painting', 'cyberpunk', 'watercolor', 'photorealistic'). Integrate the provided subject and scene details, but frame them within the context of the specified artistic approach. Highlight relevant techniques, color palettes, brushwork, or visual effects associated with that style to guide the image generator effectively." + } +} + +def load_user_presets(): + """Load user-defined presets from the JSON file.""" + if os.path.exists(USER_PRESETS_FILE): + try: + with open(USER_PRESETS_FILE, "r", encoding="utf-8") as f: + return json.load(f) + except (json.JSONDecodeError, IOError) as e: + log(f"Error loading user presets: {e}") + return {} + return {} + +def get_all_kontext_presets(): + """Get a combined dictionary of built-in and user presets.""" + # Start with built-in presets + all_presets = KONTEXT_PRESETS.copy() + # Load and add user presets, potentially overriding built-ins if names clash (usually desired) + user_presets = load_user_presets() + all_presets.update(user_presets) + return all_presets + +# Define category directly for this package +LOCAL_PROMPT_CATEGORY = "Local LLM Nodes/Prompt Generators" + +class LocalKontextPromptGenerator(object): + """ + A version of KontextPromptGenerator designed to work specifically with + the LocalLLMServiceConnector. Replicates the core prompt generation logic exactly. + """ + @classmethod + def INPUT_TYPES(cls): + # --- Replicate INPUT_TYPES exactly --- + all_presets = get_all_kontext_presets() + # Ensure there's a default if the list is somehow empty + default_preset_key = next(iter(all_presets.keys()), "") if all_presets.keys() else "" + return { + "required": { + # --- CRITICAL: Use the standard LLM_SERVICE_CONNECTOR type --- + # This matches the RETURN_TYPES of SetLocalLLMServiceConnector + # and is the same type expected by the original KontextPromptGenerator. + "llm_service_connector": ("LLMServiceConnector",), + "image1_description": ("STRING", {"default": "", "multiline": True, "tooltip": "Describe the first image"}), + "image2_description": ("STRING", {"default": "", "multiline": True, "tooltip": "Describe the second image"}), + "edit_instruction": ("STRING", {"default": "", "multiline": True}), + "preset": (list(all_presets.keys()), {"default": default_preset_key}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff, "control_after_generate": True, + "tooltip": "The random seed used for creating the noise."}), + }, + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("kontext_prompt",) + FUNCTION = "generate_kontext_prompt" + CATEGORY = LOCAL_PROMPT_CATEGORY + + # --- Replicate the core generate_kontext_prompt method exactly --- + def generate_kontext_prompt(self, llm_service_connector, image1_description, image2_description, edit_instruction, preset, seed=None): + """ + Replicates the exact core logic from KontextPromptGenerator.generate_kontext_prompt. + """ + # --- Exact replication of the original logic --- + all_presets = get_all_kontext_presets() + preset_data = all_presets.get(preset) + # --- Fixed Syntax Error --- + if not preset_data: # <-- Corrected condition + raise ValueError(f"Unknown preset: {preset}") + + # 用户输入拼到user消息中,给LLM最大上下文 + user_content = "" + if image1_description.strip(): + user_content += f"Image 1 (person) description: {image1_description.strip()}" + if image2_description.strip(): + user_content += f"Image 2 (clothing) description: {image2_description.strip()}" + if edit_instruction.strip(): + user_content += f"Edit instruction: {edit_instruction.strip()}" + if not user_content.strip(): + user_content = "No additional image description or edit instruction provided." + + messages = [ + {"role": "system", "content": preset_data["system"]}, + {"role": "user", "content": user_content}, + ] + + # --- Key Change: Use the local connector's invoke method --- + # Pass the messages and seed. + try: + kontext_prompt = llm_service_connector.invoke(messages, seed=seed) + # Ensure the output is stripped, like the original + return (kontext_prompt.strip(),) + except Exception as e: + # Handle potential errors during local LLM invocation + error_msg = f"Error generating kontext prompt with local LLM: {str(e)}" + log(error_msg) # Log to console + # Return the error message as the prompt string so the workflow doesn't crash silently + return (error_msg,) + + # --- Replicate the is_changed method for proper caching --- + def is_changed(self, llm_service_connector, image1_description, image2_description, edit_instruction, preset, seed): + """ + Replicates the exact core logic from KontextPromptGenerator.is_changed. + """ + try: + hasher = hashlib.md5() + hasher.update(image1_description.encode('utf-8')) + hasher.update(image2_description.encode('utf-8')) + hasher.update(edit_instruction.encode('utf-8')) + hasher.update(preset.encode('utf-8')) + hasher.update(str(seed).encode('utf-8')) + + # Incorporate preset system prompt + all_presets = get_all_kontext_presets() + preset_data = all_presets.get(preset) + # --- Fixed Syntax Error --- + if preset_data and "system" in preset_data: # <-- Corrected condition and variable name + hasher.update(preset_data["system"].encode('utf-8')) + + # Incorporate connector state (replicate original logic) + # Original tries get_state(), then falls back to specific attributes. + # For a local connector, str(connector) is suitable. + try: + # Try get_state if it exists on the local connector (unlikely for our simple one) + connector_state = llm_service_connector.get_state() + except AttributeError: + # Fallback: Use a string representation of the connector object + # This works if the connector's identity is tied to its instance/model path. + connector_state = str(llm_service_connector) + + hasher.update(connector_state.encode('utf-8')) + + return hasher.hexdigest() + except Exception as e: + # If hashing fails, force re-run + log(f"is_changed error: {e}") + return float("nan") # Always re-run + +# --- Replicate the AddUserKontextPreset and RemoveUserKontextPreset classes --- +# These manage the local user presets file. + +class AddUserLocalKontextPreset: # <-- Changed class name prefix + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "preset_name": ("STRING", {"default": ""}), + "system_prompt": ("STRING", {"default": "", "multiline": True}), + } + } + + RETURN_TYPES = ("BOOLEAN", "STRING") + RETURN_NAMES = ("success", "log") + FUNCTION = "add_preset" + CATEGORY = LOCAL_PROMPT_CATEGORY + + def add_preset(self, preset_name, system_prompt): + import datetime + if not preset_name or not system_prompt: + log = "Preset name and system prompt must not be empty." + return (False, log) + + user_presets = load_user_presets() + if preset_name in user_presets: + log = f"Preset '{preset_name}' already exists (custom preset)." + return (False, log) + + user_presets[preset_name] = {"system": system_prompt} + # Save to the local user presets file + try: + with open(USER_PRESETS_FILE, "w", encoding="utf-8") as f: + json.dump(user_presets, f, ensure_ascii=False, indent=2) + now = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") + log = f"Preset '{preset_name}' added successfully at {now}." + return (True, log) + except Exception as e: + log = f"Error saving preset: {e}" + return (False, log) + + @classmethod + def IS_CHANGED(cls, **kwargs): + # Always re-run when called to check for file changes + return float("nan") + +class RemoveUserLocalKontextPreset: # <-- Changed class name prefix + @classmethod + def INPUT_TYPES(cls): + # Load current presets to populate the dropdown + user_presets = load_user_presets() + preset_names = list(user_presets.keys()) + # Provide a default if the list is empty + default_name = preset_names[0] if preset_names else "" + return { + "required": { + "preset_name": (preset_names, {"default": default_name}), + } + } + + RETURN_TYPES = ("BOOLEAN", "STRING") + RETURN_NAMES = ("success", "log") + FUNCTION = "remove_preset" + CATEGORY = LOCAL_PROMPT_CATEGORY + + def remove_preset(self, preset_name): + import datetime + user_presets = load_user_presets() + if preset_name in user_presets: + del user_presets[preset_name] + try: + # Save to the local user presets file + with open(USER_PRESETS_FILE, "w", encoding="utf-8") as f: + json.dump(user_presets, f, ensure_ascii=False, indent=2) + now = datetime.datetime.now().strftime("%Y-%m-%d %H:%M:%S") + log = f"Preset '{preset_name}' removed successfully at {now}." + return (True, log) + except Exception as e: + log = f"Error saving after removal: {e}" + return (False, log) + else: + log = f"Preset '{preset_name}' not found in user presets." + return (False, log) + + @classmethod + def IS_CHANGED(cls, **kwargs): + # Always re-run when called to check for file changes + return float("nan") diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..76948ca --- /dev/null +++ b/requirements.txt @@ -0,0 +1,5 @@ +# ComfyUI/custom_nodes/ComfyUI_LocalLLMNodes/requirements.txt +transformers +torch +# bitsandbytes # Uncomment if you implement quantization +# accelerate # Often needed with transformers, especially for device_map \ No newline at end of file