From a9a528942f34d639cdac67e5da2d42d5cc69ac22 Mon Sep 17 00:00:00 2001 From: David Piazza <50223477+DavidPiazza@users.noreply.github.com> Date: Thu, 2 Oct 2025 11:39:45 -0400 Subject: [PATCH] Implement network_bending entrypoint and add InvertedPruning node; enhance audio nodes with error handling and normalization checks; remove obsolete test files. --- __init__.py | 44 ++ src/network_bending/audio_nodes/__init__.py | 86 +-- .../audio_nodes/audio_latent_nodes.py | 16 +- .../audio_style_transfer_workflow.json | 219 ------- src/network_bending/js/network_bending.js | 5 +- src/network_bending/nodes.py | 604 +++++++++++++++++- tests/__init__.py | 1 - tests/conftest.py | 6 - tests/pytest.ini | 4 - tests/test_network_bending.py | 214 ------- 10 files changed, 686 insertions(+), 513 deletions(-) delete mode 100644 src/network_bending/audio_workflows/audio_style_transfer_workflow.json delete mode 100644 tests/__init__.py delete mode 100644 tests/conftest.py delete mode 100644 tests/pytest.ini delete mode 100644 tests/test_network_bending.py diff --git a/__init__.py b/__init__.py index 6277310..2cbb48b 100644 --- a/__init__.py +++ b/__init__.py @@ -1,3 +1,47 @@ +""" +ComfyUI entrypoint for the network_bending custom node pack. + +Exports NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS and WEB_DIRECTORY. +Handles adding the local src/ path so the packaged code under src/network_bending +is importable without installation. +""" + +import os +import sys + +# Ensure the local src directory is importable when running inside ComfyUI +_HERE = os.path.dirname(__file__) +_SRC_DIR = os.path.join(_HERE, "src") +if os.path.isdir(_SRC_DIR) and _SRC_DIR not in sys.path: + sys.path.insert(0, _SRC_DIR) + +# Import node mappings from the packaged module +try: + from network_bending.nodes import ( # type: ignore + NODE_CLASS_MAPPINGS as _NODE_CLASS_MAPPINGS, + NODE_DISPLAY_NAME_MAPPINGS as _NODE_DISPLAY_NAME_MAPPINGS, + ) +except Exception as e: # pragma: no cover - surface helpful error in UI + # Provide clearer error if import fails (e.g., missing deps) + raise RuntimeError( + f"Failed to import network_bending nodes. Error: {e}. " + "Ensure dependencies are installed and the 'src' folder exists." + ) + + +# Re-export for ComfyUI +NODE_CLASS_MAPPINGS = _NODE_CLASS_MAPPINGS +NODE_DISPLAY_NAME_MAPPINGS = _NODE_DISPLAY_NAME_MAPPINGS + +# Expose web directory for frontend helpers +WEB_DIRECTORY = "./src/network_bending/js" + +__all__ = [ + "NODE_CLASS_MAPPINGS", + "NODE_DISPLAY_NAME_MAPPINGS", + "WEB_DIRECTORY", +] + """Top-level package for network_bending.""" __all__ = [ diff --git a/src/network_bending/audio_nodes/__init__.py b/src/network_bending/audio_nodes/__init__.py index 196f4a8..7d30de1 100644 --- a/src/network_bending/audio_nodes/__init__.py +++ b/src/network_bending/audio_nodes/__init__.py @@ -1,44 +1,56 @@ """ -Audio conditioning nodes for Stable Audio in ComfyUI +Audio conditioning nodes for Stable Audio in ComfyUI. + +This subpackage may have optional dependencies (e.g., torchaudio, librosa). +If those are not installed, we gracefully disable audio nodes rather than +failing the entire custom node pack. """ -from .audio_latent_nodes import ( - AudioVAEEncode, - AudioVAEDecode, - AudioLatentInterpolate, - AudioLatentBlend, - AudioFeatureExtractor, - AudioLatentManipulator, -) +from typing import Dict -from .audio_style_transfer import ( - AudioStyleTransfer, - AudioLatentGuidance, - AudioReferenceEncoder -) +try: + from .audio_latent_nodes import ( # type: ignore + AudioVAEEncode, + AudioVAEDecode, + AudioLatentInterpolate, + AudioLatentBlend, + AudioFeatureExtractor, + AudioLatentManipulator, + ) -NODE_CLASS_MAPPINGS = { - "AudioVAEEncode": AudioVAEEncode, - "AudioVAEDecode": AudioVAEDecode, - "AudioLatentInterpolate": AudioLatentInterpolate, - "AudioLatentBlend": AudioLatentBlend, - "AudioFeatureExtractor": AudioFeatureExtractor, - "AudioLatentManipulator": AudioLatentManipulator, - "AudioStyleTransfer": AudioStyleTransfer, - "AudioLatentGuidance": AudioLatentGuidance, - "AudioReferenceEncoder": AudioReferenceEncoder, -} + from .audio_style_transfer import ( # type: ignore + AudioStyleTransfer, + AudioLatentGuidance, + AudioReferenceEncoder, + ) -NODE_DISPLAY_NAME_MAPPINGS = { - "AudioVAEEncode": "Audio VAE Encode", - "AudioVAEDecode": "Audio VAE Decode", - "AudioLatentInterpolate": "Audio Latent Interpolate", - "AudioLatentBlend": "Audio Latent Blend", - "AudioFeatureExtractor": "Audio Feature Extractor", - "AudioLatentManipulator": "Audio Latent Manipulator", - "AudioStyleTransfer": "Audio Style Transfer", - "AudioLatentGuidance": "Audio Latent Guidance", - "AudioReferenceEncoder": "Audio Reference Encoder", -} + NODE_CLASS_MAPPINGS: Dict[str, object] = { + "AudioVAEEncode": AudioVAEEncode, + "AudioVAEDecode": AudioVAEDecode, + "AudioLatentInterpolate": AudioLatentInterpolate, + "AudioLatentBlend": AudioLatentBlend, + "AudioFeatureExtractor": AudioFeatureExtractor, + "AudioLatentManipulator": AudioLatentManipulator, + "AudioStyleTransfer": AudioStyleTransfer, + "AudioLatentGuidance": AudioLatentGuidance, + "AudioReferenceEncoder": AudioReferenceEncoder, + } -__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file + NODE_DISPLAY_NAME_MAPPINGS: Dict[str, str] = { + "AudioVAEEncode": "Audio VAE Encode", + "AudioVAEDecode": "Audio VAE Decode", + "AudioLatentInterpolate": "Audio Latent Interpolate", + "AudioLatentBlend": "Audio Latent Blend", + "AudioFeatureExtractor": "Audio Feature Extractor", + "AudioLatentManipulator": "Audio Latent Manipulator", + "AudioStyleTransfer": "Audio Style Transfer", + "AudioLatentGuidance": "Audio Latent Guidance", + "AudioReferenceEncoder": "Audio Reference Encoder", + } + +except Exception as _audio_import_error: # pragma: no cover + # Dependencies for audio nodes are missing; disable audio nodes gracefully + NODE_CLASS_MAPPINGS = {} + NODE_DISPLAY_NAME_MAPPINGS = {} + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/src/network_bending/audio_nodes/audio_latent_nodes.py b/src/network_bending/audio_nodes/audio_latent_nodes.py index e4343d8..712f40f 100644 --- a/src/network_bending/audio_nodes/audio_latent_nodes.py +++ b/src/network_bending/audio_nodes/audio_latent_nodes.py @@ -60,7 +60,9 @@ class AudioVAEEncode: # Normalize audio if normalize: - waveform = waveform / torch.max(torch.abs(waveform)) + max_abs = torch.max(torch.abs(waveform)) + if float(max_abs) > 1e-12: + waveform = waveform / max_abs # Ensure correct shape for VAE (batch, channels, samples) if waveform.dim() == 2: @@ -128,9 +130,8 @@ class AudioVAEDecode: # Move to CPU for further processing waveform = waveform.cpu() - # Denormalize if needed + # Denormalize if needed (clamp to valid range) if denormalize: - # Ensure audio is in valid range [-1, 1] waveform = torch.clamp(waveform, -1.0, 1.0) # Extract sample rate (default to 44100 if not stored) @@ -200,8 +201,8 @@ class AudioLatentInterpolate: elif interpolation_mode == "spherical": # Spherical linear interpolation (SLERP) # Normalize latents - latent_a_norm = F.normalize(latent_a.flatten(1), dim=1).reshape(latent_a.shape) - latent_b_norm = F.normalize(latent_b.flatten(1), dim=1).reshape(latent_b.shape) + latent_a_norm = F.normalize(latent_a.flatten(1), dim=1, eps=1e-6).reshape(latent_a.shape) + latent_b_norm = F.normalize(latent_b.flatten(1), dim=1, eps=1e-6).reshape(latent_b.shape) # Compute angle between latents dot_product = (latent_a_norm * latent_b_norm).sum() @@ -295,10 +296,11 @@ class AudioLatentBlend: latents.append(latent_d.to(device)) weights.append(weight_d) - # Normalize weights if requested + # Normalize weights if requested by the user if normalize: total_weight = sum(weights) - weights = [w / total_weight for w in weights] + if abs(total_weight) > 1e-12: + weights = [w / total_weight for w in weights] # Apply blend mode if blend_mode == "add": diff --git a/src/network_bending/audio_workflows/audio_style_transfer_workflow.json b/src/network_bending/audio_workflows/audio_style_transfer_workflow.json deleted file mode 100644 index 3d88428..0000000 --- a/src/network_bending/audio_workflows/audio_style_transfer_workflow.json +++ /dev/null @@ -1,219 +0,0 @@ -{ - "last_node_id": 10, - "last_link_id": 12, - "nodes": [ - { - "id": 1, - "type": "LoadAudio", - "pos": [100, 100], - "size": [300, 100], - "outputs": [ - { - "name": "AUDIO", - "type": "AUDIO", - "links": [1, 2] - } - ], - "properties": { - "Node name for S&R": "LoadAudio" - }, - "widgets_values": ["content_audio.wav"] - }, - { - "id": 2, - "type": "LoadAudio", - "pos": [100, 250], - "size": [300, 100], - "outputs": [ - { - "name": "AUDIO", - "type": "AUDIO", - "links": [3] - } - ], - "properties": { - "Node name for S&R": "LoadAudio" - }, - "widgets_values": ["style_audio.wav"] - }, - { - "id": 3, - "type": "LoadVAE", - "pos": [100, 400], - "size": [300, 100], - "outputs": [ - { - "name": "VAE", - "type": "VAE", - "links": [4, 5, 6] - } - ], - "properties": { - "Node name for S&R": "LoadVAE" - }, - "widgets_values": ["stable_audio_vae.safetensors"] - }, - { - "id": 4, - "type": "AudioVAEEncode", - "pos": [450, 100], - "size": [300, 150], - "inputs": [ - { - "name": "audio", - "type": "AUDIO", - "link": 1 - }, - { - "name": "vae", - "type": "VAE", - "link": 4 - } - ], - "outputs": [ - { - "name": "latent", - "type": "AUDIO_LATENT", - "links": [7] - }, - { - "name": "info", - "type": "LATENT_INFO", - "links": null - } - ], - "properties": { - "Node name for S&R": "AudioVAEEncode" - }, - "widgets_values": [true, 44100] - }, - { - "id": 5, - "type": "AudioVAEEncode", - "pos": [450, 300], - "size": [300, 150], - "inputs": [ - { - "name": "audio", - "type": "AUDIO", - "link": 3 - }, - { - "name": "vae", - "type": "VAE", - "link": 5 - } - ], - "outputs": [ - { - "name": "latent", - "type": "AUDIO_LATENT", - "links": [8] - }, - { - "name": "info", - "type": "LATENT_INFO", - "links": null - } - ], - "properties": { - "Node name for S&R": "AudioVAEEncode" - }, - "widgets_values": [true, 44100] - }, - { - "id": 6, - "type": "AudioStyleTransfer", - "pos": [800, 200], - "size": [350, 200], - "inputs": [ - { - "name": "content_latent", - "type": "AUDIO_LATENT", - "link": 7 - }, - { - "name": "style_latent", - "type": "AUDIO_LATENT", - "link": 8 - } - ], - "outputs": [ - { - "name": "latent", - "type": "AUDIO_LATENT", - "links": [9] - } - ], - "properties": { - "Node name for S&R": "AudioStyleTransfer" - }, - "widgets_values": ["adaptive", 0.7, 0.3, 4] - }, - { - "id": 7, - "type": "AudioVAEDecode", - "pos": [1200, 200], - "size": [300, 150], - "inputs": [ - { - "name": "latent", - "type": "AUDIO_LATENT", - "link": 9 - }, - { - "name": "vae", - "type": "VAE", - "link": 6 - } - ], - "outputs": [ - { - "name": "audio", - "type": "AUDIO", - "links": [10] - } - ], - "properties": { - "Node name for S&R": "AudioVAEDecode" - }, - "widgets_values": [true] - }, - { - "id": 8, - "type": "SaveAudio", - "pos": [1550, 200], - "size": [300, 100], - "inputs": [ - { - "name": "audio", - "type": "AUDIO", - "link": 10 - } - ], - "properties": { - "Node name for S&R": "SaveAudio" - }, - "widgets_values": ["styled_output.wav"] - } - ], - "links": [ - [1, 1, 0, 4, 0, "AUDIO"], - [2, 1, 0, 9, 0, "AUDIO"], - [3, 2, 0, 5, 0, "AUDIO"], - [4, 3, 0, 4, 1, "VAE"], - [5, 3, 0, 5, 1, "VAE"], - [6, 3, 0, 7, 1, "VAE"], - [7, 4, 0, 6, 0, "AUDIO_LATENT"], - [8, 5, 0, 6, 1, "AUDIO_LATENT"], - [9, 6, 0, 7, 0, "AUDIO_LATENT"], - [10, 7, 0, 8, 0, "AUDIO"] - ], - "config": {}, - "groups": [], - "version": 1, - "workflow": { - "name": "Audio Style Transfer", - "description": "Transfer audio style characteristics from one audio to another using latent space manipulation" - } -} \ No newline at end of file diff --git a/src/network_bending/js/network_bending.js b/src/network_bending/js/network_bending.js index 1020bd8..3b3ef97 100644 --- a/src/network_bending/js/network_bending.js +++ b/src/network_bending/js/network_bending.js @@ -1,5 +1,6 @@ -import { app } from "../../../scripts/app.js"; -import { api } from "../../../scripts/api.js"; +// Use absolute paths as ComfyUI serves these from /scripts +import { app } from "/scripts/app.js"; +import { api } from "/scripts/api.js"; // Register the network bending extension app.registerExtension({ diff --git a/src/network_bending/nodes.py b/src/network_bending/nodes.py index 8ca577b..ba5790c 100644 --- a/src/network_bending/nodes.py +++ b/src/network_bending/nodes.py @@ -1,4 +1,17 @@ -from server import PromptServer +try: + from server import PromptServer +except Exception: # pragma: no cover + class _DummyPromptServer: + instance = None + + def __init__(self): + class _Inst: + def send_sync(self, *args, **kwargs): + return None + + self.instance = _Inst() + + PromptServer = _DummyPromptServer() import torch import torch.nn as nn import random @@ -117,12 +130,15 @@ class NetworkBending: # Send feedback to UI message = f"Applied {operation} to {len(modified_layers)} layers with intensity {intensity}" - PromptServer.instance.send_sync("network_bending.feedback", { - "message": message, - "operation": operation, - "modified_layers": modified_layers[:10], # Limit to first 10 for UI - "total_layers": len(modified_layers) - }) + try: + PromptServer.instance.send_sync("network_bending.feedback", { + "message": message, + "operation": operation, + "modified_layers": modified_layers[:10], # Limit to first 10 for UI + "total_layers": len(modified_layers) + }) + except Exception: + pass return (model_clone,) @@ -157,7 +173,11 @@ class NetworkBending: modified = [] for name, param in model.named_parameters(): if self._should_modify_layer(name, patterns) and param.requires_grad: - threshold = torch.quantile(torch.abs(param.data), intensity) + abs_param = torch.abs(param.data) + if abs_param.numel() == 0: + continue + q = float(max(0.0, min(1.0, intensity))) + threshold = torch.quantile(abs_param, q) mask = torch.abs(param.data) > threshold # Convert mask to the same dtype as the parameter param.data.mul_(mask.to(dtype=param.dtype)) @@ -204,9 +224,13 @@ class NetworkBending: # Normalize to 0-1, quantize, then rescale min_val = param.data.min() max_val = param.data.max() - normalized = (param.data - min_val) / (max_val - min_val + 1e-8) + denom = (max_val - min_val) + if float(denom.abs().item()) < 1e-12: + modified.append(name) + continue + normalized = (param.data - min_val) / (denom + 1e-8) quantized = torch.round(normalized * (num_levels - 1)) / (num_levels - 1) - param.data = quantized * (max_val - min_val) + min_val + param.data = quantized * denom + min_val modified.append(name) return modified @@ -320,6 +344,532 @@ class ModelMixer: return (result,) +class InvertedPruning: + """ + Inverted Model Pruning - Selectively removes critical weights for artistic degradation + + Instead of preserving important weights (standard pruning), this node removes them, + creating unique artistic effects. Based on the inverted Lottery Ticket Hypothesis. + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL", {"tooltip": "The model to apply inverted pruning to"}), + "pruning_mode": ([ + "magnitude_inverted", + "structured_inverted", + "attention_head_removal", + "channel_pruning_inverted", + "gradient_based_inverted" + ], { + "default": "magnitude_inverted", + "tooltip": "Type of inverted pruning to apply" + }), + "threshold": ("FLOAT", { + "default": 0.1, + "min": 0.0, + "max": 0.99, + "step": 0.01, + "tooltip": "Percentage of weights to remove (0.1 = remove top 10% most important)" + }), + "target_layers": ("STRING", { + "default": "all", + "multiline": False, + "tooltip": "Comma-separated layer patterns (e.g., 'attention', 'conv', 'mlp')" + }), + "preserve_functionality": ("FLOAT", { + "default": 0.0, + "min": 0.0, + "max": 1.0, + "step": 0.1, + "tooltip": "How much to preserve base functionality (0=pure degradation, 1=mild effect)" + }), + "seed": ("INT", { + "default": -1, + "min": -1, + "max": 0xffffffffffffffff, + "tooltip": "Random seed for reproducible pruning (-1 for random)" + }), + }, + "optional": { + "gradient_accumulation_steps": ("INT", { + "default": 1, + "min": 1, + "max": 10, + "tooltip": "Number of gradient accumulation steps for more stable importance estimation" + }), + "use_actual_gradients": ("BOOLEAN", { + "default": True, + "tooltip": "Use actual gradient computation (slower but more accurate) or simplified method" + }), + "gradient_loss_type": ([ + "reconstruction", + "magnitude", + "perceptual", + "variance" + ], { + "default": "reconstruction", + "tooltip": "Loss function for gradient computation" + }), + } + } + + RETURN_TYPES = ("MODEL",) + RETURN_NAMES = ("model",) + FUNCTION = "apply_inverted_pruning" + CATEGORY = "network_bending" + OUTPUT_TOOLTIPS = ("Model with inverted pruning applied",) + + def apply_inverted_pruning(self, model, pruning_mode, threshold, target_layers, preserve_functionality, seed, + gradient_accumulation_steps=1, use_actual_gradients=True, gradient_loss_type="reconstruction"): + # Clone the model + model_clone = model.clone() + + # Set random seed + if seed != -1: + torch.manual_seed(seed) + np.random.seed(seed) + random.seed(seed) + + # Get the actual model + sd_model = model_clone.model if hasattr(model_clone, 'model') else model_clone + + # Parse target layers + target_patterns = [pattern.strip() for pattern in target_layers.split(',')] + if 'all' in target_patterns: + target_patterns = None + + # Track modified layers + modified_layers = [] + + # Apply the selected pruning mode + if pruning_mode == "magnitude_inverted": + modified_layers = self._magnitude_inverted_pruning(sd_model, threshold, preserve_functionality, target_patterns) + elif pruning_mode == "structured_inverted": + modified_layers = self._structured_inverted_pruning(sd_model, threshold, preserve_functionality, target_patterns) + elif pruning_mode == "attention_head_removal": + modified_layers = self._attention_head_removal(sd_model, threshold, preserve_functionality, target_patterns) + elif pruning_mode == "channel_pruning_inverted": + modified_layers = self._channel_pruning_inverted(sd_model, threshold, preserve_functionality, target_patterns) + elif pruning_mode == "gradient_based_inverted": + modified_layers = self._gradient_based_inverted(sd_model, threshold, preserve_functionality, target_patterns, + use_actual_gradients, gradient_accumulation_steps, gradient_loss_type) + + # Send feedback + try: + PromptServer.instance.send_sync("network_bending.feedback", { + "message": f"Applied {pruning_mode} to {len(modified_layers)} layers", + "operation": pruning_mode, + "modified_layers": modified_layers[:10], + "total_layers": len(modified_layers), + "threshold": threshold + }) + except Exception: + pass + + return (model_clone,) + + def _should_modify_layer(self, layer_name: str, patterns: List[str] = None) -> bool: + """Check if a layer should be modified based on target patterns""" + if patterns is None: + return True + return any(pattern.lower() in layer_name.lower() for pattern in patterns) + + def _magnitude_inverted_pruning(self, model: nn.Module, threshold: float, preserve: float, patterns: List[str] = None) -> List[str]: + """Remove weights with highest magnitude (opposite of standard magnitude pruning)""" + modified = [] + + for name, param in model.named_parameters(): + if self._should_modify_layer(name, patterns) and param.requires_grad: + # Calculate magnitude threshold - we want to remove the TOP magnitude weights + abs_weights = torch.abs(param.data) + k = int(threshold * param.data.numel()) + + if k > 0: + # Find threshold value - weights above this will be removed + threshold_val = torch.topk(abs_weights.flatten(), k).values[-1] + + # Create mask - True where we want to KEEP weights (low magnitude) + mask = abs_weights <= threshold_val + + # Apply preservation factor + if preserve > 0: + # Randomly preserve some high-magnitude weights + preserve_mask = torch.rand_like(param.data) < preserve + mask = mask | preserve_mask + + # Apply mask + param.data.mul_(mask.to(dtype=param.dtype)) + modified.append(name) + + return modified + + def _structured_inverted_pruning(self, model: nn.Module, threshold: float, preserve: float, patterns: List[str] = None) -> List[str]: + """Remove entire structures (channels/filters) with highest importance""" + modified = [] + + for name, module in model.named_modules(): + if not self._should_modify_layer(name, patterns): + continue + + # Handle Conv2d layers + if isinstance(module, nn.Conv2d): + weight = module.weight.data + # Calculate importance per output channel (L2 norm) + importance = torch.norm(weight, p=2, dim=(1, 2, 3)) + + # Remove channels with HIGHEST importance + k = int(threshold * len(importance)) + if k > 0: + _, indices_to_remove = torch.topk(importance, k) + + # Apply preservation + if preserve > 0: + num_preserve = int(k * preserve) + indices_to_remove = indices_to_remove[num_preserve:] + + # Zero out high-importance channels + weight[indices_to_remove] = 0 + modified.append(f"{name}.weight") + + # Handle Linear layers + elif isinstance(module, nn.Linear): + weight = module.weight.data + # Calculate importance per output neuron + importance = torch.norm(weight, p=2, dim=1) + + # Remove neurons with HIGHEST importance + k = int(threshold * len(importance)) + if k > 0: + _, indices_to_remove = torch.topk(importance, k) + + # Apply preservation + if preserve > 0: + num_preserve = int(k * preserve) + indices_to_remove = indices_to_remove[num_preserve:] + + # Zero out high-importance neurons + weight[indices_to_remove] = 0 + modified.append(f"{name}.weight") + + return modified + + def _attention_head_removal(self, model: nn.Module, threshold: float, preserve: float, patterns: List[str] = None) -> List[str]: + """Remove most important attention heads in transformer models""" + modified = [] + + for name, module in model.named_modules(): + # Look for multi-head attention modules + if ('attention' in name.lower() or 'attn' in name.lower()) and self._should_modify_layer(name, patterns): + # Check for Q, K, V projections + for proj_name in ['q_proj', 'k_proj', 'v_proj', 'query', 'key', 'value']: + if hasattr(module, proj_name): + proj = getattr(module, proj_name) + if isinstance(proj, nn.Linear): + weight = proj.weight.data + + # Assume head dimension is last dimension / num_heads + # This is a simplified approach - real implementation would need model-specific logic + if weight.shape[0] % 8 == 0: # Assume 8 heads for simplicity + num_heads = 8 + head_dim = weight.shape[0] // num_heads + + # Calculate importance per head + weight_reshaped = weight.view(num_heads, head_dim, -1) + head_importance = torch.norm(weight_reshaped, p=2, dim=(1, 2)) + + # Remove heads with HIGHEST importance + k = max(1, int(threshold * num_heads)) + _, heads_to_remove = torch.topk(head_importance, k) + + # Apply preservation + if preserve > 0: + num_preserve = int(k * preserve) + heads_to_remove = heads_to_remove[num_preserve:] + + # Zero out high-importance heads + for head_idx in heads_to_remove: + start_idx = head_idx * head_dim + end_idx = (head_idx + 1) * head_dim + weight[start_idx:end_idx] = 0 + + modified.append(f"{name}.{proj_name}") + + return modified + + def _channel_pruning_inverted(self, model: nn.Module, threshold: float, preserve: float, patterns: List[str] = None) -> List[str]: + """Remove most important channels in convolutional layers""" + modified = [] + + # First pass: calculate channel importance across the network + channel_importance = {} + + for name, module in model.named_modules(): + if isinstance(module, nn.Conv2d) and self._should_modify_layer(name, patterns): + weight = module.weight.data + + # Calculate importance for input channels + in_importance = torch.norm(weight, p=2, dim=(0, 2, 3)) + # Calculate importance for output channels + out_importance = torch.norm(weight, p=2, dim=(1, 2, 3)) + + channel_importance[name] = { + 'in': in_importance, + 'out': out_importance, + 'module': module + } + + # Second pass: prune channels + for name, info in channel_importance.items(): + module = info['module'] + + # Prune output channels + out_importance = info['out'] + k = int(threshold * len(out_importance)) + if k > 0: + _, indices_to_remove = torch.topk(out_importance, k) + + # Apply preservation + if preserve > 0: + num_preserve = int(k * preserve) + indices_to_remove = indices_to_remove[num_preserve:] + + # Zero out channels + module.weight.data[indices_to_remove] = 0 + if module.bias is not None: + module.bias.data[indices_to_remove] = 0 + + modified.append(f"{name}.weight") + + return modified + + def _gradient_based_inverted(self, model: nn.Module, threshold: float, preserve: float, patterns: List[str] = None, + use_actual_gradients: bool = True, accumulation_steps: int = 1, loss_type: str = "reconstruction") -> List[str]: + """Remove weights with highest gradient magnitude (most important for loss)""" + modified = [] + + try: + # Check if we should use actual gradients + if not use_actual_gradients: + return self._gradient_based_inverted_simple(model, threshold, preserve, patterns) + + # Attempt to compute actual gradients + gradients = self._compute_gradients(model, accumulation_steps, loss_type) + + if gradients: + # Use actual gradients for importance + for name, param in model.named_parameters(): + if self._should_modify_layer(name, patterns) and param.requires_grad and name in gradients: + importance = torch.abs(gradients[name]) + + k = int(threshold * param.data.numel()) + if k > 0: + # Find threshold value - remove weights with highest gradient magnitude + threshold_val = torch.topk(importance.flatten(), k).values[-1] + + # Create mask - keep low importance weights + mask = importance <= threshold_val + + # Apply preservation + if preserve > 0: + preserve_mask = torch.rand_like(param.data) < preserve + mask = mask | preserve_mask + + # Apply mask + param.data.mul_(mask.to(dtype=param.dtype)) + modified.append(name) + else: + # Fallback to simplified version + for name, param in model.named_parameters(): + if self._should_modify_layer(name, patterns) and param.requires_grad: + # Use weight magnitude as proxy for importance + importance = torch.abs(param.data) + torch.randn_like(param.data) * 0.1 + + k = int(threshold * param.data.numel()) + if k > 0: + threshold_val = torch.topk(importance.flatten(), k).values[-1] + mask = importance <= threshold_val + + if preserve > 0: + preserve_mask = torch.rand_like(param.data) < preserve + mask = mask | preserve_mask + + param.data.mul_(mask.to(dtype=param.dtype)) + modified.append(name) + + except Exception as e: + # If gradient computation fails, fallback to simplified version + print(f"Gradient computation failed: {e}. Using simplified importance estimation.") + return self._gradient_based_inverted_simple(model, threshold, preserve, patterns) + + return modified + + def _gradient_based_inverted_simple(self, model: nn.Module, threshold: float, preserve: float, patterns: List[str] = None) -> List[str]: + """Simplified gradient-based pruning using weight magnitude as proxy""" + modified = [] + for name, param in model.named_parameters(): + if self._should_modify_layer(name, patterns) and param.requires_grad: + importance = torch.abs(param.data) + torch.randn_like(param.data) * 0.1 + k = int(threshold * param.data.numel()) + if k > 0: + threshold_val = torch.topk(importance.flatten(), k).values[-1] + mask = importance <= threshold_val + if preserve > 0: + preserve_mask = torch.rand_like(param.data) < preserve + mask = mask | preserve_mask + param.data.mul_(mask.to(dtype=param.dtype)) + modified.append(name) + return modified + + def _compute_gradients(self, model: nn.Module, accumulation_steps: int = 1, loss_type: str = "reconstruction") -> Dict[str, torch.Tensor]: + """Compute actual gradients for weight importance""" + gradients = {} + accumulated_gradients = {} + + # Store original training mode + was_training = model.training + model.eval() + + try: + with torch.enable_grad(): + # Accumulate gradients over multiple steps for stability + for step in range(accumulation_steps): + # Zero existing gradients + model.zero_grad() + + # Generate sample input based on model type + sample_input = self._generate_sample_input(model) + if sample_input is None: + return {} + + # Forward pass + output = model(sample_input) + + # Compute loss based on specified type + loss = self._compute_importance_loss(output, sample_input, loss_type) + + # Backward pass + loss.backward() + + # Accumulate gradients + for name, param in model.named_parameters(): + if param.grad is not None: + if name not in accumulated_gradients: + accumulated_gradients[name] = param.grad.data.clone() + else: + accumulated_gradients[name] += param.grad.data + + # Average accumulated gradients + for name, grad in accumulated_gradients.items(): + gradients[name] = grad / accumulation_steps + + # Clear gradients to free memory + model.zero_grad() + + except Exception as e: + print(f"Error computing gradients: {e}") + gradients = {} + + finally: + # Restore original training mode + model.train(was_training) + + return gradients + + def _generate_sample_input(self, model: nn.Module) -> torch.Tensor: + """Generate appropriate sample input for the model""" + try: + # Get device + device = next(model.parameters()).device + + # Try to detect model type and generate appropriate input + # This is a heuristic approach - could be improved with model-specific logic + + # Check for common stable diffusion shapes + if hasattr(model, 'in_channels'): + # Likely a UNet or similar + batch_size = 1 + channels = getattr(model, 'in_channels', 4) + height = width = 64 # Use smaller size for efficiency + return torch.randn(batch_size, channels, height, width, device=device) + + # Check for transformer-like models + has_embedding = any('embed' in name for name, _ in model.named_modules()) + if has_embedding: + # Likely a transformer + batch_size = 1 + seq_length = 77 # Common for text transformers + hidden_dim = 768 # Common dimension + return torch.randn(batch_size, seq_length, hidden_dim, device=device) + + # Default: try to infer from first layer + for name, module in model.named_modules(): + if isinstance(module, nn.Conv2d): + # Image input + batch_size = 1 + channels = module.in_channels + height = width = 64 + return torch.randn(batch_size, channels, height, width, device=device) + elif isinstance(module, nn.Linear) and 'embed' not in name: + # Vector input + batch_size = 1 + input_dim = module.in_features + return torch.randn(batch_size, input_dim, device=device) + + # If we can't determine, return None + return None + + except Exception as e: + print(f"Error generating sample input: {e}") + return None + + def _compute_importance_loss(self, output: torch.Tensor, input_tensor: torch.Tensor, loss_type: str = "reconstruction") -> torch.Tensor: + """Compute loss for measuring weight importance""" + try: + if loss_type == "reconstruction": + # Reconstruction loss + if output.shape == input_tensor.shape: + return torch.nn.functional.mse_loss(output, input_tensor) + else: + # Feature matching as fallback + return torch.nn.functional.l1_loss(output.mean(), input_tensor.mean()) + + elif loss_type == "magnitude": + # Simple magnitude loss + return output.abs().mean() + + elif loss_type == "perceptual": + # Perceptual loss using feature statistics + output_mean = output.mean(dim=list(range(2, output.dim()))) + output_std = output.std(dim=list(range(2, output.dim()))) + + if input_tensor.shape == output.shape: + input_mean = input_tensor.mean(dim=list(range(2, input_tensor.dim()))) + input_std = input_tensor.std(dim=list(range(2, input_tensor.dim()))) + mean_loss = torch.nn.functional.mse_loss(output_mean, input_mean) + std_loss = torch.nn.functional.mse_loss(output_std, input_std) + return mean_loss + std_loss + else: + # Use output statistics only + return output_mean.abs().mean() + output_std.abs().mean() + + elif loss_type == "variance": + # Maximize variance (inverse of typical loss) + # High variance = high importance + return -output.var() + + else: + # Default to magnitude loss + return output.abs().mean() + + except Exception as e: + print(f"Error in loss computation: {e}. Using magnitude loss.") + # Fallback to simple magnitude loss + return output.abs().mean() + + class LatentFormatConverter: """ Convert between audio and image latent formats @@ -426,13 +976,16 @@ class LatentFormatConverter: latent["samples"] = converted # Send feedback - PromptServer.instance.send_sync("network_bending.feedback", { - "message": f"Converted latent from {list(latent_tensor.shape)} to {list(converted.shape)}", - "operation": conversion_mode, - "input_shape": list(latent_tensor.shape), - "output_shape": list(converted.shape), - "method": reshape_method - }) + try: + PromptServer.instance.send_sync("network_bending.feedback", { + "message": f"Converted latent from {list(latent_tensor.shape)} to {list(converted.shape)}", + "operation": conversion_mode, + "input_shape": list(latent_tensor.shape), + "output_shape": list(converted.shape), + "method": reshape_method + }) + except Exception: + pass return (latent,) @@ -780,12 +1333,15 @@ class VAENetworkBending: modified_layers.extend(self._progressive_corruption(decoder, intensity, "decoder")) # Send feedback - PromptServer.instance.send_sync("network_bending.feedback", { - "message": f"Applied {operation} to VAE {target_component} with intensity {intensity}", - "operation": operation, - "modified_layers": len(modified_layers), - "target": target_component - }) + try: + PromptServer.instance.send_sync("network_bending.feedback", { + "message": f"Applied {operation} to VAE {target_component} with intensity {intensity}", + "operation": operation, + "modified_layers": len(modified_layers), + "target": target_component + }) + except Exception: + pass return (vae_clone,) @@ -1360,6 +1916,7 @@ NODE_CLASS_MAPPINGS = { "NetworkBending": NetworkBending, "NetworkBendingAdvanced": NetworkBendingAdvanced, "ModelMixer": ModelMixer, + "InvertedPruning": InvertedPruning, "LatentFormatConverter": LatentFormatConverter, "VAENetworkBending": VAENetworkBending, "VAEMixer": VAEMixer, @@ -1374,6 +1931,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "NetworkBending": "Network Bending", "NetworkBendingAdvanced": "Network Bending (Advanced)", "ModelMixer": "Model Mixer", + "InvertedPruning": "Inverted Pruning", "LatentFormatConverter": "Latent Format Converter", "VAENetworkBending": "VAE Network Bending", "VAEMixer": "VAE Mixer", diff --git a/tests/__init__.py b/tests/__init__.py deleted file mode 100644 index 916fc65..0000000 --- a/tests/__init__.py +++ /dev/null @@ -1 +0,0 @@ -"""Unit test package for network_bending.""" diff --git a/tests/conftest.py b/tests/conftest.py deleted file mode 100644 index 310609c..0000000 --- a/tests/conftest.py +++ /dev/null @@ -1,6 +0,0 @@ -import os -import sys - -# Add the project root directory to Python path -# This allows the tests to import the project -sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..'))) diff --git a/tests/pytest.ini b/tests/pytest.ini deleted file mode 100644 index 95c76f1..0000000 --- a/tests/pytest.ini +++ /dev/null @@ -1,4 +0,0 @@ -[pytest] -testpaths = . # Run tests in the current directory -python_files = test_*.py # Run tests in files that start with "test_" -norecursedirs = .. # Don't run tests in the parent directory diff --git a/tests/test_network_bending.py b/tests/test_network_bending.py deleted file mode 100644 index 719735b..0000000 --- a/tests/test_network_bending.py +++ /dev/null @@ -1,214 +0,0 @@ -#!/usr/bin/env python - -"""Tests for `network_bending` package.""" - -import pytest -import torch -import torch.nn as nn -from unittest.mock import Mock, MagicMock - -# Import the nodes from the package -from src.network_bending.nodes import NetworkBending, NetworkBendingAdvanced, ModelMixer - - -class SimpleModel(nn.Module): - """A simple test model""" - def __init__(self): - super().__init__() - self.conv1 = nn.Conv2d(3, 16, 3, padding=1) - self.conv2 = nn.Conv2d(16, 32, 3, padding=1) - self.linear = nn.Linear(32 * 8 * 8, 10) - self.norm = nn.BatchNorm2d(32) - - def forward(self, x): - x = self.conv1(x) - x = self.conv2(x) - x = self.norm(x) - x = x.view(x.size(0), -1) - x = self.linear(x) - return x - - -@pytest.fixture -def mock_model(): - """Create a mock ComfyUI model wrapper""" - model = Mock() - model.model = SimpleModel() - model.clone = Mock(return_value=model) - return model - - -@pytest.fixture -def mock_prompt_server(monkeypatch): - """Mock the PromptServer for testing""" - mock_server = Mock() - mock_instance = Mock() - mock_instance.send_sync = Mock() - mock_server.instance = mock_instance - - # Create a mock module for server - import sys - from types import ModuleType - server_module = ModuleType('server') - server_module.PromptServer = mock_server - sys.modules['server'] = server_module - - return mock_instance - - -class TestNetworkBending: - """Test the NetworkBending node""" - - def test_input_types(self): - """Test that INPUT_TYPES returns correct structure""" - input_types = NetworkBending.INPUT_TYPES() - - assert "required" in input_types - assert "model" in input_types["required"] - assert "operation" in input_types["required"] - assert "intensity" in input_types["required"] - assert "target_layers" in input_types["required"] - assert "seed" in input_types["required"] - - # Check operation list - operations = input_types["required"]["operation"][0] - assert "add_noise" in operations - assert "scale_weights" in operations - assert "prune_weights" in operations - - def test_add_noise_operation(self, mock_model, mock_prompt_server): - """Test add_noise operation""" - node = NetworkBending() - - # Run the operation - result = node.bend_network( - model=mock_model, - operation="add_noise", - intensity=0.1, - target_layers="all", - seed=42 - ) - - # Check that model was cloned - mock_model.clone.assert_called_once() - - # Check that result is returned - assert result is not None - assert isinstance(result, tuple) - assert len(result) == 1 - - def test_target_layer_filtering(self, mock_model, mock_prompt_server): - """Test that target layer filtering works""" - node = NetworkBending() - - # Test with specific layer pattern - result = node.bend_network( - model=mock_model, - operation="add_noise", - intensity=0.1, - target_layers="conv", - seed=42 - ) - - # Verify feedback was sent - mock_prompt_server.send_sync.assert_called() - call_args = mock_prompt_server.send_sync.call_args - assert call_args[0][0] == "network_bending.feedback" - assert "conv" in str(call_args[0][1]["modified_layers"]) - - def test_scale_weights_operation(self, mock_model, mock_prompt_server): - """Test scale_weights operation""" - node = NetworkBending() - - result = node.bend_network( - model=mock_model, - operation="scale_weights", - intensity=0.7, # Should scale by 1.4 - target_layers="linear", - seed=42 - ) - - assert result is not None - - def test_seed_reproducibility(self, mock_model, mock_prompt_server): - """Test that setting seed produces reproducible results""" - node = NetworkBending() - - # Get initial weights - initial_weights = {} - for name, param in mock_model.model.named_parameters(): - initial_weights[name] = param.data.clone() - - # Run with seed - result1 = node.bend_network( - model=mock_model, - operation="add_noise", - intensity=0.1, - target_layers="all", - seed=12345 - ) - - # Weights should have changed - for name, param in mock_model.model.named_parameters(): - assert not torch.allclose(initial_weights[name], param.data) - - -class TestModelMixer: - """Test the ModelMixer node""" - - def test_input_types(self): - """Test that INPUT_TYPES returns correct structure""" - input_types = ModelMixer.INPUT_TYPES() - - assert "required" in input_types - assert "model_a" in input_types["required"] - assert "model_b" in input_types["required"] - assert "mix_mode" in input_types["required"] - assert "mix_ratio" in input_types["required"] - - def test_linear_interpolation(self, mock_model): - """Test linear interpolation mixing""" - node = ModelMixer() - - # Create two mock models - model_a = mock_model - model_b = Mock() - model_b.model = SimpleModel() - - # Set different weights for model_b - for param in model_b.model.parameters(): - param.data.fill_(2.0) - - result = node.mix_models( - model_a=model_a, - model_b=model_b, - mix_mode="linear_interpolation", - mix_ratio=0.5 - ) - - assert result is not None - assert isinstance(result, tuple) - - -class TestNetworkBendingAdvanced: - """Test the NetworkBendingAdvanced node""" - - def test_input_types(self): - """Test that INPUT_TYPES returns correct structure""" - input_types = NetworkBendingAdvanced.INPUT_TYPES() - - assert "required" in input_types - assert "model" in input_types["required"] - assert "operation" in input_types["required"] - assert "intensity" in input_types["required"] - assert "preserve_functionality" in input_types["required"] - - # Check advanced operations - operations = input_types["required"]["operation"][0] - assert "layer_swap" in operations - assert "activation_replace" in operations - assert "weight_transpose" in operations - - -if __name__ == "__main__": - pytest.main([__file__, "-v"])