From e0c97e8495041ef5c30658ec8d596f39e58db750 Mon Sep 17 00:00:00 2001 From: Martin Bukowski Date: Mon, 22 Jan 2024 14:31:05 -0600 Subject: [PATCH] cleanup --- components/clip.py | 3 +- components/dare.py | 3 +- components/dare_mbw.py | 3 +- components/normalize.py | 22 ++--------- ddare/merge.py | 69 +------------------------------- ddare/tensor.py | 87 +++++++++++++++++++++++++++++++++++++++++ 6 files changed, 97 insertions(+), 90 deletions(-) create mode 100644 ddare/tensor.py diff --git a/components/clip.py b/components/clip.py index 4d41193..c4d4a2e 100644 --- a/components/clip.py +++ b/components/clip.py @@ -3,7 +3,8 @@ from comfy.sd import CLIP import torch from typing import Optional -from ..ddare.merge import merge_tensors, dare_ties_sparsification +from ..ddare.merge import merge_tensors +from ..ddare.tensor import dare_ties_sparsification from ..ddare.util import cuda_memory_profiler, get_device from ..ddare.const import CLIP_CATEGORY diff --git a/components/dare.py b/components/dare.py index 100c427..77592d1 100644 --- a/components/dare.py +++ b/components/dare.py @@ -3,7 +3,8 @@ 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.merge import merge_tensors +from ..ddare.tensor import 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 diff --git a/components/dare_mbw.py b/components/dare_mbw.py index 1ccb69e..ab267d9 100644 --- a/components/dare_mbw.py +++ b/components/dare_mbw.py @@ -3,7 +3,8 @@ 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.merge import merge_tensors +from ..ddare.tensor import 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 diff --git a/components/normalize.py b/components/normalize.py index 55c1cc7..ada704e 100644 --- a/components/normalize.py +++ b/components/normalize.py @@ -4,7 +4,8 @@ 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 +from ..ddare.tensor import relative_norm +from ..ddare.const import UTIL_CATEGORY """ These are the layers that we are going to normalize, and how we are going to normalize them: @@ -145,7 +146,7 @@ class NormalizeUnet: 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) + scale = relative_norm(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) @@ -215,20 +216,3 @@ class NormalizeUnet: 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 - \ No newline at end of file diff --git a/ddare/merge.py b/ddare/merge.py index c70843a..08ffdba 100644 --- a/ddare/merge.py +++ b/ddare/merge.py @@ -1,7 +1,6 @@ # ddare/merge.py # Credit to https://github.com/Gryphe/MergeMonster import torch -from typing import Optional, Literal from .const import EPSILON @@ -175,70 +174,4 @@ def safe_normalize(tensor: torch.Tensor, eps: float = EPSILON): norm = tensor.norm() if norm > eps: return tensor / norm - return tensor - -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. - - 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 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) + return tensor \ No newline at end of file diff --git a/ddare/tensor.py b/ddare/tensor.py new file mode 100644 index 0000000..741d12c --- /dev/null +++ b/ddare/tensor.py @@ -0,0 +1,87 @@ +# ddare/tensor.py + +import torch +from typing import Optional, Literal + +from .const import EPSILON + +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. + + 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 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) + +def relative_norm(weight_a: torch.Tensor, weight_b: torch.Tensor, eps : float = EPSILON) -> float: + """ + Calculate the relative norm of two weight tensors. + + 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 + eps) # Adding epsilon to avoid division by zero