Implement network_bending entrypoint and add InvertedPruning node; enhance audio nodes with error handling and normalization checks; remove obsolete test files.

This commit is contained in:
David Piazza
2025-10-02 11:39:45 -04:00
parent e430d3e686
commit a9a528942f
10 changed files with 686 additions and 513 deletions
+44
View File
@@ -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__ = [
+49 -37
View File
@@ -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']
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"]
@@ -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":
@@ -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"
}
}
+3 -2
View File
@@ -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({
+581 -23
View File
@@ -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",
-1
View File
@@ -1 +0,0 @@
"""Unit test package for network_bending."""
-6
View File
@@ -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__), '..')))
-4
View File
@@ -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
-214
View File
@@ -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"])