From 0c456a67f82de2003e04965fd2bfcb0f8586351d Mon Sep 17 00:00:00 2001 From: Arctenox <69485661+Arctenox@users.noreply.github.com> Date: Mon, 5 Jan 2026 21:07:19 -0500 Subject: [PATCH] Update TensorPrism_Enhancer.py --- TensorPrism_Enhancer.py | 159 ++++++++++++++-------------------------- 1 file changed, 55 insertions(+), 104 deletions(-) diff --git a/TensorPrism_Enhancer.py b/TensorPrism_Enhancer.py index d86a074..37714c1 100644 --- a/TensorPrism_Enhancer.py +++ b/TensorPrism_Enhancer.py @@ -17,7 +17,7 @@ class TargetModule(str, Enum): """Defines the target modules within the checkpoint for enhancement.""" UNET = 'unet' VAE = 'vae' - TEXT_ENCODERS = 'text_encoders' # Renamed from 'textenc' for clarity + TEXT_ENCODERS = 'text_encoders' ALL = 'all' # Helper functions to identify keys belonging to specific model components @@ -48,7 +48,6 @@ def clamp_tensor_stats(tensor: torch.Tensor, max_abs_threshold: float = 1.0) -> device = tensor.device max_val = tensor.abs().max() if max_val > max_abs_threshold: - # Use torch.finfo for robust epsilon value based on tensor dtype epsilon = torch.finfo(tensor.dtype).eps scale = max_abs_threshold / (max_val + epsilon) return (tensor * scale).to(device) @@ -84,21 +83,14 @@ def attention_refinement_stub(tensor: torch.Tensor, strength: float) -> torch.Te Placeholder for attention-based refinement. This function emulates an attention-like effect by applying a guided non-linear boost on high-magnitude elements, followed by subtle smoothing to prevent harsh artifacts. - For a real implementation, this would involve invoking attention maps or - cross-attention reweighting using model internals. - Effective when `strength` is greater than 0. """ if strength <= 0.0: return tensor device = tensor.device - # Emulate attention emphasis by applying a guided non-linear boost on high-magnitude elements magnitudes = tensor.abs() - # Find the 75th percentile of magnitudes to identify "important" features percentile_threshold = torch.quantile(magnitudes.view(-1), 0.75).to(device) - # Create a mask for elements above the threshold mask = (magnitudes >= percentile_threshold).to(tensor.dtype).to(device) - # Boost these elements, then apply subtle smoothing to prevent harsh artifacts boosted_tensor = tensor + (tensor * mask * strength * 0.75) return apply_linear_smoothing(boosted_tensor, min(0.12 * strength, 0.5)) @@ -112,9 +104,7 @@ def apply_quality_boost(tensor: torch.Tensor, strength: float) -> torch.Tensor: return tensor device = tensor.device mean_value = tensor.mean().to(device) - # Increase local contrast around the mean boosted_tensor = (tensor - mean_value) * (1.0 + 0.6 * strength) + mean_value - # Apply gentle non-linear sharpening based on deviation from mean boosted_tensor = boosted_tensor * (1.0 + 0.12 * strength * (boosted_tensor - mean_value)) return boosted_tensor.to(device) @@ -127,11 +117,10 @@ def apply_adaptive_overbake_limiter( """ Compares the modified tensor to the original. If the global mean absolute change exceeds `max_relative_increase`, the modification is scaled back - to avoid overbaking. Prevents excessively strong modifications. + to avoid overbaking. """ device = original_tensor.device with torch.no_grad(): - # Add epsilon to prevent division by zero for tensors with all zeros epsilon = torch.finfo(original_tensor.dtype).eps original_mean_abs = original_tensor.abs().mean().item() + epsilon modified_mean_abs = modified_tensor.abs().mean().item() + epsilon @@ -141,12 +130,9 @@ def apply_adaptive_overbake_limiter( if relative_change <= max_relative_increase: return modified_tensor.to(device) - # Calculate a scale factor to bring the relative change down to the limit - # Ensure scale_back_factor is not negative scale_back_factor = 1.0 - (relative_change - max_relative_increase) / (relative_change + epsilon) scale_back_factor = max(0.0, scale_back_factor) - # Blend the original with the modified based on the scale_back_factor return (original_tensor + (modified_tensor - original_tensor) * scale_back_factor).to(device) @@ -157,12 +143,11 @@ class ModelEnhancerTensorPrism: to improve aspects like smoothing, sharpening, and overall quality. """ - # Define the input types for the ComfyUI node UI @classmethod def INPUT_TYPES(cls): return { "required": { - "checkpoint": ("MODEL",), # Input checkpoint (MODEL object) + "checkpoint": ("MODEL",), "smoothing": ("FLOAT", {"default": 0.12, "min": 0.0, "max": 1.0, "step": 0.01}), "sharpening": ("FLOAT", {"default": 0.12, "min": 0.0, "max": 1.0, "step": 0.01}), "quality_boost": ("FLOAT", {"default": 0.08, "min": 0.0, "max": 1.0, "step": 0.01}), @@ -184,27 +169,22 @@ class ModelEnhancerTensorPrism: } } - # Define the output types for the ComfyUI node UI RETURN_TYPES = ("MODEL",) RETURN_NAMES = ("enhanced_checkpoint",) FUNCTION = "enhance_checkpoint" - CATEGORY = "checkpoint/enhancement" # Or a suitable category + CATEGORY = "checkpoint/enhancement" def __init__(self): - # ComfyUI nodes typically don't need a complex __init__ if all parameters - # are passed via the INPUT_TYPES to the functional method. pass def _should_process_key(self, key: str, target_modules) -> bool: """ Determines if a given tensor key should be processed based on the - selected target modules. Now handles both single strings and lists. + selected target modules. """ - # Handle both single string and list inputs if isinstance(target_modules, str): target_modules = [target_modules] - # Convert target_modules to a set of lowercased strings for efficient lookup target_modules_set = {m.lower() for m in target_modules} if TargetModule.ALL.value in target_modules_set: @@ -213,7 +193,6 @@ class ModelEnhancerTensorPrism: return True if is_vae_key(key) and TargetModule.VAE.value in target_modules_set: return True - # Check for both new enum value and original string values for compatibility if is_text_encoder_key(key) and (TargetModule.TEXT_ENCODERS.value in target_modules_set or 'textenc' in target_modules_set or 'text' in target_modules_set): return True return False @@ -228,7 +207,7 @@ class ModelEnhancerTensorPrism: return tensor.half().to(device) elif precision == CheckpointPrecision.FP32: return tensor.float().to(device) - return tensor.to(device) # Should not happen if enum is used correctly + return tensor.to(device) def _enhance_single_tensor( self, @@ -243,39 +222,28 @@ class ModelEnhancerTensorPrism: ) -> torch.Tensor: """Applies the enhancement pipeline to a single tensor.""" - # Ensure we operate on a clone and on the correct device device = original_tensor.device processed_tensor = original_tensor.clone().to(device) - # Cast precision upfront for operations processed_tensor = self._cast_tensor_to_precision(processed_tensor, target_precision) - # Apply chosen enhancement method if enhancement_method == EnhancementMethod.LINEAR: if smoothing_strength > 0.0: processed_tensor = apply_linear_smoothing(processed_tensor, smoothing_strength) if sharpening_strength > 0.0: processed_tensor = apply_linear_sharpen(processed_tensor, sharpening_strength) elif enhancement_method == EnhancementMethod.ATTENTION: - # Run attention-based refinement for specified iterations for _ in range(max(1, attention_iterations)): combined_strength = (sharpening_strength + smoothing_strength) * 0.7 processed_tensor = attention_refinement_stub(processed_tensor, combined_strength) - # Apply overall quality boost if quality_boost_strength > 0.0: processed_tensor = apply_quality_boost(processed_tensor, quality_boost_strength) - # Clamp extreme values softly, relative to original max absolute value - # Using a factor (e.g., 3.0) to allow for growth but prevent explosion max_abs_ref_val = original_tensor.abs().max().item() - # Adding a small epsilon to max_abs_ref_val in case it's zero processed_tensor = clamp_tensor_stats(processed_tensor, max_abs_threshold=max_abs_ref_val * 3.0 + 1e-6) - # Apply adaptive overbake limiter if enabled if adaptive_overbake_enabled: - # Use a slightly higher max_relative_increase for the internal processing step - # to allow for more aggressive enhancement before final blending. processed_tensor = apply_adaptive_overbake_limiter( original_tensor, processed_tensor, max_relative_increase=0.16 ) @@ -284,97 +252,80 @@ class ModelEnhancerTensorPrism: def enhance_checkpoint( self, - checkpoint, # Changed from dict to accept MODEL objects + checkpoint, smoothing: float, sharpening: float, quality_boost: float, blend_strength: float, - precision: str, # Will be enum value string from ComfyUI - method: str, # Will be enum value string from ComfyUI - modules_to_enhance, # Can be string or list[str] + precision: str, + method: str, + modules_to_enhance, adaptive_overbake_prevention: bool, attention_iterations: int - ) -> tuple: # Changed return type + ) -> tuple: """ Main function to process and enhance a ComfyUI checkpoint (MODEL object). - This method is called by ComfyUI when the node executes. """ if checkpoint is None: raise ValueError("Input checkpoint cannot be None.") - # Extract state_dict from ComfyUI MODEL object - if hasattr(checkpoint, 'model') and hasattr(checkpoint.model, 'state_dict'): - # ComfyUI MODEL object - extract the state_dict - state_dict = checkpoint.model.state_dict() - model_wrapper = checkpoint - elif isinstance(checkpoint, dict): - # Already a state_dict (for backwards compatibility) - state_dict = checkpoint - model_wrapper = None - else: - raise TypeError("Input checkpoint must be a ComfyUI MODEL object or dictionary-like state_dict.") + # ComfyUI MODEL object - work with it directly, don't deep copy + if not hasattr(checkpoint, 'model'): + raise TypeError("Input checkpoint must be a ComfyUI MODEL object.") - # Convert string inputs from ComfyUI to Enum types for type safety and clarity + # Convert string inputs to Enum types try: target_precision = CheckpointPrecision(precision) enhancement_method = EnhancementMethod(method) except ValueError as e: raise ValueError(f"Invalid enum value provided for precision or method: {e}") - # Deep copy the state_dict to ensure no in-place modification of the original - # This can be memory-intensive for very large checkpoints, but ensures safety. - enhanced_state_dict = copy.deepcopy(state_dict) + # Clone the model using ComfyUI's clone method (this is safe) + enhanced_model = checkpoint.clone() - # Iterate over all keys in the state_dict to apply enhancements - for key, value in enhanced_state_dict.items(): - try: - # Process only PyTorch tensors that are floating-point and match selected modules - if self._should_process_key(key, modules_to_enhance) and isinstance(value, torch.Tensor) and value.is_floating_point(): - original_tensor = state_dict[key] # Reference the original tensor from the extracted state_dict - device = original_tensor.device + # Get a reference to the actual model weights + model_sd = enhanced_model.model.state_dict() - # Apply the full enhancement pipeline to the current tensor - enhanced_tensor = self._enhance_single_tensor( - original_tensor=original_tensor, - smoothing_strength=smoothing, - sharpening_strength=sharpening, - quality_boost_strength=quality_boost, - enhancement_method=enhancement_method, - adaptive_overbake_enabled=adaptive_overbake_prevention, - attention_iterations=attention_iterations, - target_precision=target_precision - ) + # Process tensors in-place on the cloned model + with torch.no_grad(): + for key in list(model_sd.keys()): + try: + value = model_sd[key] + + if self._should_process_key(key, modules_to_enhance) and isinstance(value, torch.Tensor) and value.is_floating_point(): + original_tensor = value + device = original_tensor.device - # Ensure enhanced tensor is on correct device - enhanced_tensor = enhanced_tensor.to(device) + # Apply enhancement pipeline + enhanced_tensor = self._enhance_single_tensor( + original_tensor=original_tensor, + smoothing_strength=smoothing, + sharpening_strength=sharpening, + quality_boost_strength=quality_boost, + enhancement_method=enhancement_method, + adaptive_overbake_enabled=adaptive_overbake_prevention, + attention_iterations=attention_iterations, + target_precision=target_precision + ) - # Blend the original tensor with the enhanced tensor based on blend_strength - blended_tensor = original_tensor * (1.0 - blend_strength) + enhanced_tensor * blend_strength + enhanced_tensor = enhanced_tensor.to(device) - # Apply a final, robust clamp to the blended tensor to maintain numeric stability. - # The max_abs_threshold is set to be at least 1.0 or twice the original's max abs. - final_max_abs_threshold = max(1.0, original_tensor.abs().max().item() * 2.0) - enhanced_state_dict[key] = clamp_tensor_stats(blended_tensor, max_abs_threshold=final_max_abs_threshold).to(device) - else: - # If not a tensor, not floating point, or not selected for processing, - # ensure the original value is preserved (even though deepcopy usually handles this). - enhanced_state_dict[key] = value + # Blend original with enhanced + blended_tensor = original_tensor * (1.0 - blend_strength) + enhanced_tensor * blend_strength - except Exception as e: - # Log the error and revert to the original tensor for this specific key - # This prevents a single problematic tensor from crashing the entire node. - print(f"Warning: Model Enhancer failed to process tensor '{key}'. Keeping original value. Error: {e}") - enhanced_state_dict[key] = state_dict[key] # Ensure original is used if processing failed + # Final clamp + final_max_abs_threshold = max(1.0, original_tensor.abs().max().item() * 2.0) + final_tensor = clamp_tensor_stats(blended_tensor, max_abs_threshold=final_max_abs_threshold).to(device) + + # Update the tensor in the state dict + model_sd[key] = final_tensor + + except Exception as e: + print(f"Warning: Model Enhancer failed to process tensor '{key}'. Keeping original value. Error: {e}") + continue + + return (enhanced_model,) - # Return the enhanced model in ComfyUI format - if model_wrapper is not None: - # Create a new model wrapper with the enhanced state_dict - enhanced_model = copy.deepcopy(model_wrapper) - enhanced_model.model.load_state_dict(enhanced_state_dict, strict=False) - return (enhanced_model,) - else: - # For backwards compatibility, return the state_dict - return (enhanced_state_dict,) # ComfyUI Node Class Mappings for registration NODE_CLASS_MAPPINGS = { @@ -383,4 +334,4 @@ NODE_CLASS_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = { "ModelEnhancerTensorPrism": "Model Enhancer (Tensor Prism)" -} \ No newline at end of file +}