Files
AstrionX-ComfyUI-Tensor-Pri…/TensorPrism_ModelMaskGenerator.py
T
2025-09-15 02:42:31 -04:00

346 lines
13 KiB
Python

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)"
}