802 lines
30 KiB
Python
802 lines
30 KiB
Python
"""
|
|
TensorPrism Masking System
|
|
===========================
|
|
|
|
Comprehensive masking system for selective model merging.
|
|
Includes mask generation, filtering, blending, and application.
|
|
|
|
Author: Arctenox
|
|
Version: 1.0.0
|
|
License: GPL-3.0
|
|
"""
|
|
|
|
import torch
|
|
import numpy as np
|
|
import re
|
|
import gc
|
|
import psutil
|
|
from typing import Dict, List, Tuple, Optional
|
|
from collections import defaultdict
|
|
|
|
# ==================== UTILITY FUNCTIONS ====================
|
|
|
|
def is_unet_key(key: str) -> bool:
|
|
"""Check if key belongs to UNet"""
|
|
key_lower = key.lower()
|
|
return 'unet' in key_lower or 'model.diffusion_model' in key_lower
|
|
|
|
def is_vae_key(key: str) -> bool:
|
|
"""Check if key belongs to VAE"""
|
|
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:
|
|
"""Check if key belongs to text encoder"""
|
|
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:
|
|
"""Identify UNet component type"""
|
|
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 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
|
|
|
|
def estimate_dict_memory_gb(dict_size: int) -> float:
|
|
"""Estimate memory usage of a dictionary with float values in GB"""
|
|
bytes_per_entry = 100
|
|
return (dict_size * bytes_per_entry) / (1024**3)
|
|
|
|
|
|
# ==================== MASK GENERATOR ====================
|
|
|
|
class TensorPrism_ModelMaskGenerator:
|
|
"""Generate masks with various strategies"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
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
|
|
}),
|
|
"reference_model": ("MODEL",),
|
|
},
|
|
"optional": {
|
|
"layer_start": ("INT", {
|
|
"default": 0, "min": 0, "max": 50, "step": 1
|
|
}),
|
|
"layer_end": ("INT", {
|
|
"default": -1, "min": -1, "max": 50, "step": 1
|
|
}),
|
|
"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
|
|
}),
|
|
"custom_pattern": ("STRING", {
|
|
"default": "attn,mlp.fc1"
|
|
}),
|
|
"falloff": ("FLOAT", {
|
|
"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.05
|
|
}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("MASK",)
|
|
RETURN_NAMES = ("mask",)
|
|
FUNCTION = "generate_mask"
|
|
CATEGORY = "Tensor_Prism/Mask"
|
|
|
|
def analyze_model_structure(self, state_dict):
|
|
"""Analyze model structure"""
|
|
layer_info = {
|
|
"layers": {},
|
|
"total_layers": 0,
|
|
"layer_names": list(state_dict.keys())
|
|
}
|
|
|
|
layer_patterns = [
|
|
r"layers\.(\d+)\.",
|
|
r"blocks\.(\d+)\.",
|
|
r"h\.(\d+)\.",
|
|
r"layer\.(\d+)\.",
|
|
r"encoder\.layer\.(\d+)\.",
|
|
r"decoder\.layer\.(\d+)\.",
|
|
]
|
|
|
|
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
|
|
|
|
def get_layer_number(self, param_name):
|
|
"""Extract layer number"""
|
|
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
|
|
|
|
def create_layer_range_mask(self, layer_info, start, end, intensity):
|
|
"""Create mask for layer range"""
|
|
mask = {}
|
|
if end == -1:
|
|
end = layer_info["total_layers"]
|
|
|
|
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:
|
|
mask[name] = intensity
|
|
else:
|
|
mask[name] = 0.0
|
|
return mask
|
|
|
|
def create_block_mask(self, layer_info, start, end, intensity):
|
|
"""Create mask 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:
|
|
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"""
|
|
mask = {}
|
|
for name in layer_info["layer_names"]:
|
|
should_mask = any(pattern.lower() in name.lower()
|
|
for pattern in component_patterns)
|
|
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 = any(pattern.lower() in name.lower()
|
|
for pattern in patterns)
|
|
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 name in layer_info["layer_names"]:
|
|
mask[name] = intensity if np.random.random() > sparsity else 0.0
|
|
return mask
|
|
|
|
def create_depth_gradient_mask(self, layer_info, direction, intensity, falloff):
|
|
"""Create gradient mask based on 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:
|
|
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)
|
|
|
|
if falloff > 0:
|
|
mask_value = np.power(mask_value, 1.0 / max(falloff, 0.01))
|
|
|
|
mask[name] = mask_value * intensity
|
|
else:
|
|
mask[name] = intensity * 0.1
|
|
return mask
|
|
|
|
def create_uniform_mask(self, layer_info, intensity):
|
|
"""Create uniform mask"""
|
|
return {name: intensity for name in layer_info["layer_names"]}
|
|
|
|
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 mask with model structure"""
|
|
|
|
state_dict = reference_model.model.state_dict()
|
|
layer_info = self.analyze_model_structure(state_dict)
|
|
|
|
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)
|
|
|
|
return ({
|
|
"mask_dict": mask,
|
|
"mask_type": mask_type,
|
|
"intensity": intensity,
|
|
"layer_info": layer_info
|
|
},)
|
|
|
|
|
|
# ==================== KEY FILTER ====================
|
|
|
|
class TensorPrism_ModelKeyFilter:
|
|
"""Memory-efficient model key filter"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
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
|
|
}),
|
|
"target_value": ("FLOAT", {
|
|
"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01,
|
|
"round": 0.001
|
|
}),
|
|
"memory_limit_gb": ("FLOAT", {
|
|
"default": 2.0, "min": 0.5, "max": 16.0, "step": 0.1,
|
|
"round": 0.1
|
|
}),
|
|
},
|
|
"optional": {
|
|
"custom_pattern": ("STRING", {
|
|
"default": "attn,resnets",
|
|
"multiline": True
|
|
}),
|
|
"exact_match_custom": ("BOOLEAN", {"default": False}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("MASK",)
|
|
RETURN_NAMES = ("filtered_mask",)
|
|
FUNCTION = "filter_keys_to_mask"
|
|
CATEGORY = "Tensor_Prism/Mask"
|
|
|
|
def create_key_batches(self, all_keys: List[str], memory_limit_gb: float) -> List[List[str]]:
|
|
"""Create batches of keys within memory limit"""
|
|
batches = []
|
|
current_batch = []
|
|
max_keys_per_batch = max(1000, int((memory_limit_gb * 1024**3) / 200))
|
|
|
|
for key in all_keys:
|
|
current_batch.append(key)
|
|
if len(current_batch) >= max_keys_per_batch:
|
|
batches.append(current_batch)
|
|
current_batch = []
|
|
|
|
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 key matches target 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
|
|
|
|
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 batch of keys"""
|
|
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: {target_components}")
|
|
|
|
used_memory, available_memory = get_memory_info()
|
|
print(f" Memory - Used: {used_memory:.2f}GB, Available: {available_memory:.2f}GB")
|
|
|
|
state_dict = model.model.state_dict()
|
|
all_keys = list(state_dict.keys())
|
|
print(f" Total keys: {len(all_keys)}")
|
|
|
|
patterns = [p.strip() for p in custom_pattern.split(',') if p.strip()]
|
|
batches = self.create_key_batches(all_keys, memory_limit_gb)
|
|
print(f" Created {len(batches)} batches")
|
|
|
|
mask_dict = {}
|
|
processed_keys = 0
|
|
|
|
for i, batch_keys in enumerate(batches):
|
|
batch_results = self.process_key_batch(
|
|
batch_keys, target_components, filter_mode,
|
|
default_value, target_value, patterns, exact_match_custom
|
|
)
|
|
mask_dict.update(batch_results)
|
|
processed_keys += len(batch_keys)
|
|
|
|
gc.collect()
|
|
|
|
if i % 5 == 0:
|
|
progress = (processed_keys / len(all_keys)) * 100
|
|
print(f" Progress: {progress:.1f}%")
|
|
|
|
# Analyze structure
|
|
layer_info = {"layer_names": all_keys, "total_layers": 0, "layers": {}}
|
|
|
|
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, _ = get_memory_info()
|
|
print(f" Final memory: {final_memory:.2f}GB")
|
|
print(f" Mask intensity: {mask['intensity']:.4f}")
|
|
print(f"--- Filter completed ---\n")
|
|
|
|
return (mask,)
|
|
|
|
|
|
# ==================== MASK BLENDER ====================
|
|
|
|
class TensorPrism_ModelMaskBlender:
|
|
"""Blend two masks together"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
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
|
|
}),
|
|
},
|
|
"optional": {
|
|
"blend_strength": ("FLOAT", {
|
|
"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01
|
|
}),
|
|
"clip_output": ("BOOLEAN", {"default": True}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("MASK",)
|
|
RETURN_NAMES = ("combined_mask",)
|
|
FUNCTION = "blend_masks"
|
|
CATEGORY = "Tensor_Prism/Mask"
|
|
|
|
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--- Mask Blender ---")
|
|
print(f" Mode: {blend_mode}, Strength: {blend_strength}")
|
|
|
|
# Convert to numpy
|
|
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)
|
|
|
|
# Resize if needed
|
|
if mask_A_np.shape != mask_B_np.shape:
|
|
from scipy import ndimage
|
|
if len(mask_A_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:
|
|
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]))
|
|
|
|
# Blend
|
|
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()
|
|
gc.collect()
|
|
|
|
print(f" Result shape: {result_tensor.shape}")
|
|
print(f"--- Blender completed ---\n")
|
|
|
|
return (result_tensor,)
|
|
|
|
|
|
# ==================== WEIGHTED MASK MERGE ====================
|
|
|
|
class TensorPrism_WeightedMaskMerge:
|
|
"""Apply mask to merge two models"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
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
|
|
}),
|
|
}
|
|
}
|
|
|
|
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 mask with device safety"""
|
|
|
|
merged_model = model_A.clone()
|
|
state_dict_A = model_A.model.state_dict()
|
|
state_dict_B = model_B.model.state_dict()
|
|
|
|
mask_dict = mask["mask_dict"]
|
|
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:
|
|
device = weight_A.device
|
|
weight_B = weight_B.to(device)
|
|
|
|
mask_value = mask_dict[key] * merge_ratio
|
|
merged_weight = weight_A * (1 - mask_value) + weight_B * mask_value
|
|
merged_weight = merged_weight.to(device)
|
|
|
|
if mask_value > 0:
|
|
patches[key] = ((merged_weight - weight_A).cpu(),)
|
|
|
|
if patches:
|
|
merged_model.add_patches(patches, 1.0)
|
|
|
|
return (merged_model,)
|
|
|
|
|
|
# ==================== ADVANCED WEIGHTED MERGE ====================
|
|
|
|
class TensorPrism_WeightedMaskMergeAdvanced:
|
|
"""Advanced merge with sophisticated blending"""
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"model_A": ("MODEL",),
|
|
"model_B": ("MODEL",),
|
|
"mask": ("MASK",),
|
|
"merge_ratio": ("FLOAT", {
|
|
"default": 0.5, "min": 0.0, "max": 2.0, "step": 0.01
|
|
}),
|
|
"blend_mode": ([
|
|
"linear", "sigmoid", "cosine",
|
|
"exponential", "logarithmic", "smoothstep"
|
|
], {"default": "linear"}),
|
|
},
|
|
"optional": {
|
|
"curve_power": ("FLOAT", {
|
|
"default": 1.0, "min": 0.1, "max": 5.0, "step": 0.1
|
|
}),
|
|
"preserve_extremes": ("BOOLEAN", {"default": False}),
|
|
"noise_injection": ("FLOAT", {
|
|
"default": 0.0, "min": 0.0, "max": 0.1, "step": 0.001
|
|
}),
|
|
"layer_scaling": ([
|
|
"uniform", "depth_progressive",
|
|
"shallow_bias", "deep_bias"
|
|
], {"default": "uniform"}),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("MODEL",)
|
|
RETURN_NAMES = ("merged_model",)
|
|
FUNCTION = "merge_models"
|
|
CATEGORY = "Tensor_Prism/Mask"
|
|
|
|
def apply_blend_curve(self, value, mode, power):
|
|
"""Apply blending curves"""
|
|
if mode == "linear":
|
|
return value
|
|
elif mode == "sigmoid":
|
|
return 1.0 / (1.0 + np.exp(-power * (value - 0.5) * 10))
|
|
elif mode == "cosine":
|
|
return (1.0 - np.cos(value * np.pi)) * 0.5
|
|
elif mode == "exponential":
|
|
if value < 0.5:
|
|
return 0.5 * np.power(2.0 * value, power)
|
|
else:
|
|
return 1.0 - 0.5 * np.power(2.0 * (1.0 - value), power)
|
|
elif mode == "logarithmic":
|
|
epsilon = 1e-6
|
|
return np.log(value + epsilon) / np.log(1.0 + epsilon)
|
|
elif mode == "smoothstep":
|
|
value = np.clip(value, 0.0, 1.0)
|
|
return value * value * (3.0 - 2.0 * value)
|
|
return value
|
|
|
|
def get_layer_scale(self, layer_num, total_layers, scaling_mode):
|
|
"""Calculate layer-specific scaling"""
|
|
if total_layers <= 1:
|
|
return 1.0
|
|
|
|
position = layer_num / (total_layers - 1)
|
|
|
|
if scaling_mode == "uniform":
|
|
return 1.0
|
|
elif scaling_mode == "depth_progressive":
|
|
return 0.5 + 0.5 * position
|
|
elif scaling_mode == "shallow_bias":
|
|
return 1.5 - 0.5 * position
|
|
elif scaling_mode == "deep_bias":
|
|
return 0.5 + position
|
|
return 1.0
|
|
|
|
def extract_layer_number(self, key):
|
|
"""Extract layer number from key"""
|
|
patterns = [
|
|
r"layers\.(\d+)\.", r"blocks\.(\d+)\.", r"h\.(\d+)\.",
|
|
r"layer\.(\d+)\.", r"encoder\.layer\.(\d+)\.",
|
|
r"decoder\.layer\.(\d+)\.",
|
|
]
|
|
|
|
for pattern in patterns:
|
|
match = re.search(pattern, key)
|
|
if match:
|
|
return int(match.group(1))
|
|
return None
|
|
|
|
def merge_models(self, model_A, model_B, mask, merge_ratio, blend_mode,
|
|
curve_power=1.0, preserve_extremes=False, noise_injection=0.0,
|
|
layer_scaling="uniform"):
|
|
"""Advanced merge with blending"""
|
|
|
|
print(f"\n--- Advanced Weighted Merge ---")
|
|
print(f" Mode: {blend_mode}, Ratio: {merge_ratio}")
|
|
|
|
merged_model = model_A.clone()
|
|
state_dict_A = model_A.model.state_dict()
|
|
state_dict_B = model_B.model.state_dict()
|
|
|
|
if isinstance(mask, dict) and "mask_dict" in mask:
|
|
mask_dict = mask["mask_dict"]
|
|
else:
|
|
mask_dict = {key: 1.0 for key in state_dict_A.keys()}
|
|
|
|
# Analyze layers
|
|
layer_numbers = {}
|
|
max_layer = 0
|
|
for key in state_dict_A.keys():
|
|
layer_num = self.extract_layer_number(key)
|
|
if layer_num is not None:
|
|
layer_numbers[key] = layer_num
|
|
max_layer = max(max_layer, layer_num)
|
|
|
|
total_layers = max_layer + 1 if max_layer > 0 else 1
|
|
|
|
patches = {}
|
|
processed_count = 0
|
|
|
|
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:
|
|
device = weight_A.device
|
|
weight_B = weight_B.to(device)
|
|
|
|
# Get base mask value
|
|
mask_value = float(mask_dict[key])
|
|
|
|
# Apply blending curve
|
|
blend_factor = self.apply_blend_curve(mask_value, blend_mode, curve_power)
|
|
|
|
# Apply layer scaling
|
|
if key in layer_numbers:
|
|
layer_scale = self.get_layer_scale(
|
|
layer_numbers[key], total_layers, layer_scaling
|
|
)
|
|
blend_factor *= layer_scale
|
|
|
|
# Apply global merge ratio
|
|
blend_factor *= merge_ratio
|
|
|
|
# Preserve extremes if requested
|
|
if preserve_extremes:
|
|
if mask_value < 0.05:
|
|
blend_factor = 0.0
|
|
elif mask_value > 0.95:
|
|
blend_factor = merge_ratio
|
|
|
|
# Clip to valid range
|
|
blend_factor = np.clip(blend_factor, 0.0, 2.0)
|
|
|
|
# Perform merge - all on same device
|
|
merged_weight = weight_A * (1.0 - blend_factor) + weight_B * blend_factor
|
|
|
|
# Add noise if requested
|
|
if noise_injection > 0:
|
|
noise = torch.randn_like(merged_weight, device=device)
|
|
noise *= noise_injection * merged_weight.abs().mean()
|
|
merged_weight = merged_weight + noise
|
|
|
|
merged_weight = merged_weight.to(device)
|
|
|
|
# Create patch if there's change
|
|
if blend_factor > 0.001:
|
|
patches[key] = ((merged_weight - weight_A).cpu(),)
|
|
processed_count += 1
|
|
|
|
if patches:
|
|
merged_model.add_patches(patches, 1.0)
|
|
|
|
print(f" Processed {processed_count} tensors")
|
|
print(f"--- Advanced merge completed ---\n")
|
|
|
|
return (merged_model,)
|
|
|
|
|
|
# ==================== NODE REGISTRATION ====================
|
|
|
|
NODE_CLASS_MAPPINGS = {
|
|
"TensorPrism_ModelMaskGenerator": TensorPrism_ModelMaskGenerator,
|
|
"TensorPrism_ModelKeyFilter": TensorPrism_ModelKeyFilter,
|
|
"TensorPrism_ModelMaskBlender": TensorPrism_ModelMaskBlender,
|
|
"TensorPrism_WeightedMaskMerge": TensorPrism_WeightedMaskMerge,
|
|
"TensorPrism_WeightedMaskMergeAdvanced": TensorPrism_WeightedMaskMergeAdvanced,
|
|
}
|
|
|
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
"TensorPrism_ModelMaskGenerator": "Model Mask Generator (Tensor Prism)",
|
|
"TensorPrism_ModelKeyFilter": "Model Key Filter (Tensor Prism)",
|
|
"TensorPrism_ModelMaskBlender": "Mask Blender (Tensor Prism)",
|
|
"TensorPrism_WeightedMaskMerge": "Weighted Mask Merge (Tensor Prism)",
|
|
"TensorPrism_WeightedMaskMergeAdvanced": "Advanced Weighted Mask Merge (Tensor Prism)",
|
|
} |