new nodes
This commit is contained in:
@@ -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
|
||||
@@ -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
@@ -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']
|
||||
|
||||
@@ -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,)
|
||||
@@ -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,)
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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,)
|
||||
Reference in New Issue
Block a user