new nodes

This commit is contained in:
Martin Bukowski
2024-01-22 14:12:36 -06:00
parent eab7b6bf16
commit 60706662cb
20 changed files with 1250 additions and 1195 deletions
+2
View File
@@ -150,3 +150,5 @@ cython_debug/
# and can be added to the global gitignore or merged into this file. For a more nuclear
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
#.idea/
experiment
+18 -8
View File
@@ -7,18 +7,28 @@ Merge two checkpoint models by dare ties (https://github.com/yule-BUAA/MergeLM),
|category|node name|input type|output type|desc.|
| --- | --- | --- | --- | --- |
|merge|DareModelMerger|`MODEL`, `MODEL`, `MODEL`|`MODEL`|Performs a DARE block merge|
|merge|MagnitudeModelMerger|`MODEL`, `MODEL`|`MODEL`|Performs a MP block merge|
|merge|BlockModelMergerAdv|`MODEL`, `MODEL`|`MODEL`|Performs a block merge with a custom merge type|
|unet|Model Merger (Masked)|`MODEL`, `MODEL`, `MODEL_MASK`|`MODEL`|Performs a masked block merge|
|unet|Model Merger (DARE)|`MODEL`, `MODEL`, `MODEL_MASK (optional)`|`MODEL`|Performs a DARE block merge|
|unet|MBW Merger (DARE)|`MODEL`, `MODEL`, `MODEL_MASK (optional)`|`MODEL`|Performs a DARE block merge, with full layer control (like MBW)|
|mask|Magnitude Masker|`MODEL`, `MODEL`|`MODEL_MASK`|Creates a mask based on the deltas of the parameters|
|clip|CLIP Merger (DARE)|`CLIP`, `CLIP`|`CLIP`|Performs a DARE merge on two CLIP|
|util|Normalize Model|`MODEL`, `MODEL`|`MODEL`|Normalizes one models parameter norm to another model|
### Merging
* In general, one means keep first model, zero means keep second model
* For DARE, we use the base model to determine which values to protect or include. It is optional, and if not provided will ignore exclude_a, include_b, threshold_type, and invert.
* For DARE, we use the base model to determine which values to protect or include.
* For TIES, we can use the sum of the delta ties (as in the paper), or the count, or off to disable.
* Larger density means to include more of the second model
* Larger exclude_a preserves more of the first model
* Larger include_b allows more of the second model
* input, middle, and out are the block region weights.
* Can accept a model mask, which will restrict changes to only modify the masked areas.
### Masks
* Larger thresholds on a mask means to reduce the amount of the second model more
* threshold_type is the way we determine where our threshold lies in our distribution since we use chunks, quantile will err towards the sparsity, median will use the chunk median.
* invert is whether we invert the threshold, so we keep the weights that are below the threshold instead of above.
## How to use
DARE-TIES does a stochastic selection of the parameters to keep, and then only performs updates in either the 'up' or 'down' direction. According to the paper, the workflow should be as such:
* Take our model A, and build a magnitude mask based on a base model.
* Take model B, and merge it in to model A using the mask to protect model A's largest parameters.
### Normalization
I am testing out a new normalization method, which is to normalize the norm of the parameters of one model to another. This is done by taking the ratio of the norms, and then scaling the parameters of the first model by that ratio. This is done in the `Normalize Model` node. There are a few options, of most interest is the 'q_norm' option, which only scales Q and K relative to each other.
+18 -12
View File
@@ -1,21 +1,27 @@
from .merge.dare import DareModelMerger
from .merge.dareMBW import DareModelMergerMBW
from .merge.mag import MagnitudePruningModelMerger
from .merge.block import BlockModelMergerAdv
from .components.clip import DareClipMerger
from .components.dare import DareUnetMerger
from .components.dare_mbw import DareUnetMergerMBW
from .components.block import BlockUnetMerger
from .components.normalize import NormalizeUnet
from .components.mask_model import MagnitudeMasker
NODE_CLASS_MAPPINGS = {
"DareModelMerger": DareModelMerger,
"DareModelMergerMBW": DareModelMergerMBW,
"MagnitudeModelMerger": MagnitudePruningModelMerger,
"BlockModelMergerAdv": BlockModelMergerAdv,
"DM_MaskedModelMerger": BlockUnetMerger,
"DM_DareModelMerger": DareUnetMerger,
"DM_DareModelMergerMBW": DareUnetMergerMBW,
"DM_DareClipMerger": DareClipMerger,
"DM_NormalizeModel": NormalizeUnet,
"DM_MagnitudeMasker": MagnitudeMasker,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DareModelMerger": "ModelMergeByDARE",
"DareModelMergerMBW": "ModelMergeByDAREMBW",
"MagnitudeModelMerger": "ModelMergeByMagnitudePruning",
"BlockModelMergerAdv": "ModelMergeByBlock (Advanced)",
"DM_MaskedModelMerger": "Model Merger (Masked)",
"DM_DareModelMerger": "Model Merger (DARE)",
"DM_DareModelMergerMBW": "MBW Merger (DARE)",
"DM_DareClipMerger": "CLIP Merger (DARE)",
"DM_NormalizeModel": "Normalize Model",
"DM_MagnitudeMasker": "Magnitude Masker",
}
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
+122
View File
@@ -0,0 +1,122 @@
# components/block.py
from comfy.model_patcher import ModelPatcher
import torch
from typing import Dict, Tuple, Optional
from ..ddare.merge import merge_tensors
from ..ddare.util import cuda_memory_profiler, get_device, get_patched_state
from ..ddare.mask import ModelMask
from ..ddare.const import UNET_CATEGORY
class BlockUnetMerger:
"""
A class to merge two diffusion U-Net models using m mask.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"model_mask": ("MODEL_MASK",),
"input": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"middle": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"out": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"method": (["comfy", "lerp", "slerp", "gradient"], ),
},
"optional": {
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = UNET_CATEGORY
def merge(self, model_a: ModelPatcher, model_b: ModelPatcher,
input : float, middle : float, out : float, time : float, method : str,
clear_cache : bool = True, model_mask: Optional[ModelMask] = None,
**kwargs) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
input (float): The ratio (lambda) of the input layer to keep from model_a.
middle (float): The ratio (lambda) of the middle layers to keep from model_a.
out (float): The ratio (lambda) of the output layer to keep from model_a.
time (float): The ratio (lambda) of the time layers to keep from model_a.
method (str): The method to use for merging, either "comfy", "lerp", "slerp", or "gradient".
clear_cache (bool): Whether to clear the CUDA cache after each chunk. Default is True.
model_mask (ModelMask): A ModelMask instance to use for masking the model. Default is None.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
device = get_device()
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
m = model_a.clone() # Clone model_a to keep its structure
with cuda_memory_profiler():
model_a_sd = get_patched_state(m)
model_b_sd = get_patched_state(model_b)
# Merge each parameter from model_b into model_a
for k in model_a_sd.keys():
if k not in model_b_sd:
print("could not patch. key doesn't exist in model:", k)
continue
k_unet = k[len("diffusion_model."):]
# Get our ratio for this layer
if k_unet.startswith("input"):
ratio = input
elif k_unet.startswith("middle"):
ratio = middle
elif k_unet.startswith("out"):
ratio = out
elif k_unet.startswith("time"):
ratio = time
else:
print(f"Unknown key: {k}, skipping.")
continue
# Apply sparsification by the delta for this layer
mask : torch.Tensor = model_mask.get_layer_mask(k) if model_mask is not None else None
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = model_b_sd[k]
if mask is None:
mask = torch.ones_like(a)
result_tensor = torch.where(mask.to(device), a.to(device), b.to(device))
del mask
if method == "comfy":
strength_patch = 1.0 - ratio
strength_model = ratio
else:
result_tensor = merge_tensors(method, a.to(device), result_tensor, 1 - ratio)
strength_model = 0
strength_patch = 1.0
m.add_patches({k: (result_tensor.to('cpu'),)}, strength_patch, strength_model)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return (m,)
+107
View File
@@ -0,0 +1,107 @@
# components/clip.py
from comfy.sd import CLIP
import torch
from typing import Optional
from ..ddare.merge import merge_tensors, dare_ties_sparsification
from ..ddare.util import cuda_memory_profiler, get_device
from ..ddare.const import CLIP_CATEGORY
class DareClipMerger:
"""
A class to merge two CLIP models using calculated deltas, sparsification,
and a weighted consensus method. This is the DARE-TIES method.
https://arxiv.org/pdf/2311.03099.pdf
"""
@classmethod
def INPUT_TYPES(s):
return {"required": { "clip_a": ("CLIP",),
"clip_b": ("CLIP",),
"ties": (["sum", "count", "off"], {"default": "sum"}),
"rescale": (["off", "on"], {"default": "off"}),
"ratio": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"drop_rate": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01}),
"seed": ("INT", {"default": 42}),
"method": (["comfy", "lerp", "slerp", "gradient"], ),
"iterations": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
}}
RETURN_TYPES = ("CLIP",)
FUNCTION = "merge"
CATEGORY = CLIP_CATEGORY
def merge(self, clip_a : CLIP, clip_b : CLIP, ratio : float, method : str, seed : Optional[int] = None, iterations : int = 1, clear_cache : bool = True, **kwargs):
"""
Merge two CLIP models using the DARE-TIES method.
Args:
clip_a (CLIP): The base CLIP model
clip_b (CLIP): The CLIP model to merge into the base
ratio (float): The ratio of the models to use. 1 is 100% model_a, 0 is 100% model_b.
method (str): The merge method to use
seed (Optional[int]): The seed to use for the random number generator
iterations (int): The number of iterations to use
clear_cache (bool): Whether to clear the GPU cache after merging
**kwargs: Unused
Returns:
"""
device = get_device()
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
with cuda_memory_profiler():
clip_a = clip_a.clone()
clip_a.patcher.patch_model()
clip_a_sd = clip_a.get_sd() # State dict of model_a
clip_a.patcher.unpatch_model()
clip_b = clip_b.clone()
clip_b.patcher.patch_model()
clip_b_sd = clip_b.get_sd()
clip_b.patcher.unpatch_model()
for k in clip_a_sd.keys():
if k.endswith(".position_ids") or k.endswith(".logit_scale"):
continue
if k not in clip_a_sd:
print("could not patch. key doesn't exist in model:", k)
continue
a : torch.Tensor = clip_a_sd[k]
b : torch.Tensor = clip_b_sd[k]
merged_a = a.clone()
for i in range(iterations):
if seed is not None:
torch.manual_seed(seed + i)
sparsified_delta = dare_ties_sparsification(merged_a, b, device=device, **kwargs)
if method == "comfy":
merged_a = sparsified_delta
strength_patch = 1.0 - ratio
strength_model = ratio
else:
merged_a = merge_tensors(method, merged_a.to(device), sparsified_delta.to(device), 1 - ratio)
strength_model = 0
strength_patch = 1.0
del sparsified_delta
#del a, b
# Apply the sparsified delta as a patch
nv = (merged_a.to('cpu'),)
clip_a.add_patches({k: nv}, strength_patch, strength_model)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return (clip_a,)
+217
View File
@@ -0,0 +1,217 @@
# components/dare.py
from comfy.model_patcher import ModelPatcher
import torch
from typing import Dict, Tuple, Optional
from ..ddare.merge import merge_tensors, dare_ties_sparsification
from ..ddare.util import cuda_memory_profiler, get_device, get_patched_state
from ..ddare.mask import ModelMask
from ..ddare.const import UNET_CATEGORY
class DareUnetMerger:
"""
A class to merge two diffusion U-Net models using calculated deltas, sparsification,
and a weighted consensus method. This is the DARE-TIES method.
https://arxiv.org/pdf/2311.03099.pdf
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"drop_rate": ("FLOAT", {"default": 0.90, "min": 0.0, "max": 1.0, "step": 0.01}),
"ties": (["sum", "count", "off"], {"default": "sum"}),
"rescale": (["off", "on"], {"default": "off"}),
"seed": ("INT", {"default": 1, "min":0, "max": 99999999999}),
"method": (["comfy", "lerp", "slerp", "gradient"], ),
"iterations": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"label": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"input": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"middle": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"output": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"out": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
},
"optional": {
"model_mask": ("MODEL_MASK",),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = UNET_CATEGORY
def merge(self, model_a: ModelPatcher, model_b: ModelPatcher,
method : str, seed : Optional[int] = None, clear_cache : bool = True,
model_mask: Optional[ModelMask] = None, iterations : int = 1,
**kwargs,) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
method (str): The method to use for merging, either "lerp", "slerp", or "gradient".
seed (int): The random seed to use for the merge.
clear_cache (bool): Whether to clear the CUDA cache after each chunk. Default is False.
iterations (int): The number of iterations to perform the merge. Default is 1.
model_mask (Optional[ModelMask]): The model mask to use for protection of our model_a. Default is None.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
device = get_device()
if seed is not None:
torch.manual_seed(seed)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
with cuda_memory_profiler():
m = model_a.clone() # Clone model_a to keep its structure
model_a_sd = get_patched_state(m)
model_b_sd = get_patched_state(model_b)
# Merge each parameter from model_b into model_a
for k in model_a_sd.keys():
if k not in model_b_sd:
print("could not patch. key doesn't exist in model:", k)
continue
ratio = self.calculate_layer_ratio(k, **kwargs)
if ratio is None:
continue
# Apply sparsification by the delta for this layer
mask : torch.Tensor = model_mask.get_layer_mask(k) if model_mask is not None else None
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = model_b_sd[k]
merged_a = a.clone()
for i in range(iterations):
if seed is not None:
torch.manual_seed(seed + i)
sparsified_delta = dare_ties_sparsification(merged_a, b, device=device, **kwargs)
# If we have a mask, apply it to the delta, replacing true values with our delta
if mask is not None:
sparsified_delta = torch.where(mask.to(device), sparsified_delta.to(device), merged_a.to(device))
if method == "comfy":
merged_a = sparsified_delta
strength_patch = 1.0 - ratio
strength_model = ratio
else:
merged_a = merge_tensors(method, merged_a.to(device), sparsified_delta.to(device), 1 - ratio)
strength_model = 0
strength_patch = 1.0
del sparsified_delta
# Apply the sparsified delta as a patch
m.add_patches({k: (merged_a.to('cpu'),)}, strength_patch, strength_model)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return (m,)
@classmethod
def scan_layer(cls, key : str, base : str, n : int, **kwargs):
for i in range(n):
my_key = f"{base}.{i}"
if key.startswith(my_key):
if my_key in kwargs:
return kwargs[my_key]
else:
print(f"No weight for {my_key}")
return None
elif i==(n-1):
print(f"Unknown key: {key},i={i}")
return None
return None
@classmethod
def calculate_layer_ratio(cls, key, **kwargs) -> Optional[float]:
k_unet = key[len("diffusion_model."):]
ratio = None
# Get our ratio for this layer
if k_unet.startswith(f"input_blocks."):
# use scan layer static
ratio = kwargs.get("input", None)
elif k_unet.startswith(f"middle_block."):
ratio = kwargs.get("middle", None)
elif k_unet.startswith(f"output_blocks."):
ratio = kwargs.get("output", None)
elif k_unet.startswith("out."):
ratio = kwargs.get("out", None)
elif k_unet.startswith("time"):
ratio = kwargs.get("time", None)
elif k_unet.startswith("label_emb"):
ratio = kwargs.get("label", None)
else:
print(f"Unknown key: {key}, skipping.")
return ratio
def apply_sparsification(self, base_model_param: Optional[torch.Tensor], model_a_param: torch.Tensor, model_b_param: torch.Tensor,
exclude_a: float, include_b: float, invert : str, drop_rate: float, ties : str, rescale : str,
device : torch.device, **kwargs) -> torch.Tensor:
"""
Applies sparsification to a tensor based on the specified sparsity level.
"""
model_a_flat = model_a_param.view(-1).float().to(device)
model_b_flat = model_b_param.view(-1).float().to(device)
delta_flat = model_b_flat - model_a_flat
if base_model_param is not None:
base_model_flat = base_model_param.view(-1).float().to(device)
delta_a_flat = model_a_flat - base_model_flat
delta_b_flat = model_b_flat - base_model_flat
include_mask = self.get_threshold_mask(delta_b_flat, include_b, invert, **kwargs)
exclude_mask = self.get_threshold_mask(delta_a_flat, exclude_a, invert, **kwargs)
base_mask = include_mask & (~exclude_mask)
del base_model_flat, delta_a_flat, delta_b_flat, include_mask, exclude_mask
else:
include_mask = torch.ones_like(model_a_flat).bool()
exclude_mask = torch.zeros_like(model_a_flat).bool()
base_mask = include_mask & (~exclude_mask)
del include_mask, exclude_mask
if ties != "off":
ties_mask = self.get_ties_mask(delta_flat, ties)
base_mask = base_mask & ties_mask
del ties_mask
dare_mask = torch.bernoulli(torch.full(delta_flat.shape, 1 - drop_rate, device=device)).bool()
# The paper says we should rescale, but it yields terrible results for SD
if rescale == "on":
# Rescale the remaining deltas
delta_flat = delta_flat / (1 - drop_rate)
final_mask = dare_mask & base_mask
# print(f"mask nonzero count: {torch.count_nonzero(mask)} dare nonzero count: {torch.count_nonzero(dare_mask)} base nonzero count: {torch.count_nonzero(base_mask)} include nonzero count: {torch.count_nonzero(include_mask)} exclude nonzero count: {torch.count_nonzero(exclude_mask)}")
sparsified_flat = torch.where(final_mask, model_a_flat + delta_flat, model_a_flat)
del final_mask, delta_flat, base_mask, model_a_flat, model_b_flat, dare_mask
return sparsified_flat.view_as(model_a_param)
+218
View File
@@ -0,0 +1,218 @@
# components/dare_mbw.py
from comfy.model_patcher import ModelPatcher
import torch
from typing import Dict, Tuple, Optional, Literal
from ..ddare.merge import merge_tensors, dare_ties_sparsification
from ..ddare.util import cuda_memory_profiler, get_device, get_patched_state
from ..ddare.mask import ModelMask
from ..ddare.const import UNET_CATEGORY
class DareUnetMergerMBW:
"""
A class to merge two diffusion U-Net models using calculated deltas, sparsification,
and a weighted consensus method. This is the DARE-TIES method, and allows for fine
grained control over the merge ratios for each layer like MBW.
https://arxiv.org/pdf/2311.03099.pdf
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
arg_dict = {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"drop_rate": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1.0, "step": 0.01}),
"ties": (["sum", "count", "off"], {"default": "sum"}),
"rescale": (["off", "on"], {"default": "off"}),
"seed": ("INT", {"default": 1, "min":0, "max": 99999999999}),
"method": (["comfy", "lerp", "slerp", "gradient"], ),
"iterations": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"label": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
}
argument = ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.0, "step": 0.01})
for i in range(12):
arg_dict[f"input_blocks.{i}"] = argument
for i in range(3):
arg_dict[f"middle_block.{i}"] = argument
for i in range(12):
arg_dict[f"output_blocks.{i}"] = argument
arg_dict["out"] = argument
opt = {"model_mask": ("MODEL_MASK",)}
return {"required": arg_dict ,"optional": opt}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = UNET_CATEGORY
def merge(self, model_a: ModelPatcher, model_b: ModelPatcher,
method : str, seed : Optional[int] = None, clear_cache : bool = True,
model_mask: Optional[ModelMask] = None, iterations : int = 1,
**kwargs,) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
method (str): The method to use for merging, either "lerp", "slerp", or "gradient".
seed (int): The random seed to use for the merge.
clear_cache (bool): Whether to clear the CUDA cache after each chunk. Default is False.
iterations (int): The number of iterations to perform the merge. Default is 1.
model_mask (Optional[ModelMask]): The model mask to use for protection of our model_a. Default is None.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
device = get_device()
if seed is not None:
torch.manual_seed(seed)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
with cuda_memory_profiler():
m = model_a.clone() # Clone model_a to keep its structure
model_a_sd = get_patched_state(m)
model_b_sd = get_patched_state(model_b)
# Merge each parameter from model_b into model_a
for k in model_a_sd.keys():
if k not in model_b_sd:
print("could not patch. key doesn't exist in model:", k)
continue
ratio = self.calculate_layer_ratio(k, **kwargs)
if ratio is None:
continue
# Apply sparsification by the delta for this layer
mask : torch.Tensor = model_mask.get_layer_mask(k) if model_mask is not None else None
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = model_b_sd[k]
merged_a = a.clone()
for i in range(iterations):
if seed is not None:
torch.manual_seed(seed + i)
sparsified_delta = dare_ties_sparsification(merged_a, b, device=device, **kwargs)
# If we have a mask, apply it to the delta, replacing true values with our delta
if mask is not None:
sparsified_delta = torch.where(mask.to(device), sparsified_delta.to(device), merged_a.to(device))
if method == "comfy":
merged_a = sparsified_delta
strength_patch = 1.0 - ratio
strength_model = ratio
else:
merged_a = merge_tensors(method, merged_a.to(device), sparsified_delta.to(device), 1 - ratio)
strength_model = 0
strength_patch = 1.0
del sparsified_delta
# Apply the sparsified delta as a patch
m.add_patches({k: (merged_a.to('cpu'),)}, strength_patch, strength_model)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return (m,)
@classmethod
def scan_layer(cls, key : str, base : str, n : int, **kwargs):
for i in range(n):
my_key = f"{base}.{i}"
if key.startswith(my_key):
if my_key in kwargs:
return kwargs[my_key]
else:
print(f"No weight for {my_key}")
return None
elif i==(n-1):
print(f"Unknown key: {key},i={i}")
return None
return None
@classmethod
def calculate_layer_ratio(cls, key, **kwargs) -> Optional[float]:
k_unet = key[len("diffusion_model."):]
ratio = None
# Get our ratio for this layer
if k_unet.startswith(f"input_blocks."):
# use scan layer static
ratio = cls.scan_layer(k_unet, "input_blocks", 12, **kwargs)
elif k_unet.startswith(f"middle_block."):
ratio = cls.scan_layer(k_unet, "middle_block", 3, **kwargs)
elif k_unet.startswith(f"output_blocks."):
ratio = cls.scan_layer(k_unet, "output_blocks", 12, **kwargs)
elif k_unet.startswith("out."):
ratio = kwargs.get("out", None)
elif k_unet.startswith("time"):
ratio = kwargs.get("time", None)
elif k_unet.startswith("label_emb"):
ratio = kwargs.get("label", None)
else:
print(f"Unknown key: {key}, skipping.")
return ratio
def apply_sparsification(self, base_model_param: Optional[torch.Tensor], model_a_param: torch.Tensor, model_b_param: torch.Tensor,
exclude_a: float, include_b: float, invert : str, drop_rate: float, ties : str, rescale : str,
device : torch.device, **kwargs) -> torch.Tensor:
"""
Applies sparsification to a tensor based on the specified sparsity level.
"""
model_a_flat = model_a_param.view(-1).float().to(device)
model_b_flat = model_b_param.view(-1).float().to(device)
delta_flat = model_b_flat - model_a_flat
if base_model_param is not None:
base_model_flat = base_model_param.view(-1).float().to(device)
delta_a_flat = model_a_flat - base_model_flat
delta_b_flat = model_b_flat - base_model_flat
include_mask = self.get_threshold_mask(delta_b_flat, include_b, invert, **kwargs)
exclude_mask = self.get_threshold_mask(delta_a_flat, exclude_a, invert, **kwargs)
base_mask = include_mask & (~exclude_mask)
del base_model_flat, delta_a_flat, delta_b_flat, include_mask, exclude_mask
else:
include_mask = torch.ones_like(model_a_flat).bool()
exclude_mask = torch.zeros_like(model_a_flat).bool()
base_mask = include_mask & (~exclude_mask)
del include_mask, exclude_mask
if ties != "off":
ties_mask = self.get_ties_mask(delta_flat, ties)
base_mask = base_mask & ties_mask
del ties_mask
dare_mask = torch.bernoulli(torch.full(delta_flat.shape, 1 - drop_rate, device=device)).bool()
# The paper says we should rescale, but it yields terrible results for SD
if rescale == "on":
# Rescale the remaining deltas
delta_flat = delta_flat / (1 - drop_rate)
final_mask = dare_mask & base_mask
# print(f"mask nonzero count: {torch.count_nonzero(mask)} dare nonzero count: {torch.count_nonzero(dare_mask)} base nonzero count: {torch.count_nonzero(base_mask)} include nonzero count: {torch.count_nonzero(include_mask)} exclude nonzero count: {torch.count_nonzero(exclude_mask)}")
sparsified_flat = torch.where(final_mask, model_a_flat + delta_flat, model_a_flat)
del final_mask, delta_flat, base_mask, model_a_flat, model_b_flat, dare_mask
return sparsified_flat.view_as(model_a_param)
+155
View File
@@ -0,0 +1,155 @@
# components/mask_model.py
from comfy.model_patcher import ModelPatcher
import torch
from typing import Dict, Tuple
from ..ddare.util import cuda_memory_profiler, get_device
from ..ddare.mask import ModelMask
from ..ddare.const import MASK_CATEGORY
class MagnitudeMasker:
"""
A managed state dict to allow for masking of model layers. This is used for protecting
layers from being overwritten by the merge process.
"""
CHUNK_SIZE = 10**7 # Constant chunk size for memory management
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the masking process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"threshold": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"threshold_type": (["median", "quantile"], {"default": "median"}),
"invert": (["No", "Yes"], {"default": "No"}),
}
}
RETURN_TYPES = ("MODEL_MASK",)
FUNCTION = "mask"
CATEGORY = MASK_CATEGORY
def mask(self, model_a: ModelPatcher, model_b: ModelPatcher, **kwargs) -> Tuple[ModelMask]:
"""
Uses two ModelPatcher instances to determine the deltas between the two models and then create a mask for the deltas
above and below a certain threshold.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the mask.
"""
device = get_device()
with cuda_memory_profiler():
if len(model_a.patches) > 0:
print("Model A has patches, applying them")
model_a.patch_model(None, True)
model_a_sd = model_a.model_state_dict()
model_a.unpatch_model()
else:
model_a_sd = model_a.model_state_dict()
if len(model_b.patches) > 0:
print("Model B has patches, applying them")
model_b.patch_model(None, True)
model_b_sd = model_b.model_state_dict()
model_b.unpatch_model()
else:
model_b_sd = model_b.model_state_dict()
mm = ModelMask({})
# Merge each parameter from model_b into model_a
for k in model_a_sd.keys():
if k not in model_b_sd:
print("could not patch. key doesn't exist in model:", k)
continue
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = model_b_sd[k]
layer_mask = self.get_threshold_mask(a, b, device=device, **kwargs)
mm.add_layer_mask(k, layer_mask)
return (mm,)
def get_threshold_mask(self, model_a_param: torch.Tensor, model_b_param: torch.Tensor, device : torch.device, threshold: float, invert: str, **kwargs) -> torch.Tensor:
"""
Gets a mask of the delta parameter based on the specified sparsity level.
Args:
model_a_param (torch.Tensor): The parameter from model_a.
model_b_param (torch.Tensor): The parameter from model_b.
device (torch.device): The device to use for the mask.
threshold (float): The sparsity level to use.
invert (str): Whether to invert the mask or not.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
torch.Tensor: The mask of the delta parameter.
"""
model_a_flat = model_a_param.view(-1).float().to(device)
model_b_flat = model_b_param.view(-1).float().to(device)
delta_flat = model_b_flat - model_a_flat
invertion = 1 if invert == 'No' else 0
if threshold == 1.0:
mask = torch.ones_like(delta_flat) == invertion
elif threshold == 0.0:
mask = torch.zeros_like(delta_flat) == invertion
else:
absolute_delta = torch.abs(delta_flat)
# We can easily overrun memory with large tensors, so we chunk the tensor
delta_threshold = self.process_in_chunks(tensor=absolute_delta, threshold=threshold, **kwargs)
# Create a mask for values to keep or preserve (above the threshold)
mask = absolute_delta >= delta_threshold if invert == 'No' else absolute_delta < delta_threshold
#print(f"Delta threshold: {delta_threshold} Mask: {absolute_delta.sum()} / {absolute_delta.numel()} invert: {invert} threshold: {threshold} Above: ({mask.sum()}) Below ({(mask == False).sum()})")
return mask.view_as(model_a_param)
def process_in_chunks(self, tensor: torch.Tensor, threshold: float, threshold_type : str, **kwargs) -> float:
"""
Processes the tensor in chunks and calculates the quantile thresholds for each chunk to determine our layer threshold.
Args:
tensor (torch.Tensor): The tensor to process.
threshold (float): The quantile threshold to use.
threshold_type (str): The type of threshold to use, either "median" or "quantile".
Returns:
float: The layer threshold.
"""
thresholds = []
for i in range(0, tensor.numel(), self.CHUNK_SIZE):
chunk = tensor[i:i + self.CHUNK_SIZE]
if chunk.numel() == 0:
continue
threshold = torch.quantile(torch.abs(chunk), threshold).item()
thresholds.append(threshold)
if threshold_type == "median":
global_threshold = torch.median(torch.tensor(thresholds))
else:
sorted_thresholds = sorted(thresholds)
index = int(threshold * len(sorted_thresholds))
index = max(0, min(index, len(sorted_thresholds) - 1))
global_threshold = sorted_thresholds[index]
return global_threshold
+234
View File
@@ -0,0 +1,234 @@
# components/normalize.py
import torch
from typing import Dict, Tuple
from comfy.model_patcher import ModelPatcher
from ..ddare.util import cuda_memory_profiler, get_device
from ..ddare.const import EPSILON, UTIL_CATEGORY
"""
These are the layers that we are going to normalize, and how we are going to normalize them:
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.0.in_layers.{0,2}.{weight, bias} scale by weight
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.0.emb_layers.{1}.{weight, bias} scale by weight
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.0.out_layers.{0,3}.{weight, bias} scale by weight
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.1.{norm, proj_in, proj_out}.{weight, bias} scale by weight
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.1.transformer_blocks.0.{attn1, attn2}.{to_q, to_k, to_v}.weight q, v scaled by q weight, k inverse scaled
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.1.transformer_blocks.0.{attn1, attn2}.to_out.0.{weight, bias} scale by weight
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.1.transformer_blocks.0.ff.net.{0,2}.proj.{weight, bias} scale by weight
diffusion_model.{{input_blocks, output_blocks}.{n}, {middle_block}}.1.transformer_blocks.0.{norm1, norm2, norm3}.{weight, bias} scale by weight
"""
class NormalizeUnet:
"""
A class to normalize the blocks from one model to the other, bringing them into the same scale.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"method": (["q_norm", "all", "none", "attn_only"], ),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "normalize"
CATEGORY = UTIL_CATEGORY
def generate_key_groups_sd15(self):
# Lets build our key selection strategy, and pick out the keys we want to patch from our key list
# This is currently pretty manual, but we can make it more automatic later. Right now we are just
# making sure it matches sd15.
blocks = []
for i in range(12):
blocks.append(f"input_blocks.{i}")
blocks.append(f"output_blocks.{i}")
blocks.append("middle_block")
key_groups = {}
prefix = "diffusion_model"
layers = [("0.in_layers", (0,2)), ("0.emb_layers", (1,)), ("0.out_layers", (0,3))]
layers += [("2.in_layers", (0,2)), ("2.emb_layers", (1,)), ("2.out_layers", (0,3))]
layers += [("1.norm", ()), ("1.proj_in", ()), ("1.proj_out", ())]
for b in blocks:
for lk, z in layers:
if len(z) == 0:
key_groups[f"{prefix}.{b}.{lk}.weight"] = (f"{prefix}.{b}.{lk}.weight", f"{prefix}.{b}.{lk}.bias")
else:
for i in z:
key_groups[f"{prefix}.{b}.{lk}.{i}.weight"] = (f"{prefix}.{b}.{lk}.{i}.weight", f"{prefix}.{b}.{lk}.{i}.bias")
key_groups[f"{prefix}.{b}.0.skip_connection.weight"] = (f"{prefix}.{b}.0.skip_connection.weight", f"{prefix}.{b}.0.skip_connection.bias")
key_groups[f"{prefix}.{b}.1.transformer_blocks.0.norm1.weight"] = (f"{prefix}.{b}.1.transformer_blocks.0.norm1.weight", f"{prefix}.{b}.1.transformer_blocks.0.norm1.bias")
key_groups[f"{prefix}.{b}.1.transformer_blocks.0.norm2.weight"] = (f"{prefix}.{b}.1.transformer_blocks.0.norm2.weight", f"{prefix}.{b}.1.transformer_blocks.0.norm2.bias")
key_groups[f"{prefix}.{b}.1.transformer_blocks.0.norm3.weight"] = (f"{prefix}.{b}.1.transformer_blocks.0.norm3.weight", f"{prefix}.{b}.1.transformer_blocks.0.norm3.bias")
key_groups[f"{prefix}.{b}.1.transformer_blocks.0.ff.net.0.proj.weight"] = (f"{prefix}.{b}.1.transformer_blocks.0.ff.net.0.proj.weight", f"{prefix}.{b}.1.transformer_blocks.0.ff.net.0.proj.bias")
key_groups[f"{prefix}.{b}.1.transformer_blocks.0.ff.net.2.weight"] = (f"{prefix}.{b}.1.transformer_blocks.0.ff.net.2.weight", f"{prefix}.{b}.1.transformer_blocks.0.ff.net.2.bias")
# Get our two attention blocks into our 5-tuple
key_groups[f"{prefix}.{b}.1.transformer_blocks.0.attn1.to_q.weight"] = (f"{prefix}.{b}.1.transformer_blocks.0.attn1.to_q.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn1.to_k.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn1.to_v.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn1.to_out.0.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn1.to_out.0.bias")
key_groups[f"{prefix}.{b}.1.transformer_blocks.0.attn2.to_q.weight"] = (f"{prefix}.{b}.1.transformer_blocks.0.attn2.to_q.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn2.to_k.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn2.to_v.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn2.to_out.0.weight", f"{prefix}.{b}.1.transformer_blocks.0.attn2.to_out.0.bias")
key_groups[f"{prefix}.{b}.2.conv.weight"] = (f"{prefix}.{b}.2.conv.weight", f"{prefix}.{b}.2.conv.bias")
key_groups[f"{prefix}.input_blocks.0.weight"] = (f"{prefix}.input_blocks.0.weight", f"{prefix}.input_blocks.0.bias")
key_groups[f"{prefix}.input_blocks.0.0.weight"] = (f"{prefix}.input_blocks.0.0.weight", f"{prefix}.input_blocks.0.0.bias")
key_groups[f"{prefix}.input_blocks.3.0.op.weight"] = (f"{prefix}.input_blocks.3.0.op.weight", f"{prefix}.input_blocks.3.0.op.bias")
key_groups[f"{prefix}.input_blocks.6.0.op.weight"] = (f"{prefix}.input_blocks.6.0.op.weight", f"{prefix}.input_blocks.6.0.op.bias")
key_groups[f"{prefix}.input_blocks.9.0.op.weight"] = (f"{prefix}.input_blocks.9.0.op.weight", f"{prefix}.input_blocks.9.0.op.bias")
key_groups[f"{prefix}.out.0.weight"] = (f"{prefix}.out.0.weight", f"{prefix}.out.0.bias")
key_groups[f"{prefix}.out.2.weight"] = (f"{prefix}.out.2.weight", f"{prefix}.out.2.bias")
key_groups[f"{prefix}.output_blocks.2.1.conv.weight"] = (f"{prefix}.output_blocks.2.1.conv.weight", f"{prefix}.output_blocks.2.1.conv.bias")
key_groups[f"{prefix}.time_embed.0.weight"] = (f"{prefix}.time_embed.0.weight", f"{prefix}.time_embed.0.bias")
key_groups[f"{prefix}.time_embed.2.weight"] = (f"{prefix}.time_embed.2.weight", f"{prefix}.time_embed.2.bias")
return key_groups
def normalize(self, model_a: ModelPatcher, model_b: ModelPatcher, method : str, **kwargs) -> Tuple[ModelPatcher]:
"""
Scales model A by the scaling factor calculated from model B.
Args:
model_a (ModelPatcher): Model to be scaled.
model_b (ModelPatcher): Model to be used as reference.
method (str): Method to be used for merging.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
device = get_device()
with cuda_memory_profiler():
m = model_a.clone() # Clone model_a to keep its structure
if method == "none":
if len(m.patches) > 0:
print(f"Model A has patches: {m.patches.keys()}")
return (m,)
if len(model_a.patches) > 0:
print("Model A has patches, applying them")
m.patch_model(None, True)
model_a_sd = m.model_state_dict() # State dict of model_a
m.unpatch_model() # Unpatch model_a
else:
model_a_sd = m.model_state_dict() # State dict of model_a
if len(model_b.patches) > 0:
print("Model B has patches, applying them")
model_b.patch_model(None, True)
model_b_sd = model_b.model_state_dict()
model_b.unpatch_model()
else:
model_b_sd = model_b.model_state_dict()
strength_patch = 1.0
strength_model = 0.0
processed_keys = {}
for k in sorted(model_a_sd.keys()):
processed_keys[k] = False
for key, group in self.generate_key_groups_sd15().items():
if key not in model_a_sd or key not in model_b_sd:
#print("could not patch. key doesn't exist in model:", key)
continue
elif len(group) == 2 and (method != "attn_only"):
# Normalize our weight and bias
weight_key, bias_key = group
weight_a : torch.Tensor = model_a_sd[weight_key].to(device)
weight_b : torch.Tensor = model_b_sd[weight_key].to(device)
bias_a : torch.Tensor = model_a_sd[bias_key].to(device)
scale = self._calculate_scaling_factor(weight_a, weight_b).to(device)
na = torch.empty_like(weight_a, device=device)
na = weight_a * scale
nb = torch.empty_like(weight_b, device=device)
nb = weight_b / scale
#print("normalized:", weight_key, scale)
del scale
m.add_patches({weight_key: (na.to('cpu'),)}, strength_patch, strength_model)
m.add_patches({bias_key: (nb.to('cpu'),)}, strength_patch, strength_model)
weight_a.to("cpu")
weight_b.to("cpu")
bias_a.to("cpu")
processed_keys[weight_key] = True
processed_keys[bias_key] = True
elif len(group) == 5 and (method == "q_norm" or method == "attn_only"):
# Scaled attention, we determine the scaling factor from the q weight
q, k, v, out_w, out_b = group
q_a : torch.Tensor = model_a_sd[q].to(device)
q_b : torch.Tensor = model_b_sd[q].to(device)
k_a : torch.Tensor = model_a_sd[k].to(device)
#v_a : torch.Tensor = model_a_sd[v].to(device)
out_w_a : torch.Tensor = model_a_sd[out_w].to(device)
out_w_b : torch.Tensor = model_b_sd[out_w].to(device)
out_b_a : torch.Tensor = model_a_sd[out_b].to(device)
scale_a = self._calculate_scaling_factor(q_a, q_b).to(device)
# allocate q, k, v, out_w, out_b
nq = torch.empty_like(q_a, device=device)
nq = q_a.to(device) * scale_a
nk = torch.empty_like(k_a, device=device)
nk = k_a.to(device) / scale_a
#v_a = v_a.copy_(v_a * scale).to("cpu")
scale_o = self._calculate_scaling_factor(out_w_a, out_w_b)
nout_w = torch.empty_like(out_w_a, device=device)
nout_w = out_w_a.to(device) * scale_o
nout_b = torch.empty_like(out_b_a, device=device)
nout_b = out_b_a.to(device) * scale_o
#print("normalized:", q, scale_a, scale_o)
del scale_a, scale_o
m.add_patches({q: (nq.to('cpu'),)}, strength_patch, strength_model)
m.add_patches({k: (nk.to('cpu'),)}, strength_patch, strength_model)
#m.add_patches({v: (v_a,)}, strength_patch, strength_model)
m.add_patches({out_w: (nout_w.to('cpu'),)}, strength_patch, strength_model)
m.add_patches({out_b: (nout_b.to('cpu'),)}, strength_patch, strength_model)
q_a.to("cpu")
q_b.to("cpu")
k_a.to("cpu")
#v_a.to("cpu")
out_w_a.to("cpu")
out_w_b.to("cpu")
out_b_a.to("cpu")
processed_keys[q] = True
processed_keys[k] = True
processed_keys[v] = True
processed_keys[out_w] = True
processed_keys[out_b] = True
for k, v in processed_keys.items():
if not v and method == "q_norm":
#print("key not processed:", k)
pass
return (m,)
@staticmethod
def _calculate_scaling_factor(weight_a: torch.Tensor, weight_b: torch.Tensor) -> float:
"""
Calculate the scaling factor to adjust the scale of weight_a to match weight_b.
Args:
weight_a (torch.Tensor): Weight tensor of this instance.
weight_b (torch.Tensor): Weight tensor of the other instance.
Returns:
float: Scaling factor.
"""
norm_a = torch.norm(weight_a)
norm_b = torch.norm(weight_b)
return norm_b / (norm_a + EPSILON) # Adding epsilon to avoid division by zero
View File
+7
View File
@@ -0,0 +1,7 @@
# ddare/const.py
EPSILON = 1e-8
MASK_CATEGORY="ddare/mask"
UNET_CATEGORY="ddare/unet"
CLIP_CATEGORY="ddare/clip"
UTIL_CATEGORY="ddare/util"
+24
View File
@@ -0,0 +1,24 @@
# ddare/mask.py
import torch
from typing import Dict, Optional
class ModelMask:
"""
A container to hold a state dict of masks for a model.
"""
def __init__(self, state_dict : Dict[str, torch.Tensor]):
self.state_dict = state_dict
def add_layer_mask(self, layer_name : str, mask : torch.Tensor):
self.state_dict[layer_name] = mask.clone().to("cpu")
def get_layer_mask(self, layer_name : str) -> Optional[torch.Tensor]:
if layer_name not in self.state_dict:
return None
return self.state_dict[layer_name]
def model_state_dict(self) -> Dict[str, torch.Tensor]:
return self.state_dict
+69 -16
View File
@@ -1,9 +1,9 @@
# merge/tensormerge.py
# ddare/merge.py
# Credit to https://github.com/Gryphe/MergeMonster
from comfy.model_patcher import ModelPatcher
import torch
from typing import Optional
from typing import Optional, Literal
from .const import EPSILON
def merge_tensors(method: str, v0: torch.Tensor, v1: torch.Tensor, t: float) -> torch.Tensor:
if method == "lerp":
@@ -24,7 +24,7 @@ def merge_tensors_lerp(v0: torch.Tensor, v1: torch.Tensor, t: float) -> torch.Te
return result
def merge_tensors_slerp(v0: torch.Tensor, v1: torch.Tensor, t: float, dot_threshold: float = 0.9995, eps: float = 1e-8) -> torch.Tensor:
def merge_tensors_slerp(v0: torch.Tensor, v1: torch.Tensor, t: float, dot_threshold: float = 0.9995, eps: float = EPSILON) -> torch.Tensor:
"""Spherical linear interpolation between two tensors or linear interpolation if they are one-dimensional.
Full credit to https://github.com/cg123/mergekit for the original code."""
@@ -171,21 +171,74 @@ def merge_tensors_gradient(v0: torch.Tensor, v1: torch.Tensor, t: float) -> torc
else:
return v0
def safe_normalize(tensor: torch.Tensor, eps: float):
def safe_normalize(tensor: torch.Tensor, eps: float = EPSILON):
norm = tensor.norm()
if norm > eps:
return tensor / norm
return tensor
def patcher(model: ModelPatcher, key : str) -> Optional[torch.Tensor]:
# This is slow, but seems to work
model_sd = model.model_state_dict()
if key not in model_sd:
print("could not patch. key doesn't exist in model:", key)
return None
def get_ties_mask(delta: torch.Tensor, method: Literal["sum", "count"] = "sum", mask_dtype: Optional[torch.dtype] = None, **kwargs) -> torch.Tensor:
"""
TIES-merging https://arxiv.org/abs/2306.01708 uses sign agreement, protecting from
major perturbations in the opposite direction of the base model
Returns a mask determining which delta vectors should be merged
into the final model.
weight : torch.Tensor = model_sd[key]
For the methodology described in the paper use 'sum'. For a
simpler naive count of signs, use 'count'.
"""
if mask_dtype is None:
mask_dtype = delta.dtype
temp_weight = weight.to(torch.float32, copy=True)
out_weight = model.calculate_weight(model.patches[key], temp_weight, key).to(weight.dtype)
return out_weight
sign = delta.sign().to(mask_dtype)
if method == "sum":
sign_weight = (sign * delta.abs()).sum(dim=0)
majority_sign = (sign_weight >= 0).to(mask_dtype) * 2 - 1
del sign_weight
elif method == "count":
majority_sign = (sign.sum(dim=0) >= 0).to(mask_dtype) * 2 - 1
else:
raise RuntimeError(f'Unimplemented mask method "{method}"')
return sign == majority_sign
def dare_ties_sparsification(model_a_param: torch.Tensor, model_b_param: torch.Tensor,
drop_rate: float, ties : str, rescale : str, device : torch.device,
**kwargs) -> torch.Tensor:
"""
DARE-TIES sparsification uses a stochastic mask to determine which deltas to apply
and then sign-agreement to determine which deltas to merge into the final model.
Args:
model_a_param (torch.Tensor): The base model parameter tensor.
model_b_param (torch.Tensor): The model parameter tensor to merge into the base model.
drop_rate (float): The drop rate for the stochastic mask.
ties (str): Whether to use the TIES-merging method.
rescale (str): Whether to rescale the remaining deltas.
device (torch.device): The device to use for the merge.
Returns:
torch.Tensor: The updated parameter tensor.
"""
model_a_flat = model_a_param.view(-1).float().to(device)
model_b_flat = model_b_param.view(-1).float().to(device)
delta_flat = model_b_flat - model_a_flat
dare_mask = torch.bernoulli(torch.full(delta_flat.shape, 1 - drop_rate, device=device)).bool()
# The paper says we should rescale, but it yields terrible results for SD.
if rescale == "on":
# Rescale the remaining deltas
delta_flat = delta_flat / (1 - drop_rate)
if ties != "off":
ties_mask = get_ties_mask(delta_flat, ties)
dare_mask = dare_mask & ties_mask
del ties_mask
sparsified_flat = torch.where(dare_mask, model_a_flat + delta_flat, model_a_flat)
del delta_flat, model_a_flat, model_b_flat, dare_mask
return sparsified_flat.view_as(model_a_param)
+59
View File
@@ -0,0 +1,59 @@
# merge/util.py
from comfy.model_patcher import ModelPatcher
import contextlib
import torch
from typing import Optional, Dict
def patcher(model: ModelPatcher, key : str) -> Optional[torch.Tensor]:
# This is slow, but seems to work
model_sd = model.model_state_dict()
if key not in model_sd:
print("could not patch. key doesn't exist in model:", key)
return None
weight : torch.Tensor = model_sd[key]
temp_weight = weight.to(torch.float32, copy=True)
out_weight = model.calculate_weight(model.patches[key], temp_weight, key).to(weight.dtype)
return out_weight
@contextlib.contextmanager
def cuda_memory_profiler():
"""
A context manager for profiling CUDA memory usage in PyTorch.
"""
torch.cuda.reset_peak_memory_stats()
torch.cuda.synchronize()
start_memory = torch.cuda.memory_allocated()
try:
yield
finally:
torch.cuda.synchronize()
end_memory = torch.cuda.memory_allocated()
print(f"Peak memory usage: {torch.cuda.max_memory_allocated() / (1024 ** 2):.2f} MB")
print(f"Memory allocated at start: {start_memory / (1024 ** 2):.2f} MB")
print(f"Memory allocated at end: {end_memory / (1024 ** 2):.2f} MB")
print(f"Net memory change: {(end_memory - start_memory) / (1024 ** 2):.2f} MB")
def get_device():
return torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu")
def get_patched_state(model : ModelPatcher) -> Dict[str, torch.Tensor]:
"""Uses a Comfy ModelPatcher to get the patched state dict of a model.
Args:
model (ModelPatcher): The model to get the patched state dict from.
Returns:
Dict[str, torch.Tensor]: The patched state dict.
"""
if len(model.patches) > 0:
print("Model has patches, applying them")
model.patch_model(None, True)
model_sd = model.model_state_dict()
model.unpatch_model()
else:
model_sd = model.model_state_dict()
return model_sd
-132
View File
@@ -1,132 +0,0 @@
# merge/block.py
import torch
from typing import Dict, Tuple, Optional
from comfy.model_patcher import ModelPatcher
from .mergeutil import merge_tensors
class BlockModelMergerAdv:
"""
A class to merge two diffusion U-Net models using calculated deltas, sparsification,
and a weighted consensus method. This is the Magnitude Pruning method.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"input": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"middle": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"out": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"method": (["lerp", "slerp", "gradient"], ),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = "ddare/block"
def merge(self, model_a: ModelPatcher, model_b: ModelPatcher, method : str, input : float, middle : float, out : float, time : float, **kwargs) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
input (float): The ratio (lambda) of the input layer to keep from model_a.
middle (float): The ratio (lambda) of the middle layers to keep from model_a.
out (float): The ratio (lambda) of the output layer to keep from model_a.
time (float): The ratio (lambda) of the time layer to keep from model_a.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
m = model_a.clone() # Clone model_a to keep its structure
model_a_sd = m.model_state_dict() # State dict of model_a
kp = model_b.get_key_patches("diffusion_model.") # Get the key patches from model_b
# Merge each parameter from model_b into model_a
for k in kp:
if k not in model_a_sd:
print("could not patch. key doesn't exist in model:", k)
continue
k_unet = k[len("diffusion_model."):]
# Get our ratio for this layer
if k_unet.startswith("input"):
ratio = input
elif k_unet.startswith("middle"):
ratio = middle
elif k_unet.startswith("out"):
ratio = out
elif k_unet.startswith("time"):
ratio = time
else:
print(f"Unknown key: {k}, skipping.")
continue
# Apply sparsification by the delta, I don't know if all of this cuda stuff is necessary
# but I had so many memory issues that I'm being very careful
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = kp[k][-1]
# Debugging
# our 'Tensor's might be a tuple sometimes if it's part of a chain. This logic is very hacky and could be flawed.
# typer = lambda x: type(x) if not isinstance(x, tuple) else [typer(y) for y in x]
if isinstance(a, tuple):
#print('chain', a[0], a[-1], len(a), typer(a))
a = self.patcher(model_a, k)
if a is None:
continue
else:
a = a.copy_(a)
if isinstance(b, tuple):
#print('chain', b[0], b[-1], len(b), typer(b))
b = self.patcher(model_b, k)
if b is None:
continue
else:
b = b.copy_(b)
merged_layer = merge_tensors(method, a.to(device), b.to(device), 1 - ratio)
nv = (merged_layer.to('cpu'),)
del a, b
# We have already merged the models
m.add_patches({k: nv}, 1, 0)
return (m,)
def patcher(self, model: ModelPatcher, key : str) -> Optional[torch.Tensor]:
# This is slow, but seems to work
model_sd = model.model_state_dict()
if key not in model_sd:
print("could not patch. key doesn't exist in model:", key)
return None
weight : torch.Tensor = model_sd[key]
temp_weight = weight.to(torch.float32, copy=True)
out_weight = model.calculate_weight(model.patches[key], temp_weight, key).to(weight.dtype)
return out_weight
-304
View File
@@ -1,304 +0,0 @@
# merge/dare.py
from comfy.model_patcher import ModelPatcher
import torch
from typing import Dict, Tuple, Optional, Literal
from .mergeutil import merge_tensors, patcher
class DareModelMerger:
"""
A class to merge two diffusion U-Net models using calculated deltas, sparsification,
and a weighted consensus method. This is the DARE method.
https://arxiv.org/pdf/2311.03099.pdf
"""
CHUNK_SIZE = 10**7 # Constant chunk size for memory management
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"drop_rate": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01}),
"ties": (["sum", "count", "off"], {"default": "sum"}),
"rescale": (["off", "on"], {"default": "off"}),
"seed": ("INT", {"default": 42}),
"input": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"middle": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"out": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"method": (["comfy", "lerp", "slerp", "gradient"], ),
"exclude_a": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"include_b": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"threshold_type": (["median", "quantile"], {"default": "median"}),
"invert": (["No", "Yes"], {"default": "No"}),
"iterations": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
},
"optional": {
"base_model": ("MODEL",),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = "ddare/dare"
def merge(self, model_a: ModelPatcher, model_b: ModelPatcher,
input: float, middle: float, out: float, time: float, method : str,
seed : Optional[int] = None, clear_cache : bool = True,
base_model: Optional[ModelPatcher] = None, iterations : int = 1,
**kwargs) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
input (float): The ratio (lambda) of the input layer to keep from model_a.
middle (float): The ratio (lambda) of the middle layers to keep from model_a.
out (float): The ratio (lambda) of the output layer to keep from model_a.
time (float): The ratio (lambda) of the time layers to keep from model_a.
method (str): The method to use for merging, either "lerp", "slerp", or "gradient".
seed (int): The random seed to use for the merge.
clear_cache (bool): Whether to clear the CUDA cache after each chunk. Default is False.
base_model (ModelPatcher): The base model to use for calculating the deltas. Optional.
iterations (int): The number of iterations to perform the merge. Default is 1.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
if seed is not None:
torch.manual_seed(seed)
device = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu")
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
m = model_a.clone() # Clone model_a to keep its structure
model_a_sd = m.model_state_dict() # State dict of model_a
kp = model_b.get_key_patches("diffusion_model.") # Get the key patches from model_b
if base_model is not None:
model_base_sd = base_model.model_state_dict() # State dict of base model
else:
model_base_sd = None
# Merge each parameter from model_b into model_a
for k in kp:
if k not in model_a_sd:
print("could not patch. key doesn't exist in model:", k)
continue
k_unet = k[len("diffusion_model."):]
# Get our ratio for this layer
if k_unet.startswith("input"):
ratio = input
elif k_unet.startswith("middle"):
ratio = middle
elif k_unet.startswith("out"):
ratio = out
elif k_unet.startswith("time"):
ratio = time
else:
print(f"Unknown key: {k}, skipping.")
continue
# Apply sparsification by the delta, I don't know if all of this cuda stuff is necessary
# but I had so many memory issues that I'm being very careful
base : torch.Tensor = model_base_sd[k] if model_base_sd is not None else None
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = kp[k][-1]
# Debugging
# our 'Tensor's might be a tuple sometimes if it's part of a chain. This logic is very hacky and could be flawed.
# typer = lambda x: type(x) if not isinstance(x, tuple) else [typer(y) for y in x]
if isinstance(base, tuple):
#print('chain', a[0], a[-1], len(a), typer(a))
base = patcher(base_model, k)
if base is None:
continue
elif base is not None:
base = base.copy_(base)
if isinstance(a, tuple):
#print('chain', a[0], a[-1], len(a), typer(a))
a = patcher(model_a, k)
if a is None:
continue
else:
a = a.copy_(a)
if isinstance(b, tuple):
#print('chain', b[0], b[-1], len(b), typer(b))
b = patcher(model_b, k)
if b is None:
continue
else:
b = b.copy_(b)
merged_a = a
for i in range(iterations):
sparsified_delta = self.apply_sparsification(base, merged_a, b, device=device, **kwargs)
if method == "comfy":
merged_a = sparsified_delta
strength_patch = 1.0 - ratio
strength_model = ratio
else:
merged_a = merge_tensors(method, merged_a.to(device), sparsified_delta.to(device), 1 - ratio)
strength_model = 0
strength_patch = 1.0
del base, a, b
# Apply the sparsified delta as a patch
nv = (merged_a.to('cpu'),)
m.add_patches({k: nv}, strength_patch, strength_model)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return (m,)
def apply_sparsification(self, base_model_param: Optional[torch.Tensor], model_a_param: torch.Tensor, model_b_param: torch.Tensor,
exclude_a: float, include_b: float, invert : str, drop_rate: float, ties : str, rescale : str,
device : torch.device, **kwargs) -> torch.Tensor:
"""
Applies sparsification to a tensor based on the specified sparsity level.
"""
model_a_flat = model_a_param.view(-1).float().to(device)
model_b_flat = model_b_param.view(-1).float().to(device)
delta_flat = model_b_flat - model_a_flat
if base_model_param is not None:
base_model_flat = base_model_param.view(-1).float().to(device)
delta_a_flat = model_a_flat - base_model_flat
delta_b_flat = model_b_flat - base_model_flat
include_mask = self.get_threshold_mask(delta_b_flat, include_b, invert, **kwargs)
exclude_mask = self.get_threshold_mask(delta_a_flat, exclude_a, invert, **kwargs)
base_mask = include_mask & (~exclude_mask)
del base_model_flat, delta_a_flat, delta_b_flat, include_mask, exclude_mask
else:
include_mask = torch.ones_like(model_a_flat).bool()
exclude_mask = torch.zeros_like(model_a_flat).bool()
base_mask = include_mask & (~exclude_mask)
del include_mask, exclude_mask
if ties != "off":
ties_mask = self.get_ties_mask(delta_flat, ties)
base_mask = base_mask & ties_mask
del ties_mask
dare_mask = torch.bernoulli(torch.full(delta_flat.shape, 1 - drop_rate, device=device)).bool()
# The paper says we should rescale, but it yields terrible results for SD
if rescale == "on":
# Rescale the remaining deltas
delta_flat = delta_flat / (1 - drop_rate)
final_mask = dare_mask & base_mask
# print(f"mask nonzero count: {torch.count_nonzero(mask)} dare nonzero count: {torch.count_nonzero(dare_mask)} base nonzero count: {torch.count_nonzero(base_mask)} include nonzero count: {torch.count_nonzero(include_mask)} exclude nonzero count: {torch.count_nonzero(exclude_mask)}")
sparsified_flat = torch.where(final_mask, model_a_flat + delta_flat, model_a_flat)
del final_mask, delta_flat, base_mask, model_a_flat, model_b_flat, dare_mask
return sparsified_flat.view_as(model_a_param)
def get_ties_mask(self, delta: torch.Tensor, method: Literal["sum", "count"] = "sum", mask_dtype: Optional[torch.dtype] = None, **kwargs) -> torch.Tensor:
"""
TIES-merging https://arxiv.org/abs/2306.01708 uses sign agreement, protecting from
major perturbations in the opposite direction of the base model
Returns a mask determining which delta vectors should be merged
into the final model.
For the methodology described in the paper use 'sum'. For a
simpler naive count of signs, use 'count'.
"""
if mask_dtype is None:
mask_dtype = delta.dtype
sign = delta.sign().to(mask_dtype)
if method == "sum":
sign_weight = (sign * delta.abs()).sum(dim=0)
majority_sign = (sign_weight >= 0).to(mask_dtype) * 2 - 1
del sign_weight
elif method == "count":
majority_sign = (sign.sum(dim=0) >= 0).to(mask_dtype) * 2 - 1
else:
raise RuntimeError(f'Unimplemented mask method "{method}"')
return sign == majority_sign
def process_in_chunks(self, tensor: torch.Tensor, sparsity: float, threshold_type : str, **kwargs) -> torch.Tensor:
"""
Processes the tensor in chunks and calculates the quantile thresholds for each chunk.
"""
thresholds = []
for i in range(0, tensor.numel(), self.CHUNK_SIZE):
chunk = tensor[i:i + self.CHUNK_SIZE]
if chunk.numel() == 0:
continue
threshold = torch.quantile(torch.abs(chunk), sparsity).item()
thresholds.append(threshold)
if threshold_type == "median":
global_threshold = torch.median(torch.tensor(thresholds))
else:
sorted_thresholds = sorted(thresholds)
index = int(sparsity * len(sorted_thresholds))
index = max(0, min(index, len(sorted_thresholds) - 1))
global_threshold = sorted_thresholds[index]
return global_threshold
def get_threshold_mask(self, delta_param: torch.Tensor, sparsity: float, invert: str, **kwargs) -> torch.Tensor:
"""
Gets a mask of the delta parameter based on the specified sparsity level.
Args:
delta_param (torch.Tensor): The delta parameter tensor.
sparsity (float): The fraction of elements to set to zero. 0 = include all, 1 = exclude all.
invert (str): Whether to invert the sparsification, i.e., keep the least significant changes.
Returns:
torch.Tensor: The mask of the delta parameter.
"""
invertion = 1 if invert == 'No' else 0
if sparsity == 1.0:
return torch.ones_like(delta_param) == invertion
elif sparsity == 0.0:
return torch.zeros_like(delta_param) == invertion
absolute_delta = torch.abs(delta_param)
# We can easily overrun memory with large tensors, so we chunk the tensor
delta_threshold = self.process_in_chunks(tensor=absolute_delta, sparsity=sparsity, **kwargs)
print(f"Delta threshold: {delta_threshold} Mask: {absolute_delta.sum()} / {absolute_delta.numel()} invert: {invert} sparsity: {sparsity}")
# Create a mask for values to keep or preserve (above the threshold)
mask = absolute_delta >= delta_threshold if invert == 'No' else absolute_delta < delta_threshold
return mask
-331
View File
@@ -1,331 +0,0 @@
# merge/dare.py
from comfy.model_patcher import ModelPatcher
import torch
from typing import Dict, Tuple, Optional, Literal
from .mergeutil import merge_tensors, patcher
class DareModelMergerMBW:
"""
A class to merge two diffusion U-Net models using calculated deltas, sparsification,
and a weighted consensus method. This is the DARE method.
https://arxiv.org/pdf/2311.03099.pdf
"""
CHUNK_SIZE = 10**7 # Constant chunk size for memory management
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
arg_dict = {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"drop_rate": ("FLOAT", {"default": 0.85, "min": 0.0, "max": 1.0, "step": 0.01}),
"ties": (["sum", "count", "off"], {"default": "sum"}),
"rescale": (["off", "on"], {"default": "on"}),
"seed": ("INT", {"default": 1, "min":0, "max": 99999999999}),
"method": (["comfy", "lerp", "slerp", "gradient"], ),
"exclude_a": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"include_b": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"threshold_type": (["median", "quantile"], {"default": "median"}),
"invert": (["No", "Yes"], {"default": "No"}),
"iterations": ("INT", {"default": 1, "min": 1, "max": 100, "step": 1}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"label": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
}
argument = ("FLOAT", {"default": 0.7, "min": 0.0, "max": 1.0, "step": 0.01})
for i in range(12):
arg_dict[f"input_blocks.{i}"] = argument
for i in range(3):
arg_dict[f"middle_block.{i}"] = argument
for i in range(12):
arg_dict[f"output_blocks.{i}"] = argument
arg_dict["out"] = argument
opt = {"base_model": ("MODEL",)}
return {"required": arg_dict ,"optional": opt}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = "ddare/dareMBW"
def merge(self, model_a: ModelPatcher, model_b: ModelPatcher,
time: float, label: float, method : str,
seed : Optional[int] = None, clear_cache : bool = True,
base_model: Optional[ModelPatcher] = None, iterations : int = 1,
**kwargs,) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
input (float): The ratio (lambda) of the input layer to keep from model_a.
middle (float): The ratio (lambda) of the middle layers to keep from model_a.
out (float): The ratio (lambda) of the output layer to keep from model_a.
time (float): The ratio (lambda) of the time layers to keep from model_a.
method (str): The method to use for merging, either "lerp", "slerp", or "gradient".
seed (int): The random seed to use for the merge.
clear_cache (bool): Whether to clear the CUDA cache after each chunk. Default is False.
base_model (ModelPatcher): The base model to use for calculating the deltas. Optional.
iterations (int): The number of iterations to perform the merge. Default is 1.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
if seed is not None:
torch.manual_seed(seed)
device = torch.device("cuda" if torch.cuda.is_available() else "mps" if torch.backends.mps.is_available() else "cpu")
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
m = model_a.clone() # Clone model_a to keep its structure
model_a_sd = m.model_state_dict() # State dict of model_a
kp = model_b.get_key_patches("diffusion_model.") # Get the key patches from model_b
if base_model is not None:
model_base_sd = base_model.model_state_dict() # State dict of base model
else:
model_base_sd = None
# Merge each parameter from model_b into model_a
for k in kp:
if k not in model_a_sd:
print("could not patch. key doesn't exist in model:", k)
continue
k_unet = k[len("diffusion_model."):]
# Get our ratio for this layer
#デバッグ用にごちゃごちゃしちゃったけど各break以降は画像に影響しないよ
if k_unet.startswith(f"input_blocks"):
for i in range(12):
if k_unet.startswith(f"input_blocks.{i}"):
ratio = kwargs[f"input_blocks.{i}"]
break
elif i==11:
print(f"Unknown key: {k_unet},i={i}")
continue
elif k_unet.startswith(f"middle_block"):
for i in range(3):
if k_unet.startswith(f"middle_block.{i}"):
ratio = kwargs[f"middle_block.{i}"]
break
elif i==2:
print(f"Unknown key: {k_unet},i={i}")
continue
elif k_unet.startswith(f"output_blocks"):
for i in range(12):
if k_unet.startswith(f"output_blocks.{i}"):
ratio = kwargs[f"output_blocks.{i}"]
break
elif i==11:
print(f"Unknown key: {k_unet},i={i}")
continue
elif k_unet.startswith("out."):
ratio = kwargs["out"]
elif k_unet.startswith("time"):
ratio = time
elif k_unet.startswith("label_emb"):
ratio = label
else:
print(f"Unknown key: {k}, skipping.")
continue
# Apply sparsification by the delta, I don't know if all of this cuda stuff is necessary
# but I had so many memory issues that I'm being very careful
base : torch.Tensor = model_base_sd[k] if model_base_sd is not None else None
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = kp[k][-1]
# Debugging
# our 'Tensor's might be a tuple sometimes if it's part of a chain. This logic is very hacky and could be flawed.
# typer = lambda x: type(x) if not isinstance(x, tuple) else [typer(y) for y in x]
if isinstance(base, tuple):
#print('chain', a[0], a[-1], len(a), typer(a))
base = patcher(base_model, k)
if base is None:
continue
elif base is not None:
base = base.copy_(base)
if isinstance(a, tuple):
#print('chain', a[0], a[-1], len(a), typer(a))
a = patcher(model_a, k)
if a is None:
continue
else:
a = a.copy_(a)
if isinstance(b, tuple):
#print('chain', b[0], b[-1], len(b), typer(b))
b = patcher(model_b, k)
if b is None:
continue
else:
b = b.copy_(b)
merged_a = a
for i in range(iterations):
sparsified_delta = self.apply_sparsification(base, merged_a, b, device=device, **kwargs)
if method == "comfy":
merged_a = sparsified_delta
strength_patch = 1.0 - ratio
strength_model = ratio
else:
merged_a = merge_tensors(method, merged_a.to(device), sparsified_delta.to(device), 1 - ratio)
strength_model = 0
strength_patch = 1.0
del base, a, b
# Apply the sparsified delta as a patch
nv = (merged_a.to('cpu'),)
m.add_patches({k: nv}, strength_patch, strength_model)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return (m,)
def apply_sparsification(self, base_model_param: Optional[torch.Tensor], model_a_param: torch.Tensor, model_b_param: torch.Tensor,
exclude_a: float, include_b: float, invert : str, drop_rate: float, ties : str, rescale : str,
device : torch.device, **kwargs) -> torch.Tensor:
"""
Applies sparsification to a tensor based on the specified sparsity level.
"""
model_a_flat = model_a_param.view(-1).float().to(device)
model_b_flat = model_b_param.view(-1).float().to(device)
delta_flat = model_b_flat - model_a_flat
if base_model_param is not None:
base_model_flat = base_model_param.view(-1).float().to(device)
delta_a_flat = model_a_flat - base_model_flat
delta_b_flat = model_b_flat - base_model_flat
include_mask = self.get_threshold_mask(delta_b_flat, include_b, invert, **kwargs)
exclude_mask = self.get_threshold_mask(delta_a_flat, exclude_a, invert, **kwargs)
base_mask = include_mask & (~exclude_mask)
del base_model_flat, delta_a_flat, delta_b_flat, include_mask, exclude_mask
else:
include_mask = torch.ones_like(model_a_flat).bool()
exclude_mask = torch.zeros_like(model_a_flat).bool()
base_mask = include_mask & (~exclude_mask)
del include_mask, exclude_mask
if ties != "off":
ties_mask = self.get_ties_mask(delta_flat, ties)
base_mask = base_mask & ties_mask
del ties_mask
dare_mask = torch.bernoulli(torch.full(delta_flat.shape, 1 - drop_rate, device=device)).bool()
# The paper says we should rescale, but it yields terrible results for SD
if rescale == "on":
# Rescale the remaining deltas
delta_flat = delta_flat / (1 - drop_rate)
final_mask = dare_mask & base_mask
# print(f"mask nonzero count: {torch.count_nonzero(mask)} dare nonzero count: {torch.count_nonzero(dare_mask)} base nonzero count: {torch.count_nonzero(base_mask)} include nonzero count: {torch.count_nonzero(include_mask)} exclude nonzero count: {torch.count_nonzero(exclude_mask)}")
sparsified_flat = torch.where(final_mask, model_a_flat + delta_flat, model_a_flat)
del final_mask, delta_flat, base_mask, model_a_flat, model_b_flat, dare_mask
return sparsified_flat.view_as(model_a_param)
def get_ties_mask(self, delta: torch.Tensor, method: Literal["sum", "count"] = "sum", mask_dtype: Optional[torch.dtype] = None, **kwargs) -> torch.Tensor:
"""
TIES-merging https://arxiv.org/abs/2306.01708 uses sign agreement, protecting from
major perturbations in the opposite direction of the base model
Returns a mask determining which delta vectors should be merged
into the final model.
For the methodology described in the paper use 'sum'. For a
simpler naive count of signs, use 'count'.
"""
if mask_dtype is None:
mask_dtype = delta.dtype
sign = delta.sign().to(mask_dtype)
if method == "sum":
sign_weight = (sign * delta.abs()).sum(dim=0)
majority_sign = (sign_weight >= 0).to(mask_dtype) * 2 - 1
del sign_weight
elif method == "count":
majority_sign = (sign.sum(dim=0) >= 0).to(mask_dtype) * 2 - 1
else:
raise RuntimeError(f'Unimplemented mask method "{method}"')
return sign == majority_sign
def process_in_chunks(self, tensor: torch.Tensor, sparsity: float, threshold_type : str, **kwargs) -> torch.Tensor:
"""
Processes the tensor in chunks and calculates the quantile thresholds for each chunk.
"""
thresholds = []
for i in range(0, tensor.numel(), self.CHUNK_SIZE):
chunk = tensor[i:i + self.CHUNK_SIZE]
if chunk.numel() == 0:
continue
threshold = torch.quantile(torch.abs(chunk), sparsity).item()
thresholds.append(threshold)
if threshold_type == "median":
global_threshold = torch.median(torch.tensor(thresholds))
else:
sorted_thresholds = sorted(thresholds)
index = int(sparsity * len(sorted_thresholds))
index = max(0, min(index, len(sorted_thresholds) - 1))
global_threshold = sorted_thresholds[index]
return global_threshold
def get_threshold_mask(self, delta_param: torch.Tensor, sparsity: float, invert: str, **kwargs) -> torch.Tensor:
"""
Gets a mask of the delta parameter based on the specified sparsity level.
Args:
delta_param (torch.Tensor): The delta parameter tensor.
sparsity (float): The fraction of elements to set to zero. 0 = include all, 1 = exclude all.
invert (str): Whether to invert the sparsification, i.e., keep the least significant changes.
Returns:
torch.Tensor: The mask of the delta parameter.
"""
invertion = 1 if invert == 'No' else 0
if sparsity == 1.0:
return torch.ones_like(delta_param) == invertion
elif sparsity == 0.0:
return torch.zeros_like(delta_param) == invertion
absolute_delta = torch.abs(delta_param)
# We can easily overrun memory with large tensors, so we chunk the tensor
delta_threshold = self.process_in_chunks(tensor=absolute_delta, sparsity=sparsity, **kwargs)
print(f"Delta threshold: {delta_threshold} Mask: {absolute_delta.sum()} / {absolute_delta.numel()} invert: {invert} sparsity: {sparsity}")
# Create a mask for values to keep or preserve (above the threshold)
mask = absolute_delta >= delta_threshold if invert == 'No' else absolute_delta < delta_threshold
return mask
-202
View File
@@ -1,202 +0,0 @@
# merge/mag.py
from comfy.model_patcher import ModelPatcher
import torch
from typing import Dict, Tuple
from .mergeutil import merge_tensors, patcher
class MagnitudePruningModelMerger:
"""
A class to merge two diffusion U-Net models using calculated deltas, sparsification,
and a weighted consensus method. This is the Magnitude Pruning method.
"""
CHUNK_SIZE = 10**7 # Constant chunk size for memory management
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model_a": ("MODEL",),
"model_b": ("MODEL",),
"input": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"middle": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"out": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"method": (["comfy", "lerp", "slerp", "gradient"], ),
"density": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01}),
"threshold_type": (["median", "quantile"], ),
"invert": (["No", "Yes"], ),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = "ddare/magnitude_pruning"
def merge(self, model_a: ModelPatcher, model_b: ModelPatcher,
input : float, middle : float, out : float, time : float, method : str,
clear_cache : bool = True,
**kwargs) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model_a (ModelPatcher): The base model to be merged.
model_b (ModelPatcher): The model to merge into the base model.
input (float): The ratio (lambda) of the input layer to keep from model_a.
middle (float): The ratio (lambda) of the middle layers to keep from model_a.
out (float): The ratio (lambda) of the output layer to keep from model_a.
time (float): The ratio (lambda) of the time layers to keep from model_a.
method (str): The method to use for merging, either "lerp", "slerp", or "gradient".
clear_cache (bool): Whether to clear the CUDA cache after each chunk. Default is True.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
m = model_a.clone() # Clone model_a to keep its structure
model_a_sd = m.model_state_dict() # State dict of model_a
kp = model_b.get_key_patches("diffusion_model.") # Get the key patches from model_b
# Merge each parameter from model_b into model_a
for k in kp:
if k not in model_a_sd:
print("could not patch. key doesn't exist in model:", k)
continue
k_unet = k[len("diffusion_model."):]
# Get our ratio for this layer
if k_unet.startswith("input"):
ratio = input
elif k_unet.startswith("middle"):
ratio = middle
elif k_unet.startswith("out"):
ratio = out
elif k_unet.startswith("time"):
ratio = time
else:
print(f"Unknown key: {k}, skipping.")
continue
# Apply sparsification by the delta, I don't know if all of this cuda stuff is necessary
# but I had so many memory issues that I'm being very careful
a : torch.Tensor = model_a_sd[k]
b : torch.Tensor = kp[k][-1]
# Debugging
# our 'Tensor's might be a tuple sometimes if it's part of a chain. This logic is very hacky and could be flawed.
# typer = lambda x: type(x) if not isinstance(x, tuple) else [typer(y) for y in x]
if isinstance(a, tuple):
#print('chain', a[0], a[-1], len(a), typer(a))
a = patcher(model_a, k)
if a is None:
continue
else:
a = a.copy_(a)
if isinstance(b, tuple):
#print('chain', b[0], b[-1], len(b), typer(b))
b = patcher(model_b, k)
if b is None:
continue
else:
b = b.copy_(b)
sparsified_delta = self.apply_sparsification(a, b, device=device, **kwargs)
if method == "comfy":
merged_layer = sparsified_delta
strength_patch = 1.0 - ratio
strength_model = ratio
else:
merged_layer = merge_tensors(method, a.to(device), sparsified_delta.to(device), 1 - ratio)
strength_model = 0
strength_patch = 1.0
del a, b
# Apply the sparsified delta as a patch
nv = (merged_layer.to('cpu'),)
m.add_patches({k: nv}, strength_patch, strength_model)
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return (m,)
def apply_sparsification(self, base_param: torch.Tensor, target_param: torch.Tensor, density: float,
threshold_type : str, invert : str, device : torch.device,
**kwargs) -> torch.Tensor:
"""
Applies sparsification to a tensor based on the specified sparsity level, with chunking for large tensors.
Args:
base_param (torch.Tensor): The corresponding parameter from the base model.
target_param (torch.Tensor): The corresponding parameter from the update model.
density (float): The fraction of elements to keep from the second model. 1 is keep all.
threshold_type (str): The type of threshold to use, either "median" or "quantile".
invert (str): Whether to invert the sparsification, i.e., keep the least significant changes.
Returns:
torch.Tensor: The tensor with insignificant changes replaced by the base model's values.
"""
# Ensure the delta and base_param are float tensors for quantile calculation, and on the right device
target_param = target_param.to(device)
base_param = base_param.to(device)
delta = target_param - base_param
base_param_flat = base_param.view(-1).float()
delta_flat = delta.view(-1).float().to(device)
absolute_delta = torch.abs(delta_flat)
# We can easily overrun memory with large tensors, so we chunk the tensor
# Define chunk size and prepare to collect thresholds
chunk_size = 10**7
thresholds = []
# Process each chunk to determine thresholds
for i in range(0, absolute_delta.numel(), chunk_size):
chunk = absolute_delta[i:i + chunk_size]
if chunk.numel() == 0:
continue
k = int(density * chunk.numel())
if k > 0:
threshold = torch.quantile(chunk, density)
else:
threshold = torch.tensor(0.0)
thresholds.append(threshold)
# Determine a global threshold
if threshold_type == "median":
global_threshold = torch.median(torch.tensor(thresholds))
else:
sorted_thresholds = sorted(thresholds)
index = int(density * len(sorted_thresholds))
index = max(0, min(index, len(sorted_thresholds) - 1))
global_threshold = sorted_thresholds[index]
# Create a mask for values to keep (above the threshold)
mask = absolute_delta >= global_threshold if invert == 'No' else absolute_delta < global_threshold
# Apply the mask to the delta, replace other values with the base model's parameters
sparsified_flat = torch.where(mask, base_param_flat, base_param_flat + delta_flat)
del mask, absolute_delta, delta_flat, base_param_flat, global_threshold, thresholds
return sparsified_flat.view_as(base_param)
-190
View File
@@ -1,190 +0,0 @@
# magmerge.py
import torch
from typing import Dict, Tuple, Optional
from comfy.model_patcher import ModelPatcher
class MagnitudePruningModelMerger:
"""
A class to merge two diffusion U-Net models using calculated deltas, sparsification,
and a weighted consensus method. This is the Magnitude Pruning method.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, tuple]:
"""
Defines the input types for the merging process.
Returns:
Dict[str, tuple]: A dictionary specifying the required model types and parameters.
"""
return {
"required": {
"model1": ("MODEL",),
"model2": ("MODEL",),
"input": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"middle": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"out": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"time": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
"sparsity": ("FLOAT", {"default": 0.9, "min": 0.0, "max": 1.0, "step": 0.01}),
"threshold_type": (["median", "quantile"], ),
"invert": (["No", "Yes"], ),
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "merge"
CATEGORY = "ddare/model_merging"
def apply_sparsification(self, base_param: torch.Tensor, target_param: torch.Tensor, sparsity: float,
threshold_type : str, invert : str, clear_cache : bool = False, **kwargs) -> torch.Tensor:
"""
Applies sparsification to a tensor based on the specified sparsity level, with chunking for large tensors.
Args:
base_param (torch.Tensor): The corresponding parameter from the base model.
target_param (torch.Tensor): The corresponding parameter from the update model.
sparsity (float): The fraction of elements to set to zero.
threshold_type (str): The type of threshold to use, either "median" or "quantile".
invert (str): Whether to invert the sparsification, i.e., keep the least significant changes.
clear_cache (bool): Whether to clear the CUDA cache after each chunk. Default is False.
Returns:
torch.Tensor: The tensor with insignificant changes replaced by the base model's values.
"""
# Ensure the delta and base_param are float tensors for quantile calculation, and on the right device
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
target_param = target_param.to(device)
base_param = base_param.to(device)
delta = target_param - base_param
base_param_flat = base_param.view(-1).float()
delta_flat = delta.view(-1).float().to(device)
absolute_delta = torch.abs(delta_flat)
# We can easily overrun memory with large tensors, so we chunk the tensor
# Define chunk size and prepare to collect thresholds
chunk_size = 10**7
thresholds = []
# Process each chunk to determine thresholds
for i in range(0, absolute_delta.numel(), chunk_size):
chunk = absolute_delta[i:i + chunk_size]
if chunk.numel() == 0:
continue
k = int(sparsity * chunk.numel())
if k > 0:
threshold = torch.quantile(chunk, sparsity)
else:
threshold = torch.tensor(0.0)
thresholds.append(threshold)
# Determine a global threshold
if threshold_type == "median":
global_threshold = torch.median(torch.tensor(thresholds))
else:
sorted_thresholds = sorted(thresholds)
index = int(sparsity * len(sorted_thresholds))
index = max(0, min(index, len(sorted_thresholds) - 1))
global_threshold = sorted_thresholds[index]
# Create a mask for values to keep (above the threshold)
mask = absolute_delta >= global_threshold if invert == 'No' else absolute_delta < global_threshold
# Apply the mask to the delta, replace other values with the base model's parameters
sparsified_flat = torch.where(mask, base_param_flat, base_param_flat + delta_flat)
del mask, absolute_delta, delta_flat, base_param_flat, global_threshold, thresholds
if clear_cache and torch.cuda.is_available():
torch.cuda.empty_cache()
return sparsified_flat.view_as(base_param).to('cpu')
def patcher(self, model: ModelPatcher, key : str) -> Optional[torch.Tensor]:
# This is slow, but seems to work
model_sd = model.model_state_dict()
if key not in model_sd:
print("could not patch. key doesn't exist in model:", key)
return None
weight : torch.Tensor = model_sd[key]
temp_weight = weight.to(torch.float32, copy=True)
out_weight = model.calculate_weight(model.patches[key], temp_weight, key).to(weight.dtype)
return out_weight
def merge(self, model1: ModelPatcher, model2: ModelPatcher, input : float, middle : float, out : float, time : float, **kwargs) -> Tuple[ModelPatcher]:
"""
Merges two ModelPatcher instances based on the weighted consensus of their parameters and sparsity.
Args:
model1 (ModelPatcher): The base model to be merged.
model2 (ModelPatcher): The model to merge into the base model.
input (float): The ratio (lambda) of the input layer to keep from model1.
middle (float): The ratio (lambda) of the middle layers to keep from model1.
out (float): The ratio (lambda) of the output layer to keep from model1.
**kwargs: Additional arguments specifying the merge ratios for different layers and sparsity.
Returns:
Tuple[ModelPatcher]: A tuple containing the merged ModelPatcher instance.
"""
m = model1.clone() # Clone model1 to keep its structure
model1_sd = m.model_state_dict() # State dict of model1
kp1 = model2.get_key_patches("diffusion_model.") # Get the key patches from model1
kp = model2.get_key_patches("diffusion_model.") # Get the key patches from model2
# Merge each parameter from model2 into model1
for k in kp:
if k not in model1_sd:
continue
k_unet = k[len("diffusion_model."):]
# Get our ratio for this layer
if k_unet.startswith("input"):
ratio = input
elif k_unet.startswith("middle"):
ratio = middle
elif k_unet.startswith("out"):
ratio = out
elif k_unet.startswith("time"):
ratio = time
else:
print(f"Unknown key: {k}, skipping.")
ratio = 1.0
# Apply sparsification by the delta, I don't know if all of this cuda stuff is necessary
# but I had so many memory issues that I'm being very careful
a : torch.Tensor = model1_sd[k]
b : torch.Tensor = kp[k][-1]
# Debugging
# our 'Tensor's might be a tuple sometimes if it's part of a chain. This logic is very hacky and could be flawed.
# typer = lambda x: type(x) if not isinstance(x, tuple) else [typer(y) for y in x]
if isinstance(a, tuple):
#print('chain', a[0], a[-1], len(a), typer(a))
a = self.patcher(model1, k)
if a is None:
continue
else:
a = a.copy_(a)
if isinstance(b, tuple):
#print('chain', b[0], b[-1], len(b), typer(b))
b = self.patcher(model2, k)
if b is None:
continue
else:
b = b.copy_(b)
sparsified_delta = self.apply_sparsification(a, b, **kwargs)
nv = (sparsified_delta,)
del a, b
# Apply the sparsified delta as a patch
strength_patch = 1.0 - ratio
strength_model = ratio
m.add_patches({k: nv}, strength_patch, strength_model)
return (m,)