Update TensorPrism_Enhancer.py

This commit is contained in:
Arctenox
2026-01-05 21:07:19 -05:00
committed by GitHub
parent 57c42ee3de
commit 0c456a67f8
+55 -104
View File
@@ -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)"
}
}