cleanup
This commit is contained in:
+2
-1
@@ -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
|
||||
|
||||
|
||||
+2
-1
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+3
-19
@@ -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
|
||||
|
||||
+1
-68
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user