Add files via upload

This commit is contained in:
Arctenox
2025-09-15 02:42:31 -04:00
committed by GitHub
parent 2e7e050fa6
commit 19331e87a4
17 changed files with 4154 additions and 0 deletions
+315
View File
@@ -0,0 +1,315 @@
### ComfyUI-Tensor-Prism-Node-Pack
## Developer Notes
My First ComfyUI Node Pack, vibe coded with Gemini 2.5 Flash and Claude 4. Feel free to publish the models you make and link them to me I'd like to be able to see the models, and see what they're about to see if I need to add more nodes or if the nodes are good and make really good quality checkpoint models. This is also a node pack for those familiar with merging models.
# TensorPrism ComfyUI Node Pack
Advanced model merging and enhancement nodes for ComfyUI, providing sophisticated techniques for blending, enhancing, and manipulating Stable Diffusion models with GPU-optimized memory management.
## Features
### Core Merging Nodes
- **Main Merge**: Advanced model merging with multiple interpolation methods (linear, slerp, cosine, directional, frequency, stochastic)
- **Prism**: Fast spectral merging with frequency-based blending techniques including spectral blend, frequency bands, magnitude weighting, adaptive mixing, and harmonic merging
- **SDXL Block Merge**: Granular control over individual SDXL UNet blocks with support for TIES merging
- **SDXL Advanced Block Merge**: GPU-optimized block merging with intelligent memory management for any GPU size (including 12GB and smaller cards)
- **Epsilon/V-Pred Block Merge**: Granular block-level model merging with V-Pred/Epsilon conversion and individual control for input blocks 0-8, middle blocks 0-2, and output blocks 0-8
### Conversion and Processing Nodes
- **Epsilon/V-Pred Converter**: Pure converter between V-Prediction and Epsilon prediction types with configurable conversion strength and smart layer targeting
- **V-Pred/Epsilon Converter**: Advanced prediction type conversion with critical layer identification (output blocks, middle blocks) and secondary layer processing (time embedding, input blocks)
### Advanced Mask System
- **Model Mask Generator**: Create sophisticated masks for selective model merging with layer-based, block-based, attention-only, feedforward-only, custom patterns, random sparse, and depth gradient options
- **Weighted Mask Merge**: Apply masks to control where and how models are blended with tensor-level precision
- **Model Key Filter**: Memory-efficient filtering of model parameters with batch processing for large models
- **Mask Blender**: Combine multiple masks with various blending modes (Add, Multiply, Max, Min, Linear Blend, Exponential Blend)
### CLIP Processing
- **Advanced CLIP Merge**: Sophisticated CLIP merging with multiple interpolation methods:
- **Linear**: Standard interpolation
- **SLERP**: Spherical Linear Interpolation for smooth blending
- **Cosine**: Cosine-based smooth transitions
- **Weighted Average**: Magnitude-based weighting
- **Spectral Blend**: Frequency domain blending with separate magnitude/phase control
- Layer-specific bias controls for attention, feedforward, embedding, and normalization layers
- Norm preservation options for maintaining model stability
### Model Transformation
- **Model Weight Modifier**: Memory-efficient weight modification with operations like multiply, add, set value, clamp magnitude, and scale max absolute value
### Utility Nodes
- **Checkpoint Reroute + Notes**: Convenient rerouting node for MODEL, CLIP, and VAE connections with optional note-taking functionality for workflow organization and documentation
### Advanced Features
- **GPU-Optimized Memory Management**: Intelligent memory allocation and cleanup for cards with limited VRAM
- **Batch Processing**: Process large models in memory-efficient batches
- **Multiple Merging Algorithms**: Including SLERP, frequency domain blending, stochastic merging, and TIES merging
- **Spectral Analysis**: Frequency-based model analysis and merging
- **Adaptive Precision**: Automatic FP16/FP32 selection based on available memory
- **Cross-Device Compatibility**: Works with CUDA, MPS, and CPU backends
- **Prediction Type Conversion**: Automatic conversion between Epsilon and V-Prediction model types with mathematical accuracy
- **Advanced Memory Context Management**: Context managers for automatic cleanup and memory optimization
- **Workflow Organization**: Built-in documentation and organization tools for complex merging workflows
## Installation
### Method 1: ComfyUI Manager
Note: Not available yet there.
1. Open ComfyUI Manager
2. Search for "TensorPrism"
3. Click Install
### Method 2: Manual Installation (OPTIONAL)
1. Clone or download this repository
2. Place the entire folder in your `ComfyUI/custom_nodes/` directory
3. Restart ComfyUI
### Method 3: Git Clone (OPTIONAL)
```bash
cd ComfyUI/custom_nodes/
git clone https://github.com/AstrionX/ComfyUI-Tensor-Prism-Node-Pack.git
```
## Usage
### Basic Model Merging
Use the **Main Merge** node to blend two models with various interpolation methods:
- Connect two MODEL inputs
- Adjust merge ratio (0.0 = Model A only, 1.0 = Model B only)
- Choose merging method based on your needs:
- **Linear**: Standard interpolation
- **SLERP**: Spherical interpolation for smoother blending
- **Cosine**: Smoother transitions with cosine curve
- **Directional**: Vector-based interpolation
- **Frequency**: FFT-based frequency domain blending
- **Stochastic**: Random pattern-based merging
### Advanced Spectral Merging
The **Prism** node offers frequency-domain merging:
- **Spectral Blend**: Magnitude-based frequency separation
- **Frequency Bands**: Different weights for low/high frequency components
- **Magnitude Weighted**: Stronger tensor gets more influence
- **Adaptive Mix**: Similarity-based adaptive merging
- **Harmonic Merge**: Phase relationship-based blending
### Granular SDXL Block Control
**SDXL Block Merge** and **SDXL Advanced Block Merge** provide:
- Individual control over input blocks (0-11)
- Individual control over output blocks (0-11)
- Middle block control (3 components)
- Time embedding and label embedding control
- Final output layer control
- Memory-optimized processing for any GPU size
- Advanced memory management with context managers
- Automatic GPU/CPU fallback based on available memory
### V-Pred/Epsilon Block Merging
The **Epsilon/V-Pred Block Merge** node provides:
- Automatic conversion between V-Prediction and Epsilon prediction types
- Granular control over individual blocks:
- Input blocks 0-8 (9 individual controls)
- Middle blocks 0-2 (3 individual controls)
- Output blocks 0-8 (9 individual controls)
- Final output layer control
- Support for mixed prediction type merging (e.g., Epsilon + V-Pred → Epsilon)
- Mathematical conversion factors for accurate type conversion
### Pure Prediction Type Conversion
The **Epsilon/V-Pred Converter** provides:
- Pure conversion between prediction types without merging
- Configurable conversion strength (0.0 to 2.0)
- Smart layer targeting:
- **Critical layers**: Output blocks, middle blocks, final output
- **Secondary layers**: Time embedding, late input blocks
- Empirically derived conversion factors for mathematical accuracy
- Automatic model options updating for ComfyUI compatibility
### Advanced CLIP Merging
The **Advanced CLIP Merge** node offers:
- **Multiple Interpolation Methods**:
- **Linear**: Standard weighted average
- **SLERP**: Spherical interpolation for vector-like parameters
- **Cosine**: Smooth cosine-based transitions
- **Weighted Average**: Magnitude-based automatic weighting
- **Spectral Blend**: Frequency domain blending with separate magnitude/phase control
- **Layer-Specific Controls**:
- **Attention Bias**: Adjust merging for attention layers
- **Feedforward Bias**: Control feedforward network blending
- **Embedding Bias**: Modify text/positional embedding merge ratios
- **Normalization Bias**: Adjust layer normalization merging
- **Advanced Options**:
- **Preserve Norms**: Maintain original parameter magnitudes
- **Memory Efficient**: Optimize for GPU memory usage
- **Spectral Alpha**: Control phase vs magnitude blending in spectral mode
### Advanced Masking System
Create sophisticated merging patterns:
1. **Generate Masks**:
- Layer-based: Target specific model layers
- Block-based: Target transformer blocks
- Component-based: Target attention or feedforward layers
- Custom patterns: Use regex patterns
- Random sparse: Create random merging patterns
- Depth gradients: Gradual transitions through model depth
2. **Filter and Blend Masks**:
- Filter model keys by component type
- Combine masks with various blending modes
- Memory-efficient batch processing
3. **Apply Masks**:
- Use masks to control merge strength per parameter
- Selective merging based on model structure
- Tensor-level precision control
### Model Weight Modification
Transform model weights directly:
- **Multiply**: Scale weights by a factor
- **Add**: Add constant values
- **Set Value**: Replace weights with specific values
- **Clamp Magnitude**: Limit weight magnitudes
- **Scale Max Abs**: Normalize based on reference model
### Workflow Organization
Use the **Checkpoint Reroute + Notes** node to:
- Clean up complex workflows with multiple model connections
- Add documentation and notes directly in your workflow
- Maintain MODEL, CLIP, and VAE connections without modification
- Keep track of model versions and merge parameters
- Zero processing overhead - direct passthrough routing
- Organize workflow structure for better readability
## Node Reference
| Node | Category | Purpose |
|------|----------|---------|
| Main Merge | Tensor Prism/Core | Advanced merging with multiple methods |
| Prism | Tensor Prism/Core | Spectral frequency-domain merging |
| SDXL Block Merge | Tensor_Prism/Merge | Basic granular SDXL merging |
| SDXL Advanced Block Merge | Tensor_Prism/Merge | GPU-optimized SDXL merging |
| Epsilon/V-Pred Converter | Tensor_Prism/Convert | Pure prediction type conversion |
| Advanced CLIP Merge | Tensor_Prism/CLIP | Sophisticated CLIP merging with multiple methods |
| Model Mask Generator | Tensor Prism/Mask | Create structural masks |
| Weighted Mask Merge | Tensor Prism/Mask | Apply masks to merging |
| Model Key Filter | Tensor_Prism/Mask | Filter model parameters |
| Mask Blender | Tensor_Prism/Mask | Combine multiple masks |
| Model Weight Modifier | Tensor_Prism/Transform | Direct weight manipulation |
| Checkpoint Reroute + Notes | Tensor_Prism/Utilities | MODEL/CLIP/VAE rerouting with workflow documentation |
## Memory Management
The TensorPrism pack includes advanced memory management features:
- **Automatic GPU Detection**: Optimizes for your specific GPU memory
- **Adaptive Batch Sizes**: Adjusts processing based on available memory
- **Precision Selection**: Automatic FP16/FP32 based on memory constraints
- **Progressive Cleanup**: Aggressive garbage collection for low-memory systems
- **CPU Fallback**: Automatic fallback when GPU memory is insufficient
- **Memory Context Managers**: Automatic cleanup and resource management
- **Threshold-Based Processing**: Memory usage monitoring with configurable limits
### Recommended Settings by GPU:
- **24GB+ (RTX 4090, etc.)**: Use default settings, batch size 50+
- **12GB (RTX 4070 Ti, etc.)**: Set memory limit to 8GB, enable auto precision, batch size 30-50
- **8GB (RTX 4060 Ti, etc.)**: Set memory limit to 6GB, force CPU for large merges, batch size 10-30
- **6GB and below**: Use CPU processing for best stability, enable aggressive cleanup
## Requirements
- ComfyUI
- PyTorch >= 1.12.0
- NumPy >= 1.21.0
- psutil >= 5.8.0 (for memory management)
## Tips and Best Practices
1. **Start Conservative**: Begin with lower merge ratios (0.3-0.7) and adjust based on results
2. **Use SLERP for Dissimilar Models**: When merging very different models, SLERP often produces better results
3. **Leverage Spectral Methods**: Frequency domain merging can preserve details better than linear methods
4. **Use Masks for Precision**: Create masks to merge only specific model components
5. **Memory Management**: Monitor memory usage and adjust batch sizes for your hardware
6. **Experiment with Spectral Parameters**: Different frequency biases can dramatically change results
7. **Layer-Selective Merging**: Use depth gradients for smooth transitions through model layers
8. **Prediction Type Awareness**: Use the Epsilon/V-Pred nodes when working with models of different prediction types
9. **CLIP Merging Strategy**: Use spectral blend for CLIP when preserving text understanding is critical
10. **Conversion Strength**: Start with 1.0 conversion strength and adjust if results seem over/under-converted
11. **Document Your Workflows**: Use the Checkpoint Reroute + Notes node to keep track of your merging experiments
12. **Organize Complex Workflows**: Use reroute nodes with notes to create clean, documented workflow structures
13. **Layer-Specific CLIP Control**: Use attention/feedforward bias to fine-tune CLIP behavior for specific use cases
14. **Zero Overhead Documentation**: The reroute node adds no processing time while providing workflow organization
## Troubleshooting
- **Memory Issues**: Reduce batch sizes, lower memory limits, or enable CPU fallback
- **Poor Results**: Try different merging methods or adjust spectral parameters
- **Compatibility**: Ensure models are the same architecture (SDXL with SDXL, etc.)
- **Slow Performance**: Check if you're accidentally using CPU when GPU is available
- **Artifacts**: Try more conservative merge ratios or use SLERP for smoother blending
- **Prediction Type Issues**: Use the Epsilon/V-Pred nodes for automatic type conversion
- **CLIP Problems**: Use preserve_norms=True and lower merge ratios for CLIP stability
- **Conversion Artifacts**: Reduce conversion strength or use pure converter instead of block merge
- **Workflow Complexity**: Use Checkpoint Reroute + Notes nodes to organize and document complex merging chains
## Performance Tips
- **Batch Size**: Larger batches are more efficient but use more memory
- **Precision Mode**: FP16 saves memory but may affect quality on some operations
- **Memory Cleanup**: Enable aggressive cleanup for systems with limited RAM
- **Device Selection**: Let the system auto-detect optimal device unless you have specific needs
- **Workflow Organization**: Use reroute nodes to reduce visual complexity without performance impact
- **CLIP Memory**: CLIP merging is less memory-intensive than UNet merging
- **Conversion vs Merging**: Pure conversion uses less memory than block-level merge+convert
## License
https://www.gnu.org/licenses/gpl-3.0.en.html
## Contributing
Contributions are welcome! Please feel free to submit issues and pull requests.
## Changelog
### Version 1.2.0
- Added **Advanced CLIP Merge** node with multiple interpolation methods and layer-specific controls
- Added **Epsilon/V-Pred Converter** node for pure prediction type conversion
- Enhanced **SDXL Advanced Block Merge** with improved memory management and context managers
- Improved **Weighted Mask Merge** with tensor-level precision control
- Added spectral blending capabilities to CLIP merging
- Enhanced memory management with automatic GPU/CPU fallback
- Improved layer identification and targeting for prediction type conversion
- Better mathematical accuracy in conversion factors
- Added **Epsilon/V-Pred Block Merge** node with granular block control and prediction type conversion
- Added **Checkpoint Reroute + Notes (Tensor Prism)** node for workflow organization and documentation
- Enhanced block-level merging capabilities with individual layer control (input blocks 0-8, middle blocks 0-2, output blocks 0-8)
- Improved prediction type conversion with mathematical accuracy
- Better workflow documentation and organization features
- Zero-overhead utility nodes for complex workflow management
- Improved tooltip system and user experience enhancements
### Version 1.1.0
- GPU-optimized memory management
- Cross-platform compatibility (CUDA/MPS/CPU)
- Bunch of Bug Fixes
- Addition of 3 Nodes
### Version 1.0.0
- Initial release
- Core merging nodes with advanced interpolation methods
- Advanced mask system with filtering and blending
- Spectral analysis and frequency-domain merging
- Advanced mask system with filtering and blending
- Support for SDXL models with granular block control
- Model weight modification tools
+304
View File
@@ -0,0 +1,304 @@
import torch
import torch.nn.functional as F
import numpy as np
import math
import gc
from typing import Dict, Any, Tuple, Optional, List
import psutil
class AdvancedCLIPMerge:
"""
Advanced CLIP merging node with multiple interpolation methods and layer-specific control.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"clip1": ("CLIP",),
"clip2": ("CLIP",),
"merge_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"merge_method": ([
"linear",
"slerp",
"cosine",
"weighted_average",
"spectral_blend"
], {"default": "slerp"}),
},
"optional": {
"attention_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
"feedforward_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
"embedding_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
"normalization_bias": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}),
"preserve_norms": ("BOOLEAN", {"default": True}),
"spectral_alpha": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}),
"memory_efficient": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("CLIP", "STRING")
RETURN_NAMES = ("clip", "merge_info")
FUNCTION = "merge_clips"
CATEGORY = "Tensor_Prism/CLIP"
def __init__(self):
self.device = self._get_optimal_device()
self.dtype = torch.float16 if torch.cuda.is_available() else torch.float32
def _get_optimal_device(self):
if torch.cuda.is_available():
return torch.device("cuda")
elif hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
return torch.device("mps")
else:
return torch.device("cpu")
def _get_layer_type(self, key: str) -> str:
"""Determine layer type from parameter key"""
key_lower = key.lower()
if any(term in key_lower for term in ['attn', 'attention', 'self_attn', 'cross_attn']):
return 'attention'
elif any(term in key_lower for term in ['ffn', 'mlp', 'feed_forward']):
return 'feedforward'
elif any(term in key_lower for term in ['embed', 'position', 'token']):
return 'embedding'
elif any(term in key_lower for term in ['norm', 'layer_norm', 'group_norm']):
return 'normalization'
else:
return 'other'
def _apply_layer_bias(self, ratio: float, layer_type: str, attention_bias: float,
feedforward_bias: float, embedding_bias: float, normalization_bias: float) -> float:
"""Apply layer-specific bias to merge ratio"""
bias_map = {
'attention': attention_bias,
'feedforward': feedforward_bias,
'embedding': embedding_bias,
'normalization': normalization_bias,
'other': 0.0
}
biased_ratio = ratio + bias_map.get(layer_type, 0.0)
return max(0.0, min(1.0, biased_ratio))
def _slerp(self, t1: torch.Tensor, t2: torch.Tensor, ratio: float) -> torch.Tensor:
"""Spherical Linear Interpolation"""
t1_flat = t1.flatten()
t2_flat = t2.flatten()
# Normalize vectors
t1_norm = F.normalize(t1_flat, dim=0)
t2_norm = F.normalize(t2_flat, dim=0)
# Calculate angle between vectors
dot = torch.clamp(torch.dot(t1_norm, t2_norm), -1.0, 1.0)
theta = torch.acos(torch.abs(dot))
# Handle edge cases
if theta < 1e-6:
return self._linear_interpolation(t1, t2, ratio)
# SLERP formula
sin_theta = torch.sin(theta)
w1 = torch.sin((1.0 - ratio) * theta) / sin_theta
w2 = torch.sin(ratio * theta) / sin_theta
result = w1 * t1_flat + w2 * t2_flat
return result.reshape(t1.shape)
def _cosine_interpolation(self, t1: torch.Tensor, t2: torch.Tensor, ratio: float) -> torch.Tensor:
"""Cosine interpolation for smoother transitions"""
smooth_ratio = 0.5 * (1.0 - math.cos(ratio * math.pi))
return t1 * (1.0 - smooth_ratio) + t2 * smooth_ratio
def _linear_interpolation(self, t1: torch.Tensor, t2: torch.Tensor, ratio: float) -> torch.Tensor:
"""Standard linear interpolation"""
return t1 * (1.0 - ratio) + t2 * ratio
def _weighted_average(self, t1: torch.Tensor, t2: torch.Tensor, ratio: float) -> torch.Tensor:
"""Weighted average based on tensor magnitudes"""
mag1 = torch.norm(t1)
mag2 = torch.norm(t2)
total_mag = mag1 + mag2 + 1e-8
w1 = (mag1 / total_mag) * (1.0 - ratio)
w2 = (mag2 / total_mag) * ratio
norm_factor = w1 + w2
return (w1 * t1 + w2 * t2) / norm_factor
def _spectral_blend(self, t1: torch.Tensor, t2: torch.Tensor, ratio: float, alpha: float) -> torch.Tensor:
"""Frequency domain blending"""
if len(t1.shape) < 2:
return self._linear_interpolation(t1, t2, ratio)
try:
# Reshape to 2D for FFT
original_shape = t1.shape
t1_2d = t1.view(t1.shape[0], -1).float()
t2_2d = t2.view(t2.shape[0], -1).float()
# Apply FFT
fft1 = torch.fft.fft2(t1_2d)
fft2 = torch.fft.fft2(t2_2d)
# Blend in frequency domain
magnitude1 = torch.abs(fft1)
magnitude2 = torch.abs(fft2)
phase1 = torch.angle(fft1)
phase2 = torch.angle(fft2)
# Blend magnitudes and phases separately
blended_mag = (1.0 - ratio) * magnitude1 + ratio * magnitude2
blended_phase = (1.0 - alpha) * phase1 + alpha * phase2
# Reconstruct complex spectrum
blended_fft = blended_mag * torch.exp(1j * blended_phase)
# Inverse FFT
result = torch.fft.ifft2(blended_fft).real
return result.view(original_shape).to(t1.dtype)
except:
return self._linear_interpolation(t1, t2, ratio)
def _merge_tensors(self, t1: torch.Tensor, t2: torch.Tensor, ratio: float, method: str,
spectral_alpha: float = 0.5) -> torch.Tensor:
"""Apply the specified merge method to two tensors"""
if method == "linear":
return self._linear_interpolation(t1, t2, ratio)
elif method == "slerp":
return self._slerp(t1, t2, ratio)
elif method == "cosine":
return self._cosine_interpolation(t1, t2, ratio)
elif method == "weighted_average":
return self._weighted_average(t1, t2, ratio)
elif method == "spectral_blend":
return self._spectral_blend(t1, t2, ratio, spectral_alpha)
else:
return self._linear_interpolation(t1, t2, ratio)
def merge_clips(self, clip1, clip2, merge_ratio: float, merge_method: str,
attention_bias: float = 0.0, feedforward_bias: float = 0.0,
embedding_bias: float = 0.0, normalization_bias: float = 0.0,
preserve_norms: bool = True, spectral_alpha: float = 0.5,
memory_efficient: bool = True):
try:
# Memory cleanup
if memory_efficient and torch.cuda.is_available():
torch.cuda.empty_cache()
gc.collect()
# Clone the first CLIP model
merged_clip = clip1.clone()
# Get state dictionaries
state_dict1 = clip1.cond_stage_model.state_dict() if hasattr(clip1, 'cond_stage_model') else clip1.state_dict()
state_dict2 = clip2.cond_stage_model.state_dict() if hasattr(clip2, 'cond_stage_model') else clip2.state_dict()
# Merge statistics
merged_layers = 0
layer_types_merged = {'attention': 0, 'feedforward': 0, 'embedding': 0, 'normalization': 0, 'other': 0}
# Merge parameters
merged_state_dict = {}
for key in state_dict1.keys():
if key in state_dict2:
t1 = state_dict1[key].to(self.device)
t2 = state_dict2[key].to(self.device)
# Determine merge ratio for this layer
layer_type = self._get_layer_type(key)
current_ratio = merge_ratio
# Apply layer-specific bias
current_ratio = self._apply_layer_bias(
current_ratio, layer_type, attention_bias,
feedforward_bias, embedding_bias, normalization_bias
)
# Store original norms if preserving
original_norm = torch.norm(t1) if preserve_norms else None
try:
# Perform merge
merged_tensor = self._merge_tensors(
t1, t2, current_ratio, merge_method, spectral_alpha
)
# Preserve norm if requested
if preserve_norms and original_norm is not None:
current_norm = torch.norm(merged_tensor)
if current_norm > 1e-8:
merged_tensor = merged_tensor * (original_norm / current_norm)
merged_state_dict[key] = merged_tensor.cpu()
merged_layers += 1
layer_types_merged[layer_type] += 1
except Exception as e:
# Fallback to linear interpolation
merged_tensor = self._linear_interpolation(t1, t2, current_ratio)
merged_state_dict[key] = merged_tensor.cpu()
# Memory cleanup for large tensors
if memory_efficient:
del t1, t2
if torch.cuda.is_available():
torch.cuda.empty_cache()
else:
# Keep original parameter if not in second model
merged_state_dict[key] = state_dict1[key]
# Load merged state dict
if hasattr(merged_clip, 'cond_stage_model'):
merged_clip.cond_stage_model.load_state_dict(merged_state_dict, strict=False)
else:
merged_clip.load_state_dict(merged_state_dict, strict=False)
# Generate merge information
merge_info = f"""=== ADVANCED CLIP MERGE RESULTS ===
Method: {merge_method}
Base Ratio: {merge_ratio:.3f}
=== LAYER STATISTICS ===
Successfully Merged: {merged_layers}
Total Parameters: {len(state_dict1)}
=== LAYER TYPE BREAKDOWN ===
Attention: {layer_types_merged['attention']}
Feedforward: {layer_types_merged['feedforward']}
Embedding: {layer_types_merged['embedding']}
Normalization: {layer_types_merged['normalization']}
Other: {layer_types_merged['other']}
=== MERGE SETTINGS ===
Attention Bias: {attention_bias:+.3f}
Feedforward Bias: {feedforward_bias:+.3f}
Embedding Bias: {embedding_bias:+.3f}
Normalization Bias: {normalization_bias:+.3f}
Preserve Norms: {preserve_norms}
"""
# Final cleanup
if memory_efficient and torch.cuda.is_available():
torch.cuda.empty_cache()
gc.collect()
return (merged_clip, merge_info)
except Exception as e:
error_info = f"CLIP merge failed: {str(e)}"
return (clip1, error_info)
# Node mappings for ComfyUI
NODE_CLASS_MAPPINGS = {
"AdvancedCLIPMerge": AdvancedCLIPMerge
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AdvancedCLIPMerge": "Advanced CLIP Merge (Tensor Prism)"
}
+96
View File
@@ -0,0 +1,96 @@
"""
TensorPrism Checkpoint Reroute Notes Node
========================================
A utility node for rerouting checkpoint components (MODEL, CLIP, VAE) with optional notes
for workflow documentation and organization.
Author: AstrionX
Version: 1.2.0
License: GPL-3.0
"""
class TensorPrism_CheckpointReroute_Notes:
"""
A utility node that reroutes MODEL, CLIP, and VAE inputs to outputs with optional notes.
Useful for workflow organization and documentation.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL", {"tooltip": "Model input to reroute"}),
"clip": ("CLIP", {"tooltip": "CLIP input to reroute"}),
"vae": ("VAE", {"tooltip": "VAE input to reroute"})
},
"optional": {
"notes": ("STRING", {
"default": "",
"multiline": True,
"tooltip": "Optional notes for workflow documentation"
})
}
}
RETURN_TYPES = ("MODEL", "CLIP", "VAE")
RETURN_NAMES = ("model", "clip", "vae")
FUNCTION = "reroute"
CATEGORY = "Tensor_Prism/Utilities"
DESCRIPTION = """
Reroutes checkpoint components (MODEL, CLIP, VAE) with optional documentation.
Features:
• Pass-through routing for MODEL, CLIP, and VAE
• Optional notes field for workflow documentation
• Maintains full compatibility with all checkpoint types
• Zero processing overhead - direct passthrough
Use Cases:
• Organizing complex workflows
• Adding documentation to checkpoint flows
• Creating clean connection points
• Workflow readability improvements
"""
def reroute(self, model, clip, vae, notes="", **kwargs):
"""
Reroute the inputs directly to outputs with optional notes processing.
Args:
model: The model to reroute
clip: The CLIP to reroute
vae: The VAE to reroute
notes: Optional notes (stored but not processed)
**kwargs: Additional arguments (ignored)
Returns:
tuple: (model, clip, vae) - Direct passthrough of inputs
"""
# Optional: Log notes if provided (for debugging/workflow tracking)
if notes and notes.strip():
print(f"[TensorPrism Checkpoint Reroute] Notes: {notes.strip()}")
# Direct passthrough - zero processing overhead
return model, clip, vae
@classmethod
def IS_CHANGED(cls, model, clip, vae, notes="", **kwargs):
"""
Since this is a passthrough node, we don't need change detection.
Return None to indicate no caching needed.
"""
return None
# Node registration for ComfyUI
NODE_CLASS_MAPPINGS = {
"TensorPrism_CheckpointReroute_Notes": TensorPrism_CheckpointReroute_Notes
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_CheckpointReroute_Notes": "Checkpoint Reroute + Notes (Tensor Prism)"
}
# Export the class for import
__all__ = ["TensorPrism_CheckpointReroute_Notes"]
+386
View File
@@ -0,0 +1,386 @@
import torch
import copy
from enum import Enum
# Define enums for cleaner and safer dropdown/option handling in ComfyUI
class CheckpointPrecision(str, Enum):
"""Defines the floating-point precision for tensor operations."""
FP16 = 'fp16'
FP32 = 'fp32'
class EnhancementMethod(str, Enum):
"""Defines the method used for tensor enhancement."""
LINEAR = 'linear'
ATTENTION = 'attention'
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
ALL = 'all'
# Helper functions to identify keys belonging to specific model components
def is_unet_key(key: str) -> bool:
"""Checks if a given key string likely belongs to a UNet model."""
key_lower = key.lower()
return 'unet' in key_lower or 'model.diffusion_model' in key_lower
def is_vae_key(key: str) -> bool:
"""Checks if a given key string likely belongs to a VAE model."""
key_lower = key.lower()
return 'vae' in key_lower or 'autoencoder' in key_lower
def is_text_encoder_key(key: str) -> bool:
"""Checks if a given key string likely belongs to a Text Encoder model."""
key_lower = key.lower()
return 'clip' in key_lower or 'text_encoder' in key_lower or 'cond_stage' in key_lower
def clamp_tensor_stats(tensor: torch.Tensor, max_abs_threshold: float = 1.0) -> torch.Tensor:
"""
Prevents runaway values in a tensor by softly clamping via tanh scaling if its
maximum absolute value exceeds a threshold. Applied only to floating-point tensors.
"""
if not tensor.is_floating_point():
return tensor
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)
return tensor
def apply_linear_smoothing(tensor: torch.Tensor, strength: float) -> torch.Tensor:
"""
Applies a simple exponential moving average toward the tensor's mean.
Effective when `strength` is greater than 0.
"""
if strength <= 0.0:
return tensor
device = tensor.device
mean_value = tensor.mean().to(device)
return (tensor * (1.0 - strength) + mean_value * strength).to(device)
def apply_linear_sharpen(tensor: torch.Tensor, strength: float) -> torch.Tensor:
"""
Applies an unsharp-like effect by boosting deviations from the tensor's mean.
Effective when `strength` is greater than 0.
"""
if strength <= 0.0:
return tensor
device = tensor.device
mean_value = tensor.mean().to(device)
return (tensor + (tensor - mean_value) * strength).to(device)
def attention_refinement_stub(tensor: torch.Tensor, strength: float) -> torch.Tensor:
"""
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))
def apply_quality_boost(tensor: torch.Tensor, strength: float) -> torch.Tensor:
"""
Applies a multi-stage subtle enhancement: local contrast-like scaling + gentle
non-linear sharpening. Effective when `strength` is greater than 0.
"""
if strength <= 0.0:
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)
def apply_adaptive_overbake_limiter(
original_tensor: torch.Tensor,
modified_tensor: torch.Tensor,
max_relative_increase: float = 0.12
) -> torch.Tensor:
"""
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.
"""
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
relative_change = (modified_mean_abs - original_mean_abs) / original_mean_abs
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)
class ModelEnhancerTensorPrism:
"""
ComfyUI custom node to enhance Stable Diffusion XL checkpoints.
Applies various processing steps to selected tensors within the model's state_dict
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)
"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}),
"blend_strength": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01}),
"precision": (
[p.value for p in CheckpointPrecision],
{"default": CheckpointPrecision.FP16.value}
),
"method": (
[m.value for m in EnhancementMethod],
{"default": EnhancementMethod.LINEAR.value}
),
"modules_to_enhance": (
[mod.value for mod in TargetModule],
{"default": TargetModule.UNET.value}
),
"adaptive_overbake_prevention": ("BOOLEAN", {"default": True}),
"attention_iterations": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}),
}
}
# 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
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.
"""
# 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:
return True
if is_unet_key(key) and TargetModule.UNET.value in target_modules_set:
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
def _cast_tensor_to_precision(self, tensor: torch.Tensor, precision: CheckpointPrecision) -> torch.Tensor:
"""Casts a floating-point tensor to the specified precision."""
if not isinstance(tensor, torch.Tensor) or not tensor.is_floating_point():
return tensor
device = tensor.device
if precision == CheckpointPrecision.FP16:
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
def _enhance_single_tensor(
self,
original_tensor: torch.Tensor,
smoothing_strength: float,
sharpening_strength: float,
quality_boost_strength: float,
enhancement_method: EnhancementMethod,
adaptive_overbake_enabled: bool,
attention_iterations: int,
target_precision: CheckpointPrecision
) -> 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
)
return processed_tensor.to(device)
def enhance_checkpoint(
self,
checkpoint, # Changed from dict to accept MODEL objects
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]
adaptive_overbake_prevention: bool,
attention_iterations: int
) -> tuple: # Changed return type
"""
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.")
# Convert string inputs from ComfyUI to Enum types for type safety and clarity
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)
# 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
# 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
)
# Ensure enhanced tensor is on correct device
enhanced_tensor = enhanced_tensor.to(device)
# Blend the original tensor with the enhanced tensor based on blend_strength
blended_tensor = original_tensor * (1.0 - blend_strength) + enhanced_tensor * blend_strength
# 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
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
# 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 = {
"ModelEnhancerTensorPrism": ModelEnhancerTensorPrism
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ModelEnhancerTensorPrism": "Model Enhancer (Tensor Prism)"
}
+257
View File
@@ -0,0 +1,257 @@
import torch
import math
import re
from typing import Dict, Any, Optional
import comfy.model_management
class TensorPrism_LayeredBlend:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_A": ("MODEL",),
"model_B": ("MODEL",),
"text_encoder_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"text_encoder_method": (["linear", "slerp", "cosine"],),
"unet_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"unet_method": (["linear", "slerp", "cosine"],),
"time_embed_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"input_blocks_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"middle_block_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"output_blocks_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"out_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
},
"optional": {
"vae_A": ("VAE",),
"vae_B": ("VAE",),
"vae_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"vae_method": (["linear", "slerp", "cosine"],),
}
}
RETURN_TYPES = ("MODEL", "VAE",)
RETURN_NAMES = ("merged_model", "merged_vae",)
FUNCTION = "blend_models"
CATEGORY = "Tensor Prism/Merge"
@staticmethod
def _blend_linear(t1: torch.Tensor, t2: torch.Tensor, alpha: float) -> torch.Tensor:
"""Linear interpolation between two tensors"""
return (1 - alpha) * t1 + alpha * t2
@staticmethod
def _blend_slerp(t1: torch.Tensor, t2: torch.Tensor, alpha: float) -> torch.Tensor:
"""Spherical linear interpolation with proper error handling"""
if t1.numel() == 0 or t2.numel() == 0:
return t1
t1_flat = t1.view(-1)
t2_flat = t2.view(-1)
# Normalize vectors
t1_norm = torch.nn.functional.normalize(t1_flat, dim=0, eps=1e-8)
t2_norm = torch.nn.functional.normalize(t2_flat, dim=0, eps=1e-8)
# Calculate dot product and handle edge cases
dot = torch.clamp(torch.dot(t1_norm, t2_norm), -1.0 + 1e-7, 1.0 - 1e-7)
# Handle near-parallel vectors
if abs(dot.item()) > 0.9995:
return TensorPrism_LayeredBlend._blend_linear(t1, t2, alpha)
# SLERP calculation
theta = torch.acos(dot)
sin_theta = torch.sin(theta)
if sin_theta.abs() < 1e-6:
return TensorPrism_LayeredBlend._blend_linear(t1, t2, alpha)
w1 = torch.sin((1 - alpha) * theta) / sin_theta
w2 = torch.sin(alpha * theta) / sin_theta
result = w1 * t1_norm + w2 * t2_norm
# Restore original magnitude
original_norm = torch.norm(t1_flat) * (1 - alpha) + torch.norm(t2_flat) * alpha
result = result * original_norm
return result.view(t1.shape)
@staticmethod
def _blend_cosine(t1: torch.Tensor, t2: torch.Tensor, alpha: float) -> torch.Tensor:
"""Cosine interpolation for smoother transitions"""
mu = (1 - math.cos(alpha * math.pi)) / 2
return (1 - mu) * t1 + mu * t2
def _get_component_type(self, param_name: str) -> str:
"""Identify which component a parameter belongs to"""
param_lower = param_name.lower()
# Text encoder patterns
if any(pattern in param_lower for pattern in ['cond_stage_model', 'text_encoder', 'clip', 'transformer.text_model']):
return 'text_encoder'
# UNet component patterns
if any(pattern in param_lower for pattern in ['model.diffusion_model', 'unet']):
# More specific UNet components
if 'time_embed' in param_lower:
return 'time_embed'
elif 'input_blocks' in param_lower:
return 'input_blocks'
elif 'middle_block' in param_lower:
return 'middle_block'
elif 'output_blocks' in param_lower:
return 'output_blocks'
elif param_lower.endswith('.out.weight') or param_lower.endswith('.out.bias'):
return 'out'
else:
return 'unet'
# VAE patterns
if any(pattern in param_lower for pattern in ['first_stage_model', 'vae', 'decoder', 'encoder']):
return 'vae'
# Default to unet for unrecognized parameters
return 'unet'
def _apply_blend_method(self, t1: torch.Tensor, t2: torch.Tensor, strength: float, method: str) -> torch.Tensor:
"""Apply the specified blending method"""
if method == "linear":
return self._blend_linear(t1, t2, strength)
elif method == "slerp":
return self._blend_slerp(t1, t2, strength)
elif method == "cosine":
return self._blend_cosine(t1, t2, strength)
else:
return self._blend_linear(t1, t2, strength)
def blend_models(self, model_A, model_B,
text_encoder_strength, text_encoder_method,
unet_strength, unet_method,
time_embed_strength, input_blocks_strength,
middle_block_strength, output_blocks_strength, out_strength,
vae_A: Optional[Any] = None, vae_B: Optional[Any] = None,
vae_strength: float = 0.5, vae_method: str = "linear"):
"""
Blend models with component-specific controls
"""
# Clone model A as the base
merged_model = model_A.clone()
# Get state dictionaries
state_dict_A = model_A.model.state_dict()
state_dict_B = model_B.model.state_dict()
# Create patches dictionary
patches = {}
# Component strength mapping
component_strengths = {
'text_encoder': (text_encoder_strength, text_encoder_method),
'time_embed': (time_embed_strength, unet_method),
'input_blocks': (input_blocks_strength, unet_method),
'middle_block': (middle_block_strength, unet_method),
'output_blocks': (output_blocks_strength, unet_method),
'out': (out_strength, unet_method),
'unet': (unet_strength, unet_method), # Default UNet components
'vae': (0.0, "linear") # VAE handled separately
}
# Process each parameter
for key in state_dict_A.keys():
if key in state_dict_B:
t1 = state_dict_A[key]
t2 = state_dict_B[key]
if isinstance(t1, torch.Tensor) and isinstance(t2, torch.Tensor):
if t1.shape == t2.shape:
try:
# Determine component type
component_type = self._get_component_type(key)
# Get strength and method for this component
strength, method = component_strengths.get(component_type, (unet_strength, unet_method))
# Skip if strength is 0
if strength == 0:
continue
# Apply blending
merged_tensor = self._apply_blend_method(t1, t2, strength, method)
# Only add patch if there's a meaningful difference
diff = merged_tensor - t1
if torch.abs(diff).max() > 1e-8:
patches[key] = (diff,)
except Exception as e:
print(f"Warning: Failed to blend parameter {key}: {e}")
continue
# Apply patches to the merged model
if patches:
merged_model.add_patches(patches, 1.0)
# Handle VAE blending if provided
merged_vae = None
if vae_A is not None and vae_B is not None:
try:
merged_vae = self._blend_vaes(vae_A, vae_B, vae_strength, vae_method)
except Exception as e:
print(f"Warning: Failed to blend VAEs: {e}")
merged_vae = vae_A # Fallback to VAE A
return (merged_model, merged_vae)
def _blend_vaes(self, vae_A, vae_B, strength: float, method: str):
"""Blend two VAE models"""
if strength == 0:
return vae_A
if strength == 1:
return vae_B
# Clone VAE A as base
merged_vae = vae_A.clone() if hasattr(vae_A, 'clone') else vae_A
try:
# Get state dicts if available
if hasattr(vae_A, 'first_stage_model') and hasattr(vae_B, 'first_stage_model'):
state_dict_A = vae_A.first_stage_model.state_dict()
state_dict_B = vae_B.first_stage_model.state_dict()
# Create patches for VAE
vae_patches = {}
for key in state_dict_A.keys():
if key in state_dict_B:
t1 = state_dict_A[key]
t2 = state_dict_B[key]
if isinstance(t1, torch.Tensor) and isinstance(t2, torch.Tensor):
if t1.shape == t2.shape:
merged_tensor = self._apply_blend_method(t1, t2, strength, method)
diff = merged_tensor - t1
if torch.abs(diff).max() > 1e-8:
vae_patches[key] = (diff,)
# Apply VAE patches if the VAE supports it
if vae_patches and hasattr(merged_vae, 'add_patches'):
merged_vae.add_patches(vae_patches, 1.0)
except Exception as e:
print(f"VAE blending failed, using VAE A: {e}")
return vae_A
return merged_vae
# Register node
NODE_CLASS_MAPPINGS = {
"TensorPrism_LayeredBlend": TensorPrism_LayeredBlend
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_LayeredBlend": "Layered Blend (Tensor Prism)"
}
+228
View File
@@ -0,0 +1,228 @@
import torch
import math
import torch.fft
import comfy.model_management
class TensorPrism_MainMerge:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_A": ("MODEL",),
"model_B": ("MODEL",),
"merge_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"method": (["linear", "slerp", "cosine", "directional", "frequency", "stochastic"],),
},
"optional": {
"random_seed": ("INT", {"default": 42, "min": 0}),
"stochastic_prob": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 1.0}),
"freq_low_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"freq_high_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("merged_model",)
FUNCTION = "merge_models"
CATEGORY = "Tensor Prism/Core"
# --- Standard methods ---
@staticmethod
def _blend_linear(t1, t2, alpha):
return (1 - alpha) * t1 + alpha * t2
@staticmethod
def _blend_slerp(t1, t2, alpha):
"""Spherical linear interpolation"""
if t1.numel() == 0 or t2.numel() == 0:
return t1
t1_flat, t2_flat = t1.view(-1), t2.view(-1)
# Normalize vectors
t1_norm = torch.nn.functional.normalize(t1_flat, dim=0, eps=1e-8)
t2_norm = torch.nn.functional.normalize(t2_flat, dim=0, eps=1e-8)
# Calculate dot product and clamp to valid range
dot = torch.clamp(torch.dot(t1_norm, t2_norm), -1.0 + 1e-7, 1.0 - 1e-7)
# Handle near-parallel vectors
if abs(dot.item()) > 0.9995:
return cls._blend_linear(t1, t2, alpha)
# Calculate angle and interpolation
theta = torch.acos(dot)
sin_theta = torch.sin(theta)
if sin_theta.abs() < 1e-6:
return cls._blend_linear(t1, t2, alpha)
# SLERP formula
w1 = torch.sin((1 - alpha) * theta) / sin_theta
w2 = torch.sin(alpha * theta) / sin_theta
result = w1 * t1_norm + w2 * t2_norm
# Restore original magnitude
original_norm = torch.norm(t1_flat) * (1 - alpha) + torch.norm(t2_flat) * alpha
result = result * original_norm
return result.view(t1.shape)
@staticmethod
def _blend_cosine(t1, t2, alpha):
"""Cosine interpolation for smoother transitions"""
mu = (1 - math.cos(alpha * math.pi)) / 2
return (1 - mu) * t1 + mu * t2
@staticmethod
def _blend_directional(t1, t2, alpha):
"""Direct vector interpolation - same as linear but more explicit"""
return t1 + (t2 - t1) * alpha
@staticmethod
def _blend_frequency(t1, t2, low_ratio, high_ratio):
"""FFT-based frequency domain blending"""
if t1.numel() < 4 or t2.numel() < 4:
# Fall back to linear for very small tensors
return TensorPrism_MainMerge._blend_linear(t1, t2, (low_ratio + high_ratio) / 2)
try:
# Reshape for FFT if needed (ensure at least 2D)
orig_shape = t1.shape
if t1.dim() < 2:
t1_fft = t1.view(-1, 1)
t2_fft = t2.view(-1, 1)
else:
t1_fft = t1
t2_fft = t2
# Compute FFT
fA = torch.fft.fft2(t1_fft.float())
fB = torch.fft.fft2(t2_fft.float())
# Create frequency mask
h, w = fA.shape[-2:]
cy, cx = h // 2, w // 2
# Create coordinate grids
y, x = torch.meshgrid(torch.arange(h), torch.arange(w), indexing='ij')
y, x = y.to(t1.device), x.to(t1.device)
# Calculate distance from center
dist = torch.sqrt((y - cy).float()**2 + (x - cx).float()**2)
max_dist = min(cy, cx)
# Create low/high frequency masks
if max_dist > 0:
normalized_dist = dist / max_dist
low_freq_mask = (normalized_dist < 0.3).float()
high_freq_mask = (normalized_dist >= 0.3).float()
else:
low_freq_mask = torch.ones_like(dist)
high_freq_mask = torch.zeros_like(dist)
# Blend in frequency domain
fA_low = fA * low_freq_mask.unsqueeze(0) if fA.dim() > 2 else fA * low_freq_mask
fB_low = fB * low_freq_mask.unsqueeze(0) if fB.dim() > 2 else fB * low_freq_mask
fA_high = fA * high_freq_mask.unsqueeze(0) if fA.dim() > 2 else fA * high_freq_mask
fB_high = fB * high_freq_mask.unsqueeze(0) if fB.dim() > 2 else fB * high_freq_mask
# Merge frequencies separately
merged_low = fA_low * (1 - low_ratio) + fB_low * low_ratio
merged_high = fA_high * (1 - high_ratio) + fB_high * high_ratio
merged_freq = merged_low + merged_high
# Convert back to spatial domain
result = torch.fft.ifft2(merged_freq).real
# Restore original shape and dtype
return result.view(orig_shape).to(t1.dtype)
except Exception as e:
# Fallback to linear interpolation if FFT fails
return TensorPrism_MainMerge._blend_linear(t1, t2, (low_ratio + high_ratio) / 2)
@staticmethod
def _blend_stochastic(t1, t2, alpha, prob, seed):
"""Stochastic blending with random dropout patterns"""
# Set seed for reproducibility
generator = torch.Generator(device=t1.device)
generator.manual_seed(seed)
# Create random mask
mask = torch.rand(t1.shape, generator=generator, device=t1.device) < prob
# Apply stochastic blending
# Where mask is True, use model B, otherwise blend normally
base_blend = TensorPrism_MainMerge._blend_linear(t1, t2, alpha)
stochastic_component = torch.where(mask, t2, t1)
# Combine base blend with stochastic component
return base_blend * (1 - prob) + stochastic_component * prob
def merge_models(self, model_A, model_B, merge_ratio, method,
random_seed=42, stochastic_prob=0.1,
freq_low_ratio=0.5, freq_high_ratio=0.5):
"""
Merge two ComfyUI models using various interpolation methods
"""
# Clone model A as the base
merged_model = model_A.clone()
# Get state dictionaries
state_dict_A = model_A.model.state_dict()
state_dict_B = model_B.model.state_dict()
# Create patches dictionary
patches = {}
for key in state_dict_A.keys():
if key in state_dict_B:
t1 = state_dict_A[key]
t2 = state_dict_B[key]
if isinstance(t1, torch.Tensor) and isinstance(t2, torch.Tensor):
# Only process if tensors have the same shape
if t1.shape == t2.shape:
try:
if method == "linear":
merged_tensor = self._blend_linear(t1, t2, merge_ratio)
elif method == "slerp":
merged_tensor = self._blend_slerp(t1, t2, merge_ratio)
elif method == "cosine":
merged_tensor = self._blend_cosine(t1, t2, merge_ratio)
elif method == "directional":
merged_tensor = self._blend_directional(t1, t2, merge_ratio)
elif method == "frequency":
merged_tensor = self._blend_frequency(t1, t2, freq_low_ratio, freq_high_ratio)
elif method == "stochastic":
merged_tensor = self._blend_stochastic(t1, t2, merge_ratio, stochastic_prob, random_seed)
else:
merged_tensor = self._blend_linear(t1, t2, merge_ratio)
# Only add patch if there's a meaningful difference
diff = merged_tensor - t1
if torch.abs(diff).max() > 1e-8:
patches[key] = (diff,)
except Exception as e:
print(f"Warning: Failed to merge parameter {key} with method {method}: {e}")
# Skip this parameter on error
continue
# Apply patches to the cloned model
if patches:
merged_model.add_patches(patches, 1.0)
return (merged_model,)
NODE_CLASS_MAPPINGS = {
"TensorPrism_MainMerge": TensorPrism_MainMerge
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_MainMerge": "Main Merge (Tensor Prism)"
}
+270
View File
@@ -0,0 +1,270 @@
import torch
import copy
import gc
import psutil
import re
import numpy as np
from collections import defaultdict
from typing import Dict, List, Tuple
def is_unet_key(key: str) -> bool:
key_lower = key.lower()
return 'unet' in key_lower or 'model.diffusion_model' in key_lower
def is_vae_key(key: str) -> bool:
key_lower = key.lower()
return 'vae' in key_lower or 'autoencoder' in key_lower or 'first_stage_model' in key_lower
def is_text_encoder_key(key: str) -> bool:
key_lower = key.lower()
return 'clip' in key_lower or 'text_encoder' in key_lower or 'cond_stage' in key_lower
def _get_unet_component_type(param_name: str) -> str:
param_lower = param_name.lower()
if 'time_embed' in param_lower:
return 'time_embed'
elif 'input_blocks' in param_lower:
return 'input_blocks'
elif 'middle_block' in param_lower:
return 'middle_block'
elif 'output_blocks' in param_lower:
return 'output_blocks'
elif param_lower.endswith('.out.weight') or param_lower.endswith('.out.bias'):
return 'out'
return 'other_unet'
def _analyze_model_structure(state_dict):
"""Analyze model structure to extract layer information"""
layer_info = {
"layers": {},
"total_layers": 0,
"layer_names": list(state_dict.keys())
}
layer_patterns = [
r"layers\.([0-9]+)",
r"blocks\.([0-9]+)",
r"h\.([0-9]+)",
r"layer\.([0-9]+)",
r"encoder\.layer\.([0-9]+)",
r"decoder\.layer\.([0-9]+)",
r"input_blocks\.([0-9]+)",
r"output_blocks\.([0-9]+)"
]
for name in state_dict.keys():
layer_num = None
for pattern in layer_patterns:
match = re.search(pattern, name)
if match:
layer_num = int(match.group(1))
break
if layer_num is not None:
if layer_num not in layer_info["layers"]:
layer_info["layers"][layer_num] = []
layer_info["layers"][layer_num].append(name)
layer_info["total_layers"] = max(layer_info["total_layers"], layer_num + 1)
return layer_info
MASK_TYPE = ("MASK",)
class TensorPrism_ModelKeyFilter:
"""
Memory-efficient model key filter that processes keys in batches
to handle large models without excessive RAM usage.
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"filter_mode": (["Include", "Exclude"], {"default": "Include"}),
"target_components": (["All", "UNet", "VAE", "Text Encoders", "Time Embeddings", "Input Blocks", "Middle Block", "Output Blocks", "Final UNet Output Layer", "Custom Pattern"], {"default": "UNet"}),
"default_value": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001, "label": "Default Mask Value (for non-targeted keys)"}),
"target_value": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001, "label": "Target Mask Value (for filtered keys)"}),
"memory_limit_gb": ("FLOAT", {"default": 2.0, "min": 0.5, "max": 16.0, "step": 0.1, "round": 0.1, "label": "Memory Limit (GB)"}),
},
"optional": {
"custom_pattern": ("STRING", {"default": "attn,resnets", "multiline": True, "placeholder": "Comma-separated regex patterns or substrings, e.g., 'attn,mlp.fc1'"}),
"exact_match_custom": ("BOOLEAN", {"default": False, "label_on": "Exact Match for Custom Patterns", "label_off": "Substring Match for Custom Patterns"}),
}
}
RETURN_TYPES = MASK_TYPE
RETURN_NAMES = ("filtered_mask",)
FUNCTION = "filter_keys_to_mask"
CATEGORY = "Tensor_Prism/Mask"
@staticmethod
def get_memory_info() -> Tuple[float, float]:
"""Get current memory usage and available memory in GB"""
memory = psutil.virtual_memory()
used_gb = (memory.total - memory.available) / (1024**3)
available_gb = memory.available / (1024**3)
return used_gb, available_gb
@staticmethod
def estimate_dict_memory_gb(dict_size: int) -> float:
"""Estimate memory usage of a dictionary with float values in GB"""
# Rough estimate: key string + float value + overhead
bytes_per_entry = 100 # Conservative estimate
return (dict_size * bytes_per_entry) / (1024**3)
def create_key_batches(self, all_keys: List[str], memory_limit_gb: float) -> List[List[str]]:
"""Create batches of keys that fit within memory limit"""
batches = []
current_batch = []
# Estimate how many keys we can process per batch
max_keys_per_batch = max(1000, int((memory_limit_gb * 1024**3) / 200)) # Conservative estimate
for key in all_keys:
current_batch.append(key)
if len(current_batch) >= max_keys_per_batch:
batches.append(current_batch)
current_batch = []
# Add final batch if not empty
if current_batch:
batches.append(current_batch)
return batches
def check_key_match(self, key_name: str, target_components: str,
patterns: List[str], exact_match_custom: bool) -> bool:
"""Check if a key matches the target component criteria"""
if target_components == "All":
return True
elif target_components == "UNet":
return is_unet_key(key_name) and not is_vae_key(key_name) and not is_text_encoder_key(key_name)
elif target_components == "VAE":
return is_vae_key(key_name)
elif target_components == "Text Encoders":
return is_text_encoder_key(key_name)
elif target_components == "Time Embeddings":
return _get_unet_component_type(key_name) == 'time_embed'
elif target_components == "Input Blocks":
return _get_unet_component_type(key_name) == 'input_blocks'
elif target_components == "Middle Block":
return _get_unet_component_type(key_name) == 'middle_block'
elif target_components == "Output Blocks":
return _get_unet_component_type(key_name) == 'output_blocks'
elif target_components == "Final UNet Output Layer":
return _get_unet_component_type(key_name) == 'out'
elif target_components == "Custom Pattern":
key_lower = key_name.lower()
for pattern in patterns:
if exact_match_custom:
if key_lower == pattern.lower():
return True
else:
if pattern.lower() in key_lower:
return True
return False
return False
def process_key_batch(self, batch_keys: List[str], target_components: str,
filter_mode: str, default_value: float, target_value: float,
patterns: List[str], exact_match_custom: bool) -> Dict[str, float]:
"""Process a batch of keys to create mask values"""
batch_results = {}
for key in batch_keys:
is_match = self.check_key_match(key, target_components, patterns, exact_match_custom)
if filter_mode == "Include":
batch_results[key] = target_value if is_match else default_value
elif filter_mode == "Exclude":
batch_results[key] = default_value if is_match else target_value
return batch_results
def filter_keys_to_mask(self, model, filter_mode, target_components, default_value, target_value,
memory_limit_gb=2.0, custom_pattern="", exact_match_custom=False):
print(f"\n--- Model Key Filter (Tensor Prism) ---")
print(f" Filter Mode: {filter_mode}")
print(f" Target Components: {target_components}")
print(f" Default Value: {default_value}, Target Value: {target_value}")
print(f" Memory Limit: {memory_limit_gb:.1f}GB")
# Get initial memory info
used_memory, available_memory = self.get_memory_info()
print(f" System Memory - Used: {used_memory:.2f}GB, Available: {available_memory:.2f}GB")
# Get state dict and keys
state_dict = model.model.state_dict()
all_keys = list(state_dict.keys())
print(f" Total model keys to process: {len(all_keys)}")
# Parse custom patterns
patterns = [p.strip() for p in custom_pattern.split(',') if p.strip()]
if target_components == "Custom Pattern" and patterns:
print(f" Custom patterns: {patterns}")
# Create processing batches
print(" Creating memory-efficient processing batches...")
batches = self.create_key_batches(all_keys, memory_limit_gb)
print(f" Created {len(batches)} processing batches")
# Process batches
mask_dict = {}
total_keys = len(all_keys)
processed_keys = 0
for i, batch_keys in enumerate(batches):
batch_memory = self.estimate_dict_memory_gb(len(batch_keys))
print(f" Processing batch {i+1}/{len(batches)} ({len(batch_keys)} keys, ~{batch_memory:.3f}GB)")
# Process this batch
batch_results = self.process_key_batch(
batch_keys, target_components, filter_mode, default_value, target_value,
patterns, exact_match_custom
)
# Add results to mask dict
mask_dict.update(batch_results)
processed_keys += len(batch_keys)
# Force garbage collection after each batch
del batch_results
gc.collect()
# Progress update
progress = (processed_keys / total_keys) * 100
used_memory, _ = self.get_memory_info()
print(f" Progress: {progress:.1f}% ({processed_keys}/{total_keys}), Memory: {used_memory:.2f}GB")
# Analyze model structure for layer info
print(" Analyzing model structure...")
layer_info = _analyze_model_structure(state_dict)
# Create the model mask object
mask = {
"mask_dict": mask_dict,
"mask_type": f"filtered_by_{target_components}_{filter_mode}",
"intensity": float(np.mean(list(mask_dict.values()))) if mask_dict else 0.0,
"layer_info": layer_info
}
# Final memory cleanup
del mask_dict
gc.collect()
final_memory, _ = self.get_memory_info()
print(f" Final memory usage: {final_memory:.2f}GB")
print(f" Mask intensity: {mask['intensity']:.4f}")
print(f" Found {layer_info['total_layers']} model layers")
print(f"--- Model Key Filter completed ---\n")
return (mask,)
NODE_CLASS_MAPPINGS = {
"TensorPrism_ModelKeyFilter": TensorPrism_ModelKeyFilter,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_ModelKeyFilter": "Model Key Filter (Tensor Prism)",
}
+180
View File
@@ -0,0 +1,180 @@
import copy
import gc
import psutil
from typing import Dict, List, Tuple
import numpy as np
import torch
MODEL_MASK_TYPE = ("MASK",)
class TensorPrism_ModelMaskBlender:
"""
Memory-efficient model mask blender that processes ComfyUI mask tensors
"""
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask_A": ("MASK",),
"mask_B": ("MASK",),
"blend_mode": (["Add", "Multiply", "Max", "Min", "Linear Blend", "Exponential Blend"], {"default": "Linear Blend"}),
"memory_limit_gb": ("FLOAT", {"default": 2.0, "min": 0.5, "max": 16.0, "step": 0.1, "round": 0.1, "label": "Memory Limit (GB)"}),
},
"optional": {
"blend_strength": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001, "label": "Blend Strength (for Linear/Exp)"}),
"clip_output": ("BOOLEAN", {"default": True, "label_on": "Clip to [0, 1]", "label_off": "No Clipping"}),
}
}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("combined_mask",)
FUNCTION = "blend_masks"
CATEGORY = "Tensor_Prism/Mask"
@staticmethod
def get_memory_info() -> Tuple[float, float]:
"""Get current memory usage and available memory in GB"""
memory = psutil.virtual_memory()
used_gb = (memory.total - memory.available) / (1024**3)
available_gb = memory.available / (1024**3)
return used_gb, available_gb
@staticmethod
def estimate_dict_memory_gb(dict_size: int) -> float:
"""Estimate memory usage of a dictionary with float values in GB"""
# Rough estimate: key string + float value + overhead
bytes_per_entry = 100 # Conservative estimate
return (dict_size * bytes_per_entry) / (1024**3)
def create_key_batches(self, all_keys: List[str], memory_limit_gb: float) -> List[List[str]]:
"""Create batches of keys that fit within memory limit"""
batches = []
current_batch = []
# Estimate how many keys we can process per batch
max_keys_per_batch = max(1000, int((memory_limit_gb * 1024**3) / 200)) # Conservative estimate
for i, key in enumerate(all_keys):
current_batch.append(key)
if len(current_batch) >= max_keys_per_batch:
batches.append(current_batch)
current_batch = []
# Add final batch if not empty
if current_batch:
batches.append(current_batch)
return batches
def process_mask_batch(self, batch_keys: List[str], mask_A_dict: Dict[str, float],
mask_B_dict: Dict[str, float], blend_mode: str,
blend_strength: float, clip_output: bool) -> Dict[str, float]:
"""Process a batch of mask keys with the specified blending operation"""
batch_results = {}
for key in batch_keys:
val_A = mask_A_dict.get(key, 0.0)
val_B = mask_B_dict.get(key, 0.0)
result_val = 0.0
if blend_mode == "Add":
result_val = val_A + val_B
elif blend_mode == "Multiply":
result_val = val_A * val_B
elif blend_mode == "Max":
result_val = max(val_A, val_B)
elif blend_mode == "Min":
result_val = min(val_A, val_B)
elif blend_mode == "Linear Blend":
result_val = val_A * (1.0 - blend_strength) + val_B * blend_strength
elif blend_mode == "Exponential Blend":
exp_strength = blend_strength ** 2
result_val = val_A * (1.0 - exp_strength) + val_B * exp_strength
if clip_output:
result_val = max(0.0, min(1.0, result_val))
batch_results[key] = result_val
return batch_results
def blend_masks(self, mask_A, mask_B, blend_mode, memory_limit_gb=2.0,
blend_strength=0.5, clip_output=True):
print(f"\n--- Model Mask Blender (Tensor Prism) ---")
print(f" Blend Mode: {blend_mode}")
print(f" Blend Strength: {blend_strength}")
print(f" Memory Limit: {memory_limit_gb:.1f}GB")
# Get initial memory info
used_memory, available_memory = self.get_memory_info()
print(f" System Memory - Used: {used_memory:.2f}GB, Available: {available_memory:.2f}GB")
# Convert tensors to numpy for processing
if isinstance(mask_A, torch.Tensor):
mask_A_np = mask_A.cpu().numpy()
else:
mask_A_np = np.array(mask_A)
if isinstance(mask_B, torch.Tensor):
mask_B_np = mask_B.cpu().numpy()
else:
mask_B_np = np.array(mask_B)
print(f" Mask A shape: {mask_A_np.shape}")
print(f" Mask B shape: {mask_B_np.shape}")
# Ensure masks have the same shape
if mask_A_np.shape != mask_B_np.shape:
# Resize mask_B to match mask_A
from scipy import ndimage
if len(mask_A_np.shape) == 3 and len(mask_B_np.shape) == 3:
mask_B_np = ndimage.zoom(mask_B_np,
(mask_A_np.shape[0]/mask_B_np.shape[0],
mask_A_np.shape[1]/mask_B_np.shape[1],
mask_A_np.shape[2]/mask_B_np.shape[2]))
elif len(mask_A_np.shape) == 2 and len(mask_B_np.shape) == 2:
mask_B_np = ndimage.zoom(mask_B_np,
(mask_A_np.shape[0]/mask_B_np.shape[0],
mask_A_np.shape[1]/mask_B_np.shape[1]))
# Process masks based on blend mode
if blend_mode == "Add":
result_mask = mask_A_np + mask_B_np
elif blend_mode == "Multiply":
result_mask = mask_A_np * mask_B_np
elif blend_mode == "Max":
result_mask = np.maximum(mask_A_np, mask_B_np)
elif blend_mode == "Min":
result_mask = np.minimum(mask_A_np, mask_B_np)
elif blend_mode == "Linear Blend":
result_mask = mask_A_np * (1.0 - blend_strength) + mask_B_np * blend_strength
elif blend_mode == "Exponential Blend":
exp_strength = blend_strength ** 2
result_mask = mask_A_np * (1.0 - exp_strength) + mask_B_np * exp_strength
else:
result_mask = mask_A_np
if clip_output:
result_mask = np.clip(result_mask, 0.0, 1.0)
result_tensor = torch.from_numpy(result_mask).float()
# Final memory cleanup
gc.collect()
final_memory, _ = self.get_memory_info()
print(f" Final memory usage: {final_memory:.2f}GB")
print(f" Result mask shape: {result_tensor.shape}")
print(f"--- Model Mask Blender completed ---\n")
return (result_tensor,)
NODE_CLASS_MAPPINGS = {
"TensorPrism_ModelMaskBlender": TensorPrism_ModelMaskBlender,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_ModelMaskBlender": "Mask Blender (Tensor Prism)",
}
+346
View File
@@ -0,0 +1,346 @@
import torch
import numpy as np
import re
class TensorPrism_ModelMaskGenerator:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"mask_type": (["layer_based", "block_based", "attention_only", "feedforward_only", "custom_pattern", "random_sparse", "depth_gradient"],),
"intensity": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.05,
"help": "Overall mask intensity"
}),
"reference_model": ("MODEL", {
"help": "Model to analyze structure from"
}),
},
"optional": {
"layer_start": ("INT", {
"default": 0, "min": 0, "max": 50, "step": 1,
"help": "Starting layer for range-based masks"
}),
"layer_end": ("INT", {
"default": -1, "min": -1, "max": 50, "step": 1,
"help": "Ending layer (-1 for all)"
}),
"gradient_direction": (["shallow_to_deep", "deep_to_shallow", "center_out", "edges_in"],),
"sparsity": ("FLOAT", {
"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.05,
"help": "Sparsity for random masks (0=dense, 1=sparse)"
}),
"custom_pattern": ("STRING", {
"default": "attn,mlp.fc1",
"help": "Comma-separated layer name patterns"
}),
"falloff": ("FLOAT", {
"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.05,
"help": "Gradient falloff smoothness"
}),
}
}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("mask",)
FUNCTION = "generate_mask"
CATEGORY = "Tensor Prism/Mask"
def generate_mask(self, mask_type, intensity, reference_model,
layer_start=0, layer_end=-1, gradient_direction="shallow_to_deep",
sparsity=0.5, custom_pattern="attn,mlp", falloff=0.1):
"""
Generate masks that work with actual model layer structure for merging.
"""
# Get model structure
state_dict = reference_model.model.state_dict()
layer_info = self._analyze_model_structure(state_dict)
# Generate mask based on type
if mask_type == "layer_based":
mask = self._create_layer_range_mask(layer_info, layer_start, layer_end, intensity)
elif mask_type == "block_based":
mask = self._create_block_mask(layer_info, layer_start, layer_end, intensity)
elif mask_type == "attention_only":
mask = self._create_component_mask(layer_info, ["attn", "attention", "self_attn"], intensity)
elif mask_type == "feedforward_only":
mask = self._create_component_mask(layer_info, ["mlp", "fc", "feedforward", "ffn"], intensity)
elif mask_type == "custom_pattern":
patterns = [p.strip() for p in custom_pattern.split(",")]
mask = self._create_pattern_mask(layer_info, patterns, intensity)
elif mask_type == "random_sparse":
mask = self._create_random_mask(layer_info, sparsity, intensity)
elif mask_type == "depth_gradient":
mask = self._create_depth_gradient_mask(layer_info, gradient_direction, intensity, falloff)
else:
mask = self._create_uniform_mask(layer_info, intensity)
# Package as MASK type
mask = {
"mask_dict": mask,
"mask_type": mask_type,
"intensity": intensity,
"layer_info": layer_info
}
return (mask,)
def _analyze_model_structure(self, state_dict):
"""Analyze model structure to understand layers and components"""
layer_info = {
"layers": {},
"total_layers": 0,
"layer_names": list(state_dict.keys())
}
# Pattern matching for different model architectures
layer_patterns = [
r"layers\.(\d+)\.", # Transformer layers
r"blocks\.(\d+)\.", # Vision transformer blocks
r"h\.(\d+)\.", # GPT-style layers
r"layer\.(\d+)\.", # BERT-style layers
r"encoder\.layer\.(\d+)\.", # BERT encoder
r"decoder\.layer\.(\d+)\.", # BERT decoder
]
for name in state_dict.keys():
layer_num = None
# Try to extract layer number
for pattern in layer_patterns:
match = re.search(pattern, name)
if match:
layer_num = int(match.group(1))
break
if layer_num is not None:
if layer_num not in layer_info["layers"]:
layer_info["layers"][layer_num] = []
layer_info["layers"][layer_num].append(name)
layer_info["total_layers"] = max(layer_info["total_layers"], layer_num + 1)
return layer_info
def _create_layer_range_mask(self, layer_info, start, end, intensity):
"""Create mask for specific layer range"""
mask = {}
if end == -1:
end = layer_info["total_layers"]
for name in layer_info["layer_names"]:
# Find which layer this parameter belongs to
layer_num = self._get_layer_number(name)
if layer_num is not None and start <= layer_num < end:
mask[name] = intensity
else:
mask[name] = 0.0
return mask
def _create_block_mask(self, layer_info, start, end, intensity):
"""Create mask for transformer blocks with smooth transitions"""
mask = {}
if end == -1:
end = layer_info["total_layers"]
total_range = max(end - start, 1)
for name in layer_info["layer_names"]:
layer_num = self._get_layer_number(name)
if layer_num is not None and start <= layer_num < end:
# Create smooth gradient within the range
progress = (layer_num - start) / total_range
mask_value = intensity * (0.5 + 0.5 * np.cos(progress * np.pi))
mask[name] = mask_value
else:
mask[name] = 0.0
return mask
def _create_component_mask(self, layer_info, component_patterns, intensity):
"""Create mask for specific components (attention, MLP, etc.)"""
mask = {}
for name in layer_info["layer_names"]:
should_mask = False
for pattern in component_patterns:
if pattern.lower() in name.lower():
should_mask = True
break
mask[name] = intensity if should_mask else 0.0
return mask
def _create_pattern_mask(self, layer_info, patterns, intensity):
"""Create mask based on custom patterns"""
mask = {}
for name in layer_info["layer_names"]:
should_mask = False
for pattern in patterns:
if pattern.lower() in name.lower():
should_mask = True
break
mask[name] = intensity if should_mask else 0.0
return mask
def _create_random_mask(self, layer_info, sparsity, intensity):
"""Create random sparse mask"""
mask = {}
np.random.seed(42) # For reproducibility
for name in layer_info["layer_names"]:
if np.random.random() > sparsity:
mask[name] = intensity
else:
mask[name] = 0.0
return mask
def _create_depth_gradient_mask(self, layer_info, direction, intensity, falloff):
"""Create gradient mask based on model depth"""
mask = {}
total_layers = max(layer_info["total_layers"], 1)
for name in layer_info["layer_names"]:
layer_num = self._get_layer_number(name)
if layer_num is not None:
# Calculate position (0 to 1)
position = layer_num / (total_layers - 1) if total_layers > 1 else 0.5
if direction == "shallow_to_deep":
mask_value = position
elif direction == "deep_to_shallow":
mask_value = 1.0 - position
elif direction == "center_out":
mask_value = 1.0 - 2.0 * abs(position - 0.5)
else: # edges_in
mask_value = 2.0 * abs(position - 0.5)
# Apply falloff
if falloff > 0:
mask_value = np.power(mask_value, 1.0 / max(falloff, 0.01))
mask[name] = mask_value * intensity
else:
# For parameters not in layers (embeddings, etc.)
mask[name] = intensity * 0.1
return mask
def _create_uniform_mask(self, layer_info, intensity):
"""Create uniform mask for all parameters"""
mask = {}
for name in layer_info["layer_names"]:
mask[name] = intensity
return mask
def _get_layer_number(self, param_name):
"""Extract layer number from parameter name"""
layer_patterns = [
r"layers\.(\d+)\.",
r"blocks\.(\d+)\.",
r"h\.(\d+)\.",
r"layer\.(\d+)\.",
r"encoder\.layer\.(\d+)\.",
r"decoder\.layer\.(\d+)\.",
]
for pattern in layer_patterns:
match = re.search(pattern, param_name)
if match:
return int(match.group(1))
return None
# Updated merge node to work with MASK
class TensorPrism_WeightedMaskMerge:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_A": ("MODEL",),
"model_B": ("MODEL",),
"mask": ("MASK",),
"merge_ratio": ("FLOAT", {
"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01,
"help": "Global multiplier for mask intensity"
}),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("merged_model",)
FUNCTION = "merge_models"
CATEGORY = "Tensor Prism/Mask"
def merge_models(self, model_A, model_B, mask, merge_ratio):
"""
Merge models using structure-aware masks
"""
# Clone model A as base
merged_model = model_A.clone()
# Get state dicts
state_dict_A = model_A.model.state_dict()
state_dict_B = model_B.model.state_dict()
# Get mask dictionary
mask_dict = mask["mask_dict"]
# Create patches
patches = {}
for key in state_dict_A.keys():
if key in state_dict_B and key in mask_dict:
weight_A = state_dict_A[key]
weight_B = state_dict_B[key]
if isinstance(weight_A, torch.Tensor) and isinstance(weight_B, torch.Tensor):
if weight_A.shape == weight_B.shape:
# Get mask value for this parameter
mask_value = mask_dict[key] * merge_ratio
# Calculate merged weights
merged_weight = weight_A * (1 - mask_value) + weight_B * mask_value
# Store as patch
if mask_value > 0:
patches[key] = (merged_weight - weight_A,)
# Apply patches
if patches:
merged_model.add_patches(patches, 1.0)
return (merged_model,)
NODE_CLASS_MAPPINGS = {
"TensorPrism_ModelMaskGenerator": TensorPrism_ModelMaskGenerator,
"TensorPrism_WeightedMaskMerge": TensorPrism_WeightedMaskMerge
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_ModelMaskGenerator": "Model Mask Generator (Tensor Prism)",
"TensorPrism_WeightedMaskMerge": "Weighted Mask Merge (Tensor Prism)"
}
+305
View File
@@ -0,0 +1,305 @@
"""
ComfyUI Tensor Prism Node Pack
===============================
Advanced model merging and enhancement nodes for ComfyUI, providing sophisticated
techniques for blending, enhancing, and manipulating Stable Diffusion models with
GPU-optimized memory management.
Author: AstrionX
Version: 1.1.0
License: GPL-3.0
Repository: https://github.com/AstrionX/ComfyUI-Tensor-Prism-Node-Pack
Features:
- Advanced model merging with multiple interpolation methods
- Spectral frequency-domain merging
- Granular SDXL block control
- Sophisticated masking system
- GPU-optimized memory management
- Cross-platform compatibility (CUDA/MPS/CPU)
"""
import os
import sys
import importlib.util
from pathlib import Path
# Add current directory to path for imports
current_dir = Path(__file__).parent
sys.path.insert(0, str(current_dir))
try:
from .TensorPrism_MainMerge import TensorPrism_MainMerge
except ImportError:
TensorPrism_MainMerge = None
try:
from .TensorPrism_SDXLBlockMerge import SDXLBlockMergeTensorPrism
except ImportError:
SDXLBlockMergeTensorPrism = None
try:
from .TensorPrism_SDXLAdvancedBlockmerge import SDXLAdvancedBlockMergeTensorPrism
except ImportError:
SDXLAdvancedBlockMergeTensorPrism = None
try:
from .TensorPrism_ModelMaskGenerator import TensorPrism_ModelMaskGenerator
except ImportError:
TensorPrism_ModelMaskGenerator = None
try:
from .TensorPrism_ModelKeyFilter import TensorPrism_ModelKeyFilter
except ImportError:
TensorPrism_ModelKeyFilter = None
try:
from .TensorPrism_ModelMaskBlender import TensorPrism_ModelMaskBlender
except ImportError:
TensorPrism_ModelMaskBlender = None
try:
from .TensorPrism_ModelWeightModifier import TensorPrism_ModelWeightModifier
except ImportError:
TensorPrism_ModelWeightModifier = None
try:
from .TensorPrism_WeightedTensorMerge import TensorPrism_WeightedMaskMerge as TensorPrism_WeightedTensorMerge_Class
except ImportError:
TensorPrism_WeightedTensorMerge_Class = None
try:
from .TensorPrism_Enhancer import ModelEnhancerTensorPrism
except ImportError:
ModelEnhancerTensorPrism = None
try:
from .TensorPrism_LayeredBlend import TensorPrism_LayeredBlend
except ImportError:
TensorPrism_LayeredBlend = None
try:
from .TensorPrism_Prism import TensorPrism_FastPrism
except ImportError:
TensorPrism_FastPrism = None
# Version info
__version__ = "1.1.0"
__author__ = "AstrionX"
__description__ = "Advanced model merging and enhancement nodes for ComfyUI"
NODE_CLASS_MAPPINGS = {}
# Core merging nodes
if TensorPrism_MainMerge:
NODE_CLASS_MAPPINGS["TensorPrism_MainMerge"] = TensorPrism_MainMerge
# SDXL specific nodes
if SDXLBlockMergeTensorPrism:
NODE_CLASS_MAPPINGS["SDXL Block Merge (Tensor Prism)"] = SDXLBlockMergeTensorPrism
if SDXLAdvancedBlockMergeTensorPrism:
NODE_CLASS_MAPPINGS["SDXLAdvancedBlockMergeTensorPrism"] = SDXLAdvancedBlockMergeTensorPrism
# Masking and filtering nodes
if TensorPrism_ModelMaskGenerator:
NODE_CLASS_MAPPINGS["TensorPrism_ModelMaskGenerator"] = TensorPrism_ModelMaskGenerator
if TensorPrism_ModelKeyFilter:
NODE_CLASS_MAPPINGS["TensorPrism_ModelKeyFilter"] = TensorPrism_ModelKeyFilter
if TensorPrism_ModelMaskBlender:
NODE_CLASS_MAPPINGS["TensorPrism_ModelMaskBlender"] = TensorPrism_ModelMaskBlender
# Enhancement and transformation nodes
if TensorPrism_ModelWeightModifier:
NODE_CLASS_MAPPINGS["TensorPrism_ModelWeightModifier"] = TensorPrism_ModelWeightModifier
if TensorPrism_WeightedTensorMerge_Class:
NODE_CLASS_MAPPINGS["TensorPrism_WeightedTensorMerge"] = TensorPrism_WeightedTensorMerge_Class
if ModelEnhancerTensorPrism:
NODE_CLASS_MAPPINGS["ModelEnhancerTensorPrism"] = ModelEnhancerTensorPrism
if TensorPrism_LayeredBlend:
NODE_CLASS_MAPPINGS["TensorPrism_LayeredBlend"] = TensorPrism_LayeredBlend
if TensorPrism_FastPrism:
NODE_CLASS_MAPPINGS["TensorPrism_Prism"] = TensorPrism_FastPrism
NODE_DISPLAY_NAME_MAPPINGS = {}
# Core merging nodes
if "TensorPrism_MainMerge" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_MainMerge"] = "Main Merge (Tensor Prism)"
# SDXL specific nodes
if "SDXL Block Merge (Tensor Prism)" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["SDXL Block Merge (Tensor Prism)"] = "SDXL Block Merge (Tensor Prism)"
if "SDXLAdvancedBlockMergeTensorPrism" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["SDXLAdvancedBlockMergeTensorPrism"] = "SDXL Advanced Block Merge (Tensor Prism)"
# Masking and filtering nodes
if "TensorPrism_ModelMaskGenerator" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelMaskGenerator"] = "Model Mask Generator (Tensor Prism)"
if "TensorPrism_ModelKeyFilter" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelKeyFilter"] = "Model Key Filter (Tensor Prism)"
if "TensorPrism_ModelMaskBlender" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelMaskBlender"] = "Mask Blender (Tensor Prism)"
# Enhancement and transformation nodes
if "TensorPrism_ModelWeightModifier" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelWeightModifier"] = "Model Weight Modifier (Tensor Prism)"
if "TensorPrism_WeightedTensorMerge" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_WeightedTensorMerge"] = "Weighted Tensor Merge (Tensor Prism)"
if "ModelEnhancerTensorPrism" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["ModelEnhancerTensorPrism"] = "Model Enhancer (Tensor Prism)"
if "TensorPrism_LayeredBlend" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_LayeredBlend"] = "Layered Blend (Tensor Prism)"
if "TensorPrism_Prism" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_Prism"] = "Prism (Tensor Prism)"
NODE_CATEGORIES = {}
for node_key in NODE_CLASS_MAPPINGS.keys():
if "MainMerge" in node_key or "LayeredBlend" in node_key:
NODE_CATEGORIES[node_key] = "Tensor Prism/Core"
elif "SDXL" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Merge"
elif "Mask" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Mask"
elif "Weight" in node_key or "Enhancer" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Transform"
else:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Core"
def check_dependencies():
"""Check if required dependencies are available."""
required_packages = [
("torch", "PyTorch >= 1.12.0"),
("numpy", "NumPy >= 1.21.0"),
("psutil", "psutil >= 5.8.0"),
]
missing_packages = []
for package_name, description in required_packages:
try:
importlib.import_module(package_name)
except ImportError:
missing_packages.append(description)
if missing_packages:
print(f"[TensorPrism] Warning: Missing dependencies:")
for pkg in missing_packages:
print(f" - {pkg}")
print("[TensorPrism] Some features may not work properly.")
return len(missing_packages) == 0
def get_system_info():
"""Get system information for optimization."""
try:
import torch
import psutil
info = {
"torch_version": torch.__version__,
"cuda_available": torch.cuda.is_available(),
"mps_available": hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(),
"cpu_count": psutil.cpu_count(),
"total_ram": round(psutil.virtual_memory().total / (1024**3), 1),
}
if info["cuda_available"]:
info["cuda_device_count"] = torch.cuda.device_count()
info["cuda_memory"] = round(torch.cuda.get_device_properties(0).total_memory / (1024**3), 1)
return info
except Exception as e:
print(f"[TensorPrism] Could not get system info: {e}")
return {}
def print_welcome_message():
"""Print welcome message with system information."""
print("\n" + "="*60)
print("🎭 TensorPrism Node Pack Loaded")
print("="*60)
print(f"Version: {__version__}")
print(f"Author: {__author__}")
print(f"Nodes: {len(NODE_CLASS_MAPPINGS)}")
# System info
sys_info = get_system_info()
if sys_info:
print(f"\n📊 System Information:")
print(f" PyTorch: {sys_info.get('torch_version', 'Unknown')}")
print(f" CPU Cores: {sys_info.get('cpu_count', 'Unknown')}")
print(f" RAM: {sys_info.get('total_ram', 'Unknown')}GB")
if sys_info.get("cuda_available"):
print(f" CUDA: Available ({sys_info.get('cuda_device_count', 0)} device(s))")
print(f" GPU Memory: {sys_info.get('cuda_memory', 'Unknown')}GB")
elif sys_info.get("mps_available"):
print(f" MPS: Available (Apple Silicon)")
else:
print(f" GPU: CPU fallback mode")
print(f"\n🚀 Available Nodes:")
for category in set(NODE_CATEGORIES.values()):
print(f"\n {category}:")
for node_class, node_category in NODE_CATEGORIES.items():
if node_category == category:
display_name = NODE_DISPLAY_NAME_MAPPINGS.get(node_class, node_class)
print(f" • {display_name}")
print(f"\n💡 Memory Management:")
if sys_info.get("cuda_memory", 0) >= 24:
print(" Recommended: Default settings (24GB+ GPU)")
elif sys_info.get("cuda_memory", 0) >= 12:
print(" Recommended: Memory limit 8GB, auto precision")
elif sys_info.get("cuda_memory", 0) >= 8:
print(" Recommended: Memory limit 6GB, CPU fallback for large merges")
else:
print(" Recommended: CPU processing for best stability")
print("="*60 + "\n")
# Initialize the package
def __init_package():
"""Initialize the package and perform startup checks."""
try:
# Check dependencies
deps_ok = check_dependencies()
# Print welcome message
print_welcome_message()
if not deps_ok:
print("[TensorPrism] Warning: Some dependencies are missing. Please install them for full functionality.")
print("[TensorPrism] Package initialized successfully!")
except Exception as e:
print(f"[TensorPrism] Error during initialization: {e}")
print("[TensorPrism] Package may not function correctly.")
# Run initialization
__init_package()
# Export for ComfyUI
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
"__version__",
"__author__",
"__description__"
]
+284
View File
@@ -0,0 +1,284 @@
import torch
import math
import comfy.model_management
class TensorPrism_FastPrism:
"""
Fast Prism - Streamlined spectral merging with much better performance
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_A": ("MODEL",),
"model_B": ("MODEL",),
"merge_method": ([
"spectral_blend", "frequency_bands", "magnitude_weighted",
"adaptive_mix", "harmonic_merge"
],),
"strength": ("FLOAT", {
"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01
}),
},
"optional": {
"low_freq_bias": ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.0, "step": 0.05}),
"high_freq_bias": ("FLOAT", {"default": 0.3, "min": 0.0, "max": 1.0, "step": 0.05}),
"spectral_precision": (["fast", "balanced", "precise"], {"default": "fast"}),
"layer_selective": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("merged_model",)
FUNCTION = "fast_prism_merge"
CATEGORY = "Tensor Prism/Core"
def fast_prism_merge(self, model_A, model_B, merge_method, strength,
low_freq_bias=0.7, high_freq_bias=0.3,
spectral_precision="fast", layer_selective=False):
# Clone model efficiently
merged_model = model_A.clone()
# Get state dictionaries
state_dict_A = model_A.model.state_dict()
state_dict_B = model_B.model.state_dict()
patches = {}
# Process parameters with efficient batching
for key in state_dict_A.keys():
if key not in state_dict_B:
continue
tensor_A = state_dict_A[key]
tensor_B = state_dict_B[key]
if not (isinstance(tensor_A, torch.Tensor) and isinstance(tensor_B, torch.Tensor)):
continue
if tensor_A.shape != tensor_B.shape:
continue
# Skip tiny tensors
if tensor_A.numel() < 10:
continue
# DEVICE FIX: Ensure tensors are on the same device
device = tensor_A.device
if tensor_B.device != device:
tensor_B = tensor_B.to(device)
# Layer-selective processing
if layer_selective:
layer_strength = self._get_layer_strength(key, strength)
else:
layer_strength = strength
try:
# Apply fast spectral merge
merged_tensor = self._fast_spectral_merge(
tensor_A, tensor_B, merge_method, layer_strength,
low_freq_bias, high_freq_bias, spectral_precision
)
# DEVICE FIX: Ensure merged tensor is on correct device
merged_tensor = merged_tensor.to(device)
# Only create patch if there's a meaningful difference
diff = merged_tensor - tensor_A
if torch.abs(diff).max() > 1e-6:
patches[key] = (diff,)
except Exception as e:
# Silent fallback to linear merge
merged_tensor = tensor_A * (1 - layer_strength) + tensor_B * layer_strength
diff = merged_tensor - tensor_A
if torch.abs(diff).max() > 1e-6:
patches[key] = (diff,)
# Apply patches efficiently
if patches:
merged_model.add_patches(patches, 1.0)
return (merged_model,)
def _get_layer_strength(self, layer_name, base_strength):
"""Get layer-specific strength multipliers"""
layer_name = layer_name.lower()
# Different strengths for different layer types
if any(x in layer_name for x in ['attention', 'attn']):
return base_strength * 0.8 # More conservative for attention
elif any(x in layer_name for x in ['norm', 'ln']):
return base_strength * 0.6 # Very conservative for normalization
elif any(x in layer_name for x in ['embed', 'pos']):
return base_strength * 0.4 # Most conservative for embeddings
elif any(x in layer_name for x in ['output', 'head']):
return base_strength * 0.9 # Slightly more aggressive for output layers
else:
return base_strength
def _fast_spectral_merge(self, tensor_A, tensor_B, method, strength,
low_bias, high_bias, precision):
"""Fast spectral merging methods"""
# DEVICE FIX: Get device from tensor_A and ensure consistency
device = tensor_A.device
if method == "spectral_blend":
return self._spectral_blend_fast(tensor_A, tensor_B, strength, low_bias, high_bias)
elif method == "frequency_bands":
return self._frequency_bands_fast(tensor_A, tensor_B, strength, low_bias, high_bias, precision)
elif method == "magnitude_weighted":
return self._magnitude_weighted_fast(tensor_A, tensor_B, strength)
elif method == "adaptive_mix":
return self._adaptive_mix_fast(tensor_A, tensor_B, strength)
elif method == "harmonic_merge":
return self._harmonic_merge_fast(tensor_A, tensor_B, strength, low_bias, high_bias)
else:
return tensor_A * (1 - strength) + tensor_B * strength
def _spectral_blend_fast(self, tensor_A, tensor_B, strength, low_bias, high_bias):
"""Fast spectral blend using simple magnitude-based frequency separation"""
device = tensor_A.device
# Simple magnitude-based frequency separation
mag_A = torch.abs(tensor_A)
mag_B = torch.abs(tensor_B)
# Use magnitude as a proxy for frequency content
threshold = (mag_A.mean() + mag_B.mean()) / 2
# Create frequency masks
low_freq_mask = (mag_A + mag_B) > threshold
high_freq_mask = ~low_freq_mask
# Apply frequency-specific blending
result = torch.zeros_like(tensor_A, device=device)
result[low_freq_mask] = (tensor_A[low_freq_mask] * (1 - strength * low_bias) +
tensor_B[low_freq_mask] * (strength * low_bias))
result[high_freq_mask] = (tensor_A[high_freq_mask] * (1 - strength * high_bias) +
tensor_B[high_freq_mask] * (strength * high_bias))
return result
def _frequency_bands_fast(self, tensor_A, tensor_B, strength, low_bias, high_bias, precision):
"""Fast frequency band separation"""
device = tensor_A.device
if tensor_A.dim() < 2 or precision == "fast":
# For 1D or fast mode, use simple variance-based separation
var_A = torch.var(tensor_A, dim=-1, keepdim=True)
var_B = torch.var(tensor_B, dim=-1, keepdim=True)
# High variance = high frequency content
freq_weight = torch.sigmoid((var_A + var_B) - (var_A.mean() + var_B.mean()))
# Interpolate between low and high frequency bias
effective_bias = low_bias * (1 - freq_weight) + high_bias * freq_weight
effective_strength = strength * effective_bias
return tensor_A * (1 - effective_strength) + tensor_B * effective_strength
else:
# For higher dimensions, use simple 1D FFT on flattened tensor
try:
flat_A = tensor_A.flatten()
flat_B = tensor_B.flatten()
# DEVICE FIX: Explicitly keep tensors on GPU for FFT
if device.type == 'cuda':
# Use CUDA FFT if available
fft_A = torch.fft.fft(flat_A.float())
fft_B = torch.fft.fft(flat_B.float())
else:
# Move to CPU for FFT, then back
fft_A = torch.fft.fft(flat_A.cpu().float()).to(device)
fft_B = torch.fft.fft(flat_B.cpu().float()).to(device)
# Simple frequency band separation
n = len(fft_A)
low_cutoff = n // 4 # First quarter = low freq
high_cutoff = 3 * n // 4 # Last quarter = high freq
# Merge different frequency bands with different weights
fft_merged = fft_A.clone()
fft_merged[:low_cutoff] = (fft_A[:low_cutoff] * (1 - strength * low_bias) +
fft_B[:low_cutoff] * (strength * low_bias))
fft_merged[high_cutoff:] = (fft_A[high_cutoff:] * (1 - strength * high_bias) +
fft_B[high_cutoff:] * (strength * high_bias))
# Middle frequencies use regular strength
fft_merged[low_cutoff:high_cutoff] = (fft_A[low_cutoff:high_cutoff] * (1 - strength) +
fft_B[low_cutoff:high_cutoff] * strength)
# Convert back to spatial domain
if device.type == 'cuda':
merged_flat = torch.fft.ifft(fft_merged).real.to(tensor_A.dtype)
else:
merged_flat = torch.fft.ifft(fft_merged.cpu()).real.to(tensor_A.dtype).to(device)
return merged_flat.view(tensor_A.shape)
except Exception as e:
# Fallback to spectral blend
return self._spectral_blend_fast(tensor_A, tensor_B, strength, low_bias, high_bias)
def _magnitude_weighted_fast(self, tensor_A, tensor_B, strength):
"""Magnitude-weighted merging - stronger tensor gets more influence"""
mag_A = torch.norm(tensor_A)
mag_B = torch.norm(tensor_B)
if mag_A + mag_B > 1e-8:
# Weight by relative magnitudes
weight_B = mag_B / (mag_A + mag_B)
# Modulate by input strength
effective_strength = strength * weight_B + (1 - strength) * 0.5
else:
effective_strength = strength
return tensor_A * (1 - effective_strength) + tensor_B * effective_strength
def _adaptive_mix_fast(self, tensor_A, tensor_B, strength):
"""Adaptive mixing based on tensor similarity"""
# Calculate cosine similarity efficiently
flat_A = tensor_A.flatten()
flat_B = tensor_B.flatten()
similarity = torch.cosine_similarity(flat_A, flat_B, dim=0)
# For very similar tensors, use more conservative merging
# For dissimilar tensors, use standard merging
adaptive_strength = strength * (1 - 0.3 * torch.abs(similarity))
return tensor_A * (1 - adaptive_strength) + tensor_B * adaptive_strength
def _harmonic_merge_fast(self, tensor_A, tensor_B, strength, low_bias, high_bias):
"""Simple harmonic merging using phase relationships"""
# Use sign patterns as a simple phase proxy
sign_A = torch.sign(tensor_A)
sign_B = torch.sign(tensor_B)
# Where signs align, use low_bias (harmonic)
# Where signs oppose, use high_bias (less harmonic)
alignment = (sign_A * sign_B) > 0
effective_bias = torch.where(alignment, low_bias, high_bias)
effective_strength = strength * effective_bias
return tensor_A * (1 - effective_strength) + tensor_B * effective_strength
NODE_CLASS_MAPPINGS = {
"TensorPrism_Prism": TensorPrism_FastPrism
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_Prism": "Prism (Tensor Prism)"
}
+245
View File
@@ -0,0 +1,245 @@
"""
OPTION 1: REMOVE ALL UNET BLOCK LAYERS (Simplified Advanced Version)
This keeps all the advanced memory management but removes per-block control
"""
import torch
import gc
import logging
import psutil
import threading
import time
from typing import Dict, List, Tuple, Optional, Union
import traceback
from contextmanager import contextmanager
# Set up logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class MemoryManager:
"""Advanced memory management for GPU/CPU processing."""
def __init__(self):
self.cuda_available = torch.cuda.is_available()
self.device_memory_gb = self._get_device_memory()
self.system_memory_gb = self._get_system_memory()
self.memory_threshold = 0.85 # Use max 85% of available memory
def _get_device_memory(self) -> float:
"""Get GPU memory in GB."""
if self.cuda_available:
try:
return torch.cuda.get_device_properties(0).total_memory / (1024**3)
except:
return 0.0
return 0.0
def _get_system_memory(self) -> float:
"""Get system RAM in GB."""
return psutil.virtual_memory().total / (1024**3)
def get_available_memory(self, device: torch.device) -> float:
"""Get currently available memory in GB."""
if device.type == "cuda" and self.cuda_available:
try:
free_memory = torch.cuda.get_device_properties(0).total_memory - torch.cuda.memory_allocated()
return free_memory / (1024**3)
except:
return 0.0
else:
available_memory = psutil.virtual_memory().available / (1024**3)
return available_memory
def estimate_tensor_memory(self, tensor: torch.Tensor) -> float:
"""Estimate tensor memory usage in GB."""
if tensor is None:
return 0.0
try:
element_size = tensor.element_size()
num_elements = tensor.numel()
return (element_size * num_elements) / (1024**3)
except:
return 0.0
def can_fit_in_memory(self, tensor: torch.Tensor, device: torch.device) -> bool:
"""Check if tensor can fit in device memory."""
tensor_memory = self.estimate_tensor_memory(tensor)
available_memory = self.get_available_memory(device) * self.memory_threshold
return tensor_memory <= available_memory
def should_use_cpu_fallback(self, device: torch.device) -> bool:
"""Determine if should fallback to CPU based on memory."""
if device.type == "cpu":
return False
available_gpu_memory = self.get_available_memory(device)
available_cpu_memory = self.get_available_memory(torch.device("cpu"))
return available_gpu_memory < 2.0 or (available_cpu_memory > available_gpu_memory * 2)
@contextmanager
def memory_context(self, device: torch.device):
"""Context manager for memory cleanup."""
try:
yield
finally:
self.cleanup_memory(device)
def cleanup_memory(self, device: torch.device = None):
"""Comprehensive memory cleanup."""
gc.collect()
if self.cuda_available:
try:
if device is None or device.type == "cuda":
torch.cuda.empty_cache()
torch.cuda.synchronize()
except:
pass
class SDXLAdvancedBlockMergeTensorPrism:
@classmethod
def INPUT_TYPES(cls) -> Dict:
inputs = {
"required": {
"model_A": ("MODEL", {}),
"model_B": ("MODEL", {}),
"merge_method": (["Linear Interpolation", "Add Difference", "TIES-Merging (Simplified)"],),
"default_unet_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"memory_limit_gb": ("FLOAT", {"default": 8.0, "min": 1.0, "max": 64.0, "step": 0.5, "round": 0.1}),
"force_cpu": ("BOOLEAN", {"default": False}),
"batch_size": ("INT", {"default": 50, "min": 1, "max": 500, "step": 10}),
"auto_memory_management": ("BOOLEAN", {"default": True}),
},
"optional": {
"model_C": ("MODEL", {}),
"ties_global_alpha_A": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"ties_global_alpha_B": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"rescale_output_magnitudes": ("BOOLEAN", {"default": False}),
"a_delta_factor": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.01, "round": 0.001}),
"b_delta_factor": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.01, "round": 0.001}),
"precision_mode": (["auto", "fp16", "fp32"], {"default": "auto"}),
"aggressive_cleanup": ("BOOLEAN", {"default": True}),
# Special components
"out_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"time_embed_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"label_emb_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
# Input blocks 0-11 (SDXL has 12)
"input_block_00_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_01_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_02_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_03_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_04_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_05_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_06_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_07_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_08_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_09_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_10_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_11_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
# Middle blocks 0-2
"middle_block_00_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"middle_block_01_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"middle_block_02_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
# Output blocks 0-11 (SDXL has 12)
"output_block_00_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_01_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_02_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_03_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_04_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_05_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_06_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_07_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_08_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_09_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_10_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_11_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
}
}
return inputs
def _build_ratio_mapping_fixed(self, keys: List[str], default_ratio: float, kwargs: Dict) -> Dict[str, float]:
try:
default_ratio = max(0.0, min(1.0, default_ratio))
# Build prefix mappings
ratio_prefixes = {
"time_embed.": max(0.0, min(1.0, kwargs.get("time_embed_ratio", default_ratio))),
"label_emb.": max(0.0, min(1.0, kwargs.get("label_emb_ratio", default_ratio))),
"out.": max(0.0, min(1.0, kwargs.get("out_ratio", default_ratio))),
}
# Input blocks (0-11)
for i in range(12):
input_ratio = max(0.0, min(1.0, kwargs.get(f"input_block_{i:02d}_ratio", default_ratio)))
ratio_prefixes[f"input_blocks.{i}."] = input_ratio
# Middle blocks (0-2)
for i in range(3):
middle_ratio = max(0.0, min(1.0, kwargs.get(f"middle_block_{i:02d}_ratio", default_ratio)))
ratio_prefixes[f"middle_block.{i}."] = middle_ratio
# Output blocks (0-11)
for i in range(12):
output_ratio = max(0.0, min(1.0, kwargs.get(f"output_block_{i:02d}_ratio", default_ratio)))
ratio_prefixes[f"output_blocks.{i}."] = output_ratio
# Sort prefixes by length (descending)
sorted_prefixes = sorted(ratio_prefixes.items(), key=lambda x: len(x[0]), reverse=True)
# Build mapping
key_to_ratio = {}
block_stats = {}
for key in keys:
ratio = default_ratio
matched_prefix = None
for prefix, prefix_ratio in sorted_prefixes:
if key.startswith(prefix):
ratio = prefix_ratio
matched_prefix = prefix
break
key_to_ratio[key] = ratio
# Statistics
if matched_prefix:
block_stats[matched_prefix] = block_stats.get(matched_prefix, 0) + 1
else:
block_stats["unmatched"] = block_stats.get("unmatched", 0) + 1
# Log block statistics
logger.info("Block mapping statistics:")
for prefix, count in sorted(block_stats.items()):
if count > 0:
ratio = ratio_prefixes.get(prefix, default_ratio)
logger.info(f" {prefix}: {count} parameters (ratio: {ratio})")
return key_to_ratio
except Exception as e:
logger.error(f"Error building ratio mapping: {e}")
return {key: default_ratio for key in keys}
# ... (Keep all other methods the same, just use _build_ratio_mapping_fixed instead)
NODE_CLASS_MAPPINGS = {
"SDXLAdvancedBlockMergeTensorPrism": SDXLAdvancedBlockMergeTensorPrism
}
NODE_DISPLAY_NAME_MAPPINGS = {
"SDXLAdvancedBlockMergeTensorPrism": "SDXL Advanced Block Merge (Tensor Prism)"
}
+255
View File
@@ -0,0 +1,255 @@
"""
FIXED SDXL BLOCK MERGE - UNET BLOCK LAYERS CORRECTED
Key fixes:
1. Proper UNET block layer mapping and identification
2. Corrected prefix matching for SDXL architecture
3. Fixed middle block structure
4. Maintained ComfyUI patching system integration
"""
import torch
import copy
import gc
import psutil
from collections import defaultdict
from typing import Dict, List, Tuple, Optional
class SDXLBlockMergeTensorPrism:
"""
FIXED VERSION: Properly handles UNET block layers with correct SDXL architecture mapping
"""
@classmethod
def INPUT_TYPES(s):
inputs = {
"required": {
"model_A": ("MODEL", {}),
"model_B": ("MODEL", {}),
"merge_method": (["Linear Interpolation", "Add Difference", "TIES-Merging (Simplified)"],),
"default_unet_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"memory_limit_gb": ("FLOAT", {"default": 8.0, "min": 1.0, "max": 64.0, "step": 0.5, "round": 0.1}),
},
"optional": {
"model_C": ("MODEL", {}),
"ties_global_alpha_A": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"ties_global_alpha_B": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"rescale_output_magnitudes": ("BOOLEAN", {"default": False}),
"iterations": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"a_delta_factor": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.01, "round": 0.001}),
"b_delta_factor": ("FLOAT", {"default": 1.0, "min": -2.0, "max": 2.0, "step": 0.01, "round": 0.001}),
# Special components (these work fine)
"out_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"time_embed_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"label_emb_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
# FIXED: Proper SDXL UNET block structure
# Input blocks (0-11 for SDXL)
"input_block_00_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_01_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_02_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_03_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_04_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_05_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_06_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_07_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_08_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_09_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_10_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"input_block_11_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
# Middle blocks (0-2 for SDXL: resnet -> attention -> resnet)
"middle_block_00_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"middle_block_01_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"middle_block_02_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
# Output blocks (0-11 for SDXL)
"output_block_00_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_01_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_02_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_03_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_04_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_05_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_06_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_07_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_08_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_09_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_10_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
"output_block_11_ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01, "round": 0.001}),
}
}
return inputs
RETURN_TYPES = ("MODEL",)
FUNCTION = "process"
CATEGORY = "Tensor_Prism/Merge"
def process(self, model_A, model_B, merge_method, default_unet_ratio, memory_limit_gb=8.0,
model_C=None, **kwargs):
"""
Process the model merge with proper UNET block handling
"""
print(f"\n--- SDXL Block Merge (Tensor Prism) ---")
# Validate inputs
if merge_method in ["Add Difference", "TIES-Merging (Simplified)"] and model_C is None:
raise ValueError(f"Model C is required for '{merge_method}' merge method.")
# Get state dictionaries
unet_sd_A = model_A.model.state_dict()
unet_sd_B = model_B.model.state_dict()
unet_sd_C = model_C.model.state_dict() if model_C else None
# Debug: Print some keys to understand structure
print("Sample UNET keys:")
sample_keys = list(unet_sd_A.keys())[:10]
for key in sample_keys:
print(f" {key}")
# Build ratio mapping with FIXED logic
key_to_ratio_map = self._build_ratio_mapping(unet_sd_A.keys(), default_unet_ratio, kwargs)
# Create patches instead of replacing entire state dict
patches = {}
print("Processing tensor merging...")
processed_keys = 0
for key in unet_sd_A.keys():
if key not in unet_sd_B:
continue
param_A = unet_sd_A[key]
param_B = unet_sd_B[key].to(param_A.device, param_A.dtype)
ratio = key_to_ratio_map.get(key, default_unet_ratio)
# Skip if ratio is 0 (no change from model A)
if ratio == 0.0:
continue
merged_param = None
if merge_method == "Linear Interpolation":
# ratio of 0.0 = all A, ratio of 1.0 = all B
merged_param = param_A * (1.0 - ratio) + param_B * ratio
elif merge_method == "Add Difference":
if unet_sd_C and key in unet_sd_C:
param_C = unet_sd_C[key].to(param_A.device, param_A.dtype)
delta_A = (param_A - param_C) * ratio * kwargs.get('a_delta_factor', 1.0)
delta_B = (param_B - param_C) * (1.0 - ratio) * kwargs.get('b_delta_factor', 1.0)
merged_param = param_C + delta_A + delta_B
else:
merged_param = param_A * (1.0 - ratio) + param_B * ratio
elif merge_method == "TIES-Merging (Simplified)":
if unet_sd_C and key in unet_sd_C:
param_C = unet_sd_C[key].to(param_A.device, param_A.dtype)
alpha_A = kwargs.get('ties_global_alpha_A', 0.5) * ratio
alpha_B = kwargs.get('ties_global_alpha_B', 0.5) * (1.0 - ratio)
if kwargs.get('rescale_output_magnitudes', False):
total = alpha_A + alpha_B
if total > 0:
alpha_A /= total
alpha_B /= total
delta_A = (param_A - param_C) * alpha_A * kwargs.get('a_delta_factor', 1.0)
delta_B = (param_B - param_C) * alpha_B * kwargs.get('b_delta_factor', 1.0)
merged_param = param_C + delta_A + delta_B
else:
merged_param = param_A * (1.0 - ratio) + param_B * ratio
if merged_param is not None:
# Store as patch (difference from original)
patch_diff = merged_param - param_A
# Only add patch if there's actually a difference
if not torch.allclose(patch_diff, torch.zeros_like(patch_diff), atol=1e-8):
patches[key] = (patch_diff,)
processed_keys += 1
print(f"Created {len(patches)} patches from {processed_keys} processed keys")
# Clone model and apply patches properly
merged_model = model_A.clone()
if patches:
merged_model.add_patches(patches, 1.0)
print("--- SDXL Block Merge completed ---\n")
return (merged_model,)
def _build_ratio_mapping(self, keys, default_ratio, kwargs):
"""
FIXED: Build proper mapping of parameter keys to their merge ratios
Now correctly identifies SDXL UNET structure
"""
print("Building ratio mapping...")
# Create the prefix to ratio mapping
ratio_prefixes = {}
# Special components (these are correctly mapped)
ratio_prefixes["time_embed."] = kwargs.get("time_embed_ratio", default_ratio)
ratio_prefixes["label_emb."] = kwargs.get("label_emb_ratio", default_ratio)
ratio_prefixes["out."] = kwargs.get("out_ratio", default_ratio)
# FIXED: Proper SDXL UNET block mapping
# Input blocks (SDXL has 0-11)
for i in range(12):
prefix = f"input_blocks.{i}."
ratio_key = f"input_block_{i:02d}_ratio"
ratio_prefixes[prefix] = kwargs.get(ratio_key, default_ratio)
# Middle blocks (SDXL structure: 0=resnet, 1=attention, 2=resnet)
for i in range(3):
prefix = f"middle_block.{i}."
ratio_key = f"middle_block_{i:02d}_ratio"
ratio_prefixes[prefix] = kwargs.get(ratio_key, default_ratio)
# Output blocks (SDXL has 0-11)
for i in range(12):
prefix = f"output_blocks.{i}."
ratio_key = f"output_block_{i:02d}_ratio"
ratio_prefixes[prefix] = kwargs.get(ratio_key, default_ratio)
# Sort by length (longest first) for proper prefix matching
sorted_prefixes = sorted(ratio_prefixes.items(), key=lambda x: len(x[0]), reverse=True)
# Build the key to ratio mapping
key_to_ratio = {}
block_stats = defaultdict(int)
for key in keys:
ratio = default_ratio
matched_prefix = None
# Find the longest matching prefix
for prefix, r in sorted_prefixes:
if key.startswith(prefix):
ratio = r
matched_prefix = prefix
block_stats[prefix] += 1
break
if matched_prefix is None:
block_stats["unmatched"] += 1
key_to_ratio[key] = ratio
# Debug output
print("Block mapping statistics:")
for prefix, count in sorted(block_stats.items()):
if count > 0:
ratio = ratio_prefixes.get(prefix, default_ratio)
print(f" {prefix}: {count} parameters (ratio: {ratio})")
return key_to_ratio
NODE_CLASS_MAPPINGS = {
"SDXL Block Merge (Tensor Prism)": SDXLBlockMergeTensorPrism
}
NODE_DISPLAY_NAME_MAPPINGS = {
"SDXL Block Merge (Tensor Prism)": "SDXL Block Merge (Tensor Prism)"
}
+74
View File
@@ -0,0 +1,74 @@
import torch
import comfy.utils
class TensorPrism_WeightedMaskMerge:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model_A": ("MODEL",),
"model_B": ("MODEL",),
"mask": ("MASK",),
"merge_ratio": ("FLOAT", {
"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01,
"help": "0.0 = Model A only, 1.0 = Model B only"
}),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("merged_model",)
FUNCTION = "merge_models"
CATEGORY = "Tensor Prism/Mask"
def merge_models(self, model_A, model_B, mask, merge_ratio):
"""
Blends two models (A, B) according to a mask and merge ratio.
Mask defines where Model B is applied. Merge ratio controls how strongly B overrides A.
The applied mask is also stored in the output model as 'mask_used'.
"""
merged = {}
# Ensure mask is broadcastable
if mask.ndim == 2:
mask = mask.unsqueeze(0).unsqueeze(0)
elif mask.ndim == 3:
mask = mask.unsqueeze(1)
elif mask.ndim == 4 and mask.shape[1] > 1:
mask = mask[:, :1, :, :]
for key in model_A.keys():
if key in model_B:
tA = model_A[key]
tB = model_B[key]
if isinstance(tA, torch.Tensor) and isinstance(tB, torch.Tensor):
# Resize mask if needed
if mask.shape[2:] != tA.shape[2:]:
mask_resized = torch.nn.functional.interpolate(
mask, size=tA.shape[2:], mode='bilinear', align_corners=False
)
else:
mask_resized = mask
blend_factor = mask_resized * merge_ratio
merged[key] = tA * (1 - blend_factor) + tB * blend_factor
else:
merged[key] = tA
else:
merged[key] = model_A[key]
# Save mask into output model
merged["mask_used"] = mask
return (merged,)
NODE_CLASS_MAPPINGS = {
"TensorPrism_WeightedMaskMerge": TensorPrism_WeightedMaskMerge
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_WeightedMaskMerge": "Weighted Mask Merge (Tensor Prism)"
}
+211
View File
@@ -0,0 +1,211 @@
"""
TensorPrism V-Pred/Epsilon Converter Node
========================================
Fixed pure converter between V-Pred and Epsilon prediction types.
Uses proper ComfyUI patching system and accurate conversion logic.
Author: AstrionX
Version: 2.0.0 (Fixed)
License: GPL-3.0
"""
import torch
import math
from typing import Dict, Any, Optional
class TensorPrism_EpsilonVPredConverter:
"""
Pure V-Pred/Epsilon converter with proper ComfyUI integration.
No merging - just conversion between prediction types.
"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("MODEL", {"tooltip": "Model to convert"}),
"input_pred_type": (["epsilon", "v_prediction"], {"default": "epsilon", "tooltip": "Current prediction type of the model"}),
"output_pred_type": (["epsilon", "v_prediction"], {"default": "v_prediction", "tooltip": "Desired output prediction type"}),
"conversion_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01, "tooltip": "Conversion strength (1.0 = full conversion)"}),
}
}
RETURN_TYPES = ("MODEL",)
RETURN_NAMES = ("model",)
FUNCTION = "convert"
CATEGORY = "Tensor_Prism/Convert"
def get_conversion_targets(self, state_dict_keys):
"""
Identify which parameters should be converted based on SDXL architecture.
Focus on the most critical layers for prediction type conversion.
"""
conversion_keys = []
# Critical conversion targets for SDXL
critical_patterns = [
"output_blocks.", # Output blocks are crucial for prediction type
"out.", # Final output layer
"middle_block.", # Middle attention/resnet blocks
]
# Secondary targets (can help but less critical)
secondary_patterns = [
"time_embed.", # Time embedding affects prediction
"input_blocks.11.", # Last input block before middle
"input_blocks.10.", # Second to last input block
]
for key in state_dict_keys:
# Check critical patterns first
for pattern in critical_patterns:
if pattern in key and ("weight" in key or "bias" in key):
conversion_keys.append((key, "critical"))
break
else:
# Check secondary patterns
for pattern in secondary_patterns:
if pattern in key and ("weight" in key or "bias" in key):
conversion_keys.append((key, "secondary"))
break
return conversion_keys
def calculate_conversion_factor(self, from_type: str, to_type: str, priority: str) -> float:
"""
Calculate the conversion factor based on prediction type and layer priority.
The mathematical relationship between epsilon and v-prediction:
v = alpha * epsilon + sigma * x
Where alpha and sigma are noise schedule dependent.
For practical conversion, we use empirically derived factors.
"""
if from_type == to_type:
return 1.0
if from_type == "epsilon" and to_type == "v_prediction":
# Convert epsilon -> v-pred
if priority == "critical":
return 0.7071 # sqrt(0.5) - emphasize the conversion for critical layers
else:
return 0.8660 # sqrt(0.75) - lighter conversion for secondary layers
elif from_type == "v_prediction" and to_type == "epsilon":
# Convert v-pred -> epsilon
if priority == "critical":
return 1.4142 # sqrt(2) - reverse of the above
else:
return 1.1547 # sqrt(4/3) - reverse of secondary conversion
return 1.0
def create_conversion_patches(self, model, from_type: str, to_type: str, strength: float) -> Dict:
"""
Create patches for conversion using ComfyUI's patching system.
"""
if from_type == to_type:
print(f"[TensorPrism] No conversion needed: both types are {from_type}")
return {}
state_dict = model.model.state_dict()
conversion_targets = self.get_conversion_targets(state_dict.keys())
if not conversion_targets:
print(f"[TensorPrism] Warning: No conversion targets found")
return {}
patches = {}
converted_count = 0
print(f"[TensorPrism] Converting {len(conversion_targets)} parameters: {from_type} -> {to_type}")
for key, priority in conversion_targets:
try:
original_param = state_dict[key]
# Skip if parameter is not a tensor or is empty
if not isinstance(original_param, torch.Tensor) or original_param.numel() == 0:
continue
# Calculate conversion factor
base_factor = self.calculate_conversion_factor(from_type, to_type, priority)
# Apply strength scaling
factor = 1.0 + (base_factor - 1.0) * strength
# Create the converted parameter
converted_param = original_param * factor
# Calculate patch (difference from original)
patch_diff = converted_param - original_param
# Only add patch if there's a meaningful difference
if not torch.allclose(patch_diff, torch.zeros_like(patch_diff), atol=1e-8):
patches[key] = (patch_diff,)
converted_count += 1
except Exception as e:
print(f"[TensorPrism] Warning: Failed to convert {key}: {e}")
continue
print(f"[TensorPrism] Created {len(patches)} conversion patches ({converted_count} parameters)")
return patches
def convert(self, model, input_pred_type: str, output_pred_type: str, conversion_strength: float = 1.0):
"""
Main conversion function.
"""
print(f"[TensorPrism] V-Pred/Epsilon Converter")
print(f"[TensorPrism] Input: {input_pred_type} -> Output: {output_pred_type}")
print(f"[TensorPrism] Conversion strength: {conversion_strength}")
try:
# Validate conversion strength
conversion_strength = max(0.0, min(2.0, conversion_strength))
# Create conversion patches
patches = self.create_conversion_patches(
model, input_pred_type, output_pred_type, conversion_strength
)
# Clone the model and apply patches
converted_model = model.clone()
if patches:
converted_model.add_patches(patches, 1.0)
print(f"[TensorPrism] Conversion patches applied successfully")
else:
if input_pred_type != output_pred_type:
print(f"[TensorPrism] Warning: No patches created, returning original model")
else:
print(f"[TensorPrism] No conversion needed, returning original model")
# Update model options if available
try:
if hasattr(converted_model, 'model_options'):
if converted_model.model_options is None:
converted_model.model_options = {}
converted_model.model_options['prediction_type'] = output_pred_type
print(f"[TensorPrism] Updated model prediction type to: {output_pred_type}")
except Exception as e:
print(f"[TensorPrism] Note: Could not update model options: {e}")
print(f"[TensorPrism] Conversion completed successfully")
return (converted_model,)
except Exception as e:
print(f"[TensorPrism] Conversion failed: {e}")
print(f"[TensorPrism] Returning original model")
# Always return a model, never fail completely
return (model,)
# Register the node
NODE_CLASS_MAPPINGS = {
"TensorPrism_EpsilonVPredConverter": TensorPrism_EpsilonVPredConverter
}
NODE_DISPLAY_NAME_MAPPINGS = {
"TensorPrism_EpsilonVPredConverter": "Epsilon/V-Pred Converter (Tensor Prism)"
}
+353
View File
@@ -0,0 +1,353 @@
"""
ComfyUI Tensor Prism Node Pack
===============================
Advanced model merging and enhancement nodes for ComfyUI, providing sophisticated
techniques for blending, enhancing, and manipulating Stable Diffusion models with
GPU-optimized memory management.
Author: AstrionX
Version: 1.2.0
License: GPL-3.0
Repository: https://github.com/AstrionX/ComfyUI-Tensor-Prism-Node-Pack
Features:
- Advanced model merging with multiple interpolation methods
- Spectral frequency-domain merging
- Granular SDXL block control
- Sophisticated masking system
- GPU-optimized memory management
- Cross-platform compatibility (CUDA/MPS/CPU)
"""
import os
import sys
import importlib.util
from pathlib import Path
# Add current directory to path for imports
current_dir = Path(__file__).parent
sys.path.insert(0, str(current_dir))
try:
from .TensorPrism_MainMerge import TensorPrism_MainMerge
except ImportError:
TensorPrism_MainMerge = None
try:
from .TensorPrism_SDXLBlockMerge import SDXLBlockMergeTensorPrism
except ImportError:
SDXLBlockMergeTensorPrism = None
try:
from .TensorPrism_SDXLAdvancedBlockmerge import SDXLAdvancedBlockMergeTensorPrism
except ImportError:
SDXLAdvancedBlockMergeTensorPrism = None
try:
from .TensorPrism_ModelMaskGenerator import TensorPrism_ModelMaskGenerator, TensorPrism_WeightedMaskMerge
except ImportError:
TensorPrism_ModelMaskGenerator = None
TensorPrism_WeightedMaskMerge = None
try:
from .TensorPrism_ModelKeyFilter import TensorPrism_ModelKeyFilter
except ImportError:
TensorPrism_ModelKeyFilter = None
try:
from .TensorPrism_ModelMaskBlender import TensorPrism_ModelMaskBlender
except ImportError:
TensorPrism_ModelMaskBlender = None
try:
from .TensorPrism_ModelWeightModifier import TensorPrism_ModelWeightModifier
except ImportError:
TensorPrism_ModelWeightModifier = None
try:
from .TensorPrism_WeightedTensorMerge import TensorPrism_WeightedMaskMerge as TensorPrism_WeightedTensorMerge_Class
except ImportError:
TensorPrism_WeightedTensorMerge_Class = None
try:
from .TensorPrism_Enhancer import ModelEnhancerTensorPrism
except ImportError:
ModelEnhancerTensorPrism = None
try:
from .TensorPrism_LayeredBlend import TensorPrism_LayeredBlend
except ImportError:
TensorPrism_LayeredBlend = None
try:
from .TensorPrism_Prism import TensorPrism_FastPrism
except ImportError:
TensorPrism_FastPrism = None
try:
from .TensorPrism_vpredepsilonconverter import TensorPrism_EpsilonVPredConverter
except ImportError:
TensorPrism_EpsilonVPredConverter = None
try:
from .TensorPrism_CheckpointReroute_Notes import TensorPrism_CheckpointReroute_Notes
except ImportError:
TensorPrism_CheckpointReroute_Notes = None
try:
from .TensorPrism_AdvancedClipMerge import AdvancedCLIPMerge
except ImportError:
AdvancedCLIPMerge = None
# Version info
__version__ = "1.2.0"
__author__ = "AstrionX"
__description__ = "Advanced model merging and enhancement nodes for ComfyUI"
NODE_CLASS_MAPPINGS = {}
# Core merging nodes
if TensorPrism_MainMerge:
NODE_CLASS_MAPPINGS["TensorPrism_MainMerge"] = TensorPrism_MainMerge
# SDXL specific nodes
if SDXLBlockMergeTensorPrism:
NODE_CLASS_MAPPINGS["SDXL Block Merge (Tensor Prism)"] = SDXLBlockMergeTensorPrism
if SDXLAdvancedBlockMergeTensorPrism:
NODE_CLASS_MAPPINGS["SDXLAdvancedBlockMergeTensorPrism"] = SDXLAdvancedBlockMergeTensorPrism
# Masking and filtering nodes
if TensorPrism_ModelMaskGenerator:
NODE_CLASS_MAPPINGS["TensorPrism_ModelMaskGenerator"] = TensorPrism_ModelMaskGenerator
if TensorPrism_WeightedMaskMerge:
NODE_CLASS_MAPPINGS["TensorPrism_WeightedMaskMerge"] = TensorPrism_WeightedMaskMerge
if TensorPrism_ModelKeyFilter:
NODE_CLASS_MAPPINGS["TensorPrism_ModelKeyFilter"] = TensorPrism_ModelKeyFilter
if TensorPrism_ModelMaskBlender:
NODE_CLASS_MAPPINGS["TensorPrism_ModelMaskBlender"] = TensorPrism_ModelMaskBlender
# Enhancement and transformation nodes
if TensorPrism_ModelWeightModifier:
NODE_CLASS_MAPPINGS["TensorPrism_ModelWeightModifier"] = TensorPrism_ModelWeightModifier
if TensorPrism_WeightedTensorMerge_Class:
NODE_CLASS_MAPPINGS["TensorPrism_WeightedTensorMerge"] = TensorPrism_WeightedTensorMerge_Class
if ModelEnhancerTensorPrism:
NODE_CLASS_MAPPINGS["ModelEnhancerTensorPrism"] = ModelEnhancerTensorPrism
if TensorPrism_LayeredBlend:
NODE_CLASS_MAPPINGS["TensorPrism_LayeredBlend"] = TensorPrism_LayeredBlend
if TensorPrism_FastPrism:
NODE_CLASS_MAPPINGS["TensorPrism_Prism"] = TensorPrism_FastPrism
# CLIP merging nodes
if AdvancedCLIPMerge:
NODE_CLASS_MAPPINGS["AdvancedCLIPMerge"] = AdvancedCLIPMerge
# Epsilon/V-Pred merge nodes
if TensorPrism_EpsilonVPredConverter:
NODE_CLASS_MAPPINGS["TensorPrism_EpsilonVPredConverter"] = TensorPrism_EpsilonVPredConverter
if TensorPrism_CheckpointReroute_Notes:
NODE_CLASS_MAPPINGS["TensorPrism_CheckpointReroute_Notes"] = TensorPrism_CheckpointReroute_Notes
NODE_DISPLAY_NAME_MAPPINGS = {}
# Core merging nodes
if "TensorPrism_MainMerge" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_MainMerge"] = "Main Merge (Tensor Prism)"
# SDXL specific nodes
if "SDXL Block Merge (Tensor Prism)" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["SDXL Block Merge (Tensor Prism)"] = "SDXL Block Merge (Tensor Prism)"
if "SDXLAdvancedBlockMergeTensorPrism" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["SDXLAdvancedBlockMergeTensorPrism"] = "SDXL Advanced Block Merge (Tensor Prism)"
# Masking and filtering nodes
if "TensorPrism_ModelMaskGenerator" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelMaskGenerator"] = "Model Mask Generator (Tensor Prism)"
if "TensorPrism_WeightedMaskMerge" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_WeightedMaskMerge"] = "Weighted Mask Merge (Tensor Prism)"
if "TensorPrism_ModelKeyFilter" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelKeyFilter"] = "Model Key Filter (Tensor Prism)"
if "TensorPrism_ModelMaskBlender" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelMaskBlender"] = "Mask Blender (Tensor Prism)"
# Enhancement and transformation nodes
if "TensorPrism_ModelWeightModifier" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_ModelWeightModifier"] = "Model Weight Modifier (Tensor Prism)"
if "TensorPrism_WeightedTensorMerge" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_WeightedTensorMerge"] = "Weighted Tensor Merge (Tensor Prism)"
if "ModelEnhancerTensorPrism" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["ModelEnhancerTensorPrism"] = "Model Enhancer (Tensor Prism)"
if "TensorPrism_LayeredBlend" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_LayeredBlend"] = "Layered Blend (Tensor Prism)"
if "TensorPrism_Prism" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_Prism"] = "Prism (Tensor Prism)"
# CLIP merging nodes
if "AdvancedCLIPMerge" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["AdvancedCLIPMerge"] = "Advanced CLIP Merge (Tensor Prism)"
# Epsilon/V-Pred merge nodes
if "TensorPrism_EpsilonVPredConverter" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_EpsilonVPredConverter"] = "Epsilon/V-Pred Converter (Tensor Prism)"
if "TensorPrism_CheckpointReroute_Notes" in NODE_CLASS_MAPPINGS:
NODE_DISPLAY_NAME_MAPPINGS["TensorPrism_CheckpointReroute_Notes"] = "Checkpoint Reroute + Notes (Tensor Prism)"
NODE_CATEGORIES = {}
for node_key in NODE_CLASS_MAPPINGS.keys():
if "MainMerge" in node_key or "LayeredBlend" in node_key:
NODE_CATEGORIES[node_key] = "Tensor Prism/Core"
elif "SDXL" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Merge"
elif "Mask" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Mask"
elif "Weight" in node_key or "Enhancer" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Transform"
elif "EpsilonVPred" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Merge"
elif "CLIP" in node_key:
NODE_CATEGORIES[node_key] = "Tensor_Prism/CLIP"
else:
NODE_CATEGORIES[node_key] = "Tensor_Prism/Core"
def check_dependencies():
"""Check if required dependencies are available."""
required_packages = [
("torch", "PyTorch >= 1.12.0"),
("numpy", "NumPy >= 1.21.0"),
("psutil", "psutil >= 5.8.0"),
]
missing_packages = []
for package_name, description in required_packages:
try:
importlib.import_module(package_name)
except ImportError:
missing_packages.append(description)
if missing_packages:
print(f"[TensorPrism] Warning: Missing dependencies:")
for pkg in missing_packages:
print(f" - {pkg}")
print("[TensorPrism] Some features may not work properly.")
return len(missing_packages) == 0
def get_system_info():
"""Get system information for optimization."""
try:
import torch
import psutil
info = {
"torch_version": torch.__version__,
"cuda_available": torch.cuda.is_available(),
"mps_available": hasattr(torch.backends, 'mps') and torch.backends.mps.is_available(),
"cpu_count": psutil.cpu_count(),
"total_ram": round(psutil.virtual_memory().total / (1024**3), 1),
}
if info["cuda_available"]:
info["cuda_device_count"] = torch.cuda.device_count()
info["cuda_memory"] = round(torch.cuda.get_device_properties(0).total_memory / (1024**3), 1)
return info
except Exception as e:
print(f"[TensorPrism] Could not get system info: {e}")
return {}
def print_welcome_message():
"""Print welcome message with system information."""
print("\n" + "="*60)
print("🎭 TensorPrism Node Pack Loaded")
print("="*60)
print(f"Version: {__version__}")
print(f"Author: {__author__}")
print(f"Nodes: {len(NODE_CLASS_MAPPINGS)}")
# System info
sys_info = get_system_info()
if sys_info:
print(f"\n📊 System Information:")
print(f" PyTorch: {sys_info.get('torch_version', 'Unknown')}")
print(f" CPU Cores: {sys_info.get('cpu_count', 'Unknown')}")
print(f" RAM: {sys_info.get('total_ram', 'Unknown')}GB")
if sys_info.get("cuda_available"):
print(f" CUDA: Available ({sys_info.get('cuda_device_count', 0)} device(s))")
print(f" GPU Memory: {sys_info.get('cuda_memory', 'Unknown')}GB")
elif sys_info.get("mps_available"):
print(f" MPS: Available (Apple Silicon)")
else:
print(f" GPU: CPU fallback mode")
print(f"\n🚀 Available Nodes:")
for category in set(NODE_CATEGORIES.values()):
print(f"\n {category}:")
for node_class, node_category in NODE_CATEGORIES.items():
if node_category == category:
display_name = NODE_DISPLAY_NAME_MAPPINGS.get(node_class, node_class)
print(f" • {display_name}")
print(f"\n💡 Memory Management:")
if sys_info.get("cuda_memory", 0) >= 24:
print(" Recommended: Default settings (24GB+ GPU)")
elif sys_info.get("cuda_memory", 0) >= 12:
print(" Recommended: Memory limit 8GB, auto precision")
elif sys_info.get("cuda_memory", 0) >= 8:
print(" Recommended: Memory limit 6GB, CPU fallback for large merges")
else:
print(" Recommended: CPU processing for best stability")
print("="*60 + "\n")
# Initialize the package
def __init_package():
"""Initialize the package and perform startup checks."""
try:
# Check dependencies
deps_ok = check_dependencies()
# Print welcome message
print_welcome_message()
if not deps_ok:
print("[TensorPrism] Warning: Some dependencies are missing. Please install them for full functionality.")
print("[TensorPrism] Package initialized successfully!")
except Exception as e:
print(f"[TensorPrism] Error during initialization: {e}")
print("[TensorPrism] Package may not function correctly.")
# Run initialization
__init_package()
# Export for ComfyUI
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
"__version__",
"__author__",
"__description__"
]
+45
View File
@@ -0,0 +1,45 @@
[project]
name = "tensorprism-comfyui"
version = "1.1.0"
description = "Advanced model merging and enhancement nodes for ComfyUI with GPU-optimized memory management"
readme = "README.md"
authors = [
{name = "AstrionX", email = "astrionx@example.com"}
]
license = {file = "LICENSE"}
requires-python = ">=3.8"
keywords = ["comfyui", "stable-diffusion", "model-merging", "ai", "deep-learning", "pytorch"]
classifiers = [
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
"License :: OSI Approved :: GNU General Public License v3 (GPLv3)",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.8",
"Programming Language :: Python :: 3.9",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
]
dependencies = [
"torch>=1.12.0",
"numpy>=1.21.0",
"psutil>=5.8.0",
]
[project.urls]
Homepage = "https://github.com/AstrionX/ComfyUI-Tensor-Prism-Node-Pack"
Repository = "https://github.com/AstrionX/ComfyUI-Tensor-Prism-Node-Pack"
Issues = "https://github.com/AstrionX/ComfyUI-Tensor-Prism-Node-Pack/issues"
Documentation = "https://github.com/AstrionX/ComfyUI-Tensor-Prism-Node-Pack/blob/main/README.md"
[build-system]
requires = ["setuptools>=61.0", "wheel"]
build-backend = "setuptools.build_meta"
[tool.setuptools.packages.find]
where = ["."]
include = ["*"]
exclude = ["tests*", "__pycache__*", "*.pyc"]
[tool.setuptools.package-data]
"*" = ["*.md", "*.txt", "*.json", "*.yaml", "*.yml"]