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:
+44
@@ -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__ = [
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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 +0,0 @@
|
||||
"""Unit test package for network_bending."""
|
||||
@@ -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__), '..')))
|
||||
@@ -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
|
||||
@@ -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"])
|
||||
Reference in New Issue
Block a user