diff --git a/brekel_enhance_prompt.py b/brekel_enhance_prompt.py index 7d796aa..d03f265 100644 --- a/brekel_enhance_prompt.py +++ b/brekel_enhance_prompt.py @@ -12,7 +12,7 @@ # - Logs the model's memory footprint upon loading. # - Defaults to SDPA (Scaled Dot-Product Attention) for optimal memory and speed. # - Flexible memory management with performance reporting for offloading. -# - Controllable creativity, seed, and max_length for prompt generation. +# - Controllable creativity, seed, and target_length for prompt generation. import os @@ -158,7 +158,7 @@ class BrekelEnhancePrompt: "quantization": (["Disabled", "8-bit (Int8)", "4-bit (NF4)"], {"default": "Disabled", "tooltip": "Quantization method to save VRAM."}), "memory_management": (["Offload to CPU after use", "Keep in VRAM", "Unload completely after use"], {"tooltip": "Memory management strategy."}), "system_prompt": (system_prompt_names, {"tooltip": "System prompt from files in the 'prompt_enhancer' subfolder."}), - "max_length": ("INT", {"default": 128, "min": 64, "max": 512, "step": 8, "display": "slider", "tooltip": "Maximum length of the generated prompt."}), + "target_length": ("INT", {"default": 150, "min": 64, "max": 512, "step": 8, "display": "slider", "tooltip": "Target length for the generated prompt. Used as a goal (characters) and a safe token limit."}), "creativity": ("FLOAT", {"default": 0.8, "min": 0.0, "max": 2.0, "step": 0.1, "display": "slider", "tooltip": "Creativity level (temperature)."}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "postfix": ("STRING", {"multiline": False, "default": "", "tooltip": "Postfix to append at the end of the prompt, for example to add your Lora trigger word(s)."}), @@ -304,7 +304,7 @@ class BrekelEnhancePrompt: return cleaned_text - def brekel_enhance_prompt(self, prompt: str, prefix: str, model_name: str, quantization: str, memory_management: str, system_prompt: str, max_length: int, creativity: float, seed: int, postfix: str): + def brekel_enhance_prompt(self, prompt: str, prefix: str, model_name: str, quantization: str, memory_management: str, system_prompt: str, target_length: int, creativity: float, seed: int, postfix: str): """ The main execution function of the node. It orchestrates loading the model, generating the enhanced prompt, and handling memory management. @@ -322,59 +322,51 @@ class BrekelEnhancePrompt: raise ValueError(f"Selected system prompt '{system_prompt}' could not be loaded. Check for errors above.") # --- MODEL LOADING --- - # Ensure the correct model is loaded and ready on the GPU. self._load_model(model_name, quantization) if _loaded_model is None or _loaded_tokenizer is None: raise RuntimeError("Model and/or tokenizer failed to load, but no exception was caught. Check logs.") - # Set the random seed for reproducibility if not set to 0. if seed != 0: torch.manual_seed(seed) # --- PROMPT GENERATION --- logger.info(f"\n--- Enhancing Prompt ---") logger.info(f"Original: '{prompt}'") - logger.info(f"Settings: model={model_name}, creativity={creativity}, max_length={max_length}, seed={'random' if seed == 0 else seed}") + logger.info(f"Settings: model={model_name}, creativity={creativity}, target_length={target_length}, seed={'random' if seed == 0 else seed}") - # Dynamically add the length constraint to the system prompt. - # This instructs the LLM to be concise, rather than just cutting off its output. - length_instruction = f"\n\nIMPORTANT: The total length of your enhanced prompt should not exceed {max_length} characters. Be concise and impactful." + length_instruction = ( + f"\n\nIMPORTANT: Expand the user's idea into a rich, detailed, and evocative prompt for a text-to-image model. " + f"Aim for a total length of approximately {target_length} characters. " + "Add creative details about lighting, style, composition, and mood. " + "Your response must be ONLY the enhanced prompt itself, without any conversational introduction." + ) final_system_prompt = active_system_prompt + length_instruction + logger.info(f"\nSystem prompt: '{final_system_prompt}'\n") - # Format the input for the LLM using a chat template. messages = [{"role": "system", "content": final_system_prompt}, {"role": "user", "content": f"user_prompt: {prompt}"}] text_input_for_llm = _loaded_tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) model_inputs = _loaded_tokenizer(text_input_for_llm, return_tensors="pt").to(_loaded_model.device) - # Set generation parameters based on user inputs. - # The 'max_new_tokens' acts as a safety net in case the LLM ignores the instruction. - generation_params = {"max_new_tokens": max_length, "pad_token_id": _loaded_tokenizer.eos_token_id} + # Set generation parameters. We only need `max_new_tokens` to control the output length. + generation_params = {"max_new_tokens": target_length, "pad_token_id": _loaded_tokenizer.eos_token_id} + if creativity > 0.0: - # Use sampling for creative, non-deterministic output. generation_params.update({"do_sample": True, "temperature": creativity, "top_p": 0.9}) else: - # Use greedy decoding for deterministic output. generation_params["do_sample"] = False - # Generate the prompt in inference mode to save memory and improve speed. with torch.inference_mode(): - # Suppress a common but harmless warning from the generate function. with warnings.catch_warnings(): warnings.filterwarnings("ignore", "The following generation flags are not valid.*") outputs = _loaded_model.generate(**model_inputs, **generation_params) - # Decode the generated token IDs into a text string. - generated_ids = outputs[0][len(model_inputs.input_ids[0]):] # Exclude the input prompt tokens. + input_token_length = model_inputs.input_ids.shape[-1] + generated_ids = outputs[0][input_token_length:] raw_enhanced_prompt = _loaded_tokenizer.decode(generated_ids, skip_special_tokens=True) - # Clean the raw output to remove LLM conversational boilerplate. enhanced_prompt = self._clean_llm_output(raw_enhanced_prompt) - - # Final cleanup: remove extra spaces, quotes, and trailing punctuation. enhanced_prompt = enhanced_prompt.strip().strip('"').rstrip(' ,.').strip() - # Robustly combine prefix, enhanced prompt, and postfix. - # This filters out empty parts to avoid extra spaces. prompt_parts = [p.strip() for p in [prefix, enhanced_prompt, postfix] if p.strip()] final_prompt = ", ".join(prompt_parts) @@ -383,16 +375,12 @@ class BrekelEnhancePrompt: except Exception as e: logger.error(f"!!! PROMPT ENHANCEMENT FAILED: {e}") - raise # Re-raise the exception to make the node fail visibly in ComfyUI. + raise finally: - # This block always executes, whether the `try` block succeeded or failed. - # It's used here for memory management cleanup. if _loaded_model is not None: if memory_management == "Offload to CPU after use": - # If not already offloaded, move the model to CPU to free up VRAM. if not _model_is_offloaded: - # Only non-quantized models on CUDA need manual moving. if quantization == "Disabled" and _loaded_model.device.type == 'cuda': logger.info(f"Offloading model '{_loaded_model_name}' to CPU RAM...") start_time = time.perf_counter() @@ -404,14 +392,12 @@ class BrekelEnhancePrompt: _model_is_offloaded = True elif memory_management == "Unload completely after use": - # Remove the model and tokenizer from memory entirely. logger.info(f"Unloading model '{_loaded_model_name}' from all memory...") del _loaded_model, _loaded_tokenizer _loaded_model = _loaded_tokenizer = _loaded_model_name = _loaded_quantization = None _model_is_offloaded = False torch.cuda.empty_cache() - # --- ComfyUI Node Registration --- NODE_CLASS_MAPPINGS = {"BrekelEnhancePrompt": BrekelEnhancePrompt} NODE_DISPLAY_NAME_MAPPINGS = {"BrekelEnhancePrompt": "Brekel Prompt Enhancer (LLM)"} \ No newline at end of file