mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-02 14:27:08 +08:00
Use custom fwd/bwd func and better init for loha
This commit is contained in:
parent
c025308918
commit
28a236d505
@ -6,6 +6,74 @@ import comfy.model_management
|
|||||||
from .base import WeightAdapterBase, WeightAdapterTrainBase, weight_decompose
|
from .base import WeightAdapterBase, WeightAdapterTrainBase, weight_decompose
|
||||||
|
|
||||||
|
|
||||||
|
class HadaWeight(torch.autograd.Function):
|
||||||
|
@staticmethod
|
||||||
|
def forward(ctx, w1u, w1d, w2u, w2d, scale=torch.tensor(1)):
|
||||||
|
ctx.save_for_backward(w1d, w1u, w2d, w2u, scale)
|
||||||
|
diff_weight = ((w1u @ w1d) * (w2u @ w2d)) * scale
|
||||||
|
return diff_weight
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def backward(ctx, grad_out):
|
||||||
|
(w1d, w1u, w2d, w2u, scale) = ctx.saved_tensors
|
||||||
|
grad_out = grad_out * scale
|
||||||
|
temp = grad_out * (w2u @ w2d)
|
||||||
|
grad_w1u = temp @ w1d.T
|
||||||
|
grad_w1d = w1u.T @ temp
|
||||||
|
|
||||||
|
temp = grad_out * (w1u @ w1d)
|
||||||
|
grad_w2u = temp @ w2d.T
|
||||||
|
grad_w2d = w2u.T @ temp
|
||||||
|
|
||||||
|
del temp
|
||||||
|
return grad_w1u, grad_w1d, grad_w2u, grad_w2d, None
|
||||||
|
|
||||||
|
|
||||||
|
class HadaWeightTucker(torch.autograd.Function):
|
||||||
|
@staticmethod
|
||||||
|
def forward(ctx, t1, w1u, w1d, t2, w2u, w2d, scale=torch.tensor(1)):
|
||||||
|
ctx.save_for_backward(t1, w1d, w1u, t2, w2d, w2u, scale)
|
||||||
|
|
||||||
|
rebuild1 = torch.einsum("i j ..., j r, i p -> p r ...", t1, w1d, w1u)
|
||||||
|
rebuild2 = torch.einsum("i j ..., j r, i p -> p r ...", t2, w2d, w2u)
|
||||||
|
|
||||||
|
return rebuild1 * rebuild2 * scale
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def backward(ctx, grad_out):
|
||||||
|
(t1, w1d, w1u, t2, w2d, w2u, scale) = ctx.saved_tensors
|
||||||
|
grad_out = grad_out * scale
|
||||||
|
|
||||||
|
temp = torch.einsum("i j ..., j r -> i r ...", t2, w2d)
|
||||||
|
rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w2u)
|
||||||
|
|
||||||
|
grad_w = rebuild * grad_out
|
||||||
|
del rebuild
|
||||||
|
|
||||||
|
grad_w1u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w)
|
||||||
|
grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w1u.T)
|
||||||
|
del grad_w, temp
|
||||||
|
|
||||||
|
grad_w1d = torch.einsum("i r ..., i j ... -> r j", t1, grad_temp)
|
||||||
|
grad_t1 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w1d.T)
|
||||||
|
del grad_temp
|
||||||
|
|
||||||
|
temp = torch.einsum("i j ..., j r -> i r ...", t1, w1d)
|
||||||
|
rebuild = torch.einsum("i j ..., i r -> r j ...", temp, w1u)
|
||||||
|
|
||||||
|
grad_w = rebuild * grad_out
|
||||||
|
del rebuild
|
||||||
|
|
||||||
|
grad_w2u = torch.einsum("r j ..., i j ... -> r i", temp, grad_w)
|
||||||
|
grad_temp = torch.einsum("i j ..., i r -> r j ...", grad_w, w2u.T)
|
||||||
|
del grad_w, temp
|
||||||
|
|
||||||
|
grad_w2d = torch.einsum("i r ..., i j ... -> r j", t2, grad_temp)
|
||||||
|
grad_t2 = torch.einsum("i j ..., j r -> i r ...", grad_temp, w2d.T)
|
||||||
|
del grad_temp
|
||||||
|
return grad_t1, grad_w1u, grad_w1d, grad_t2, grad_w2u, grad_w2d, None
|
||||||
|
|
||||||
|
|
||||||
class LohaDiff(WeightAdapterTrainBase):
|
class LohaDiff(WeightAdapterTrainBase):
|
||||||
def __init__(self, weights):
|
def __init__(self, weights):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
@ -36,24 +104,17 @@ class LohaDiff(WeightAdapterTrainBase):
|
|||||||
def __call__(self, w):
|
def __call__(self, w):
|
||||||
org_dtype = w.dtype
|
org_dtype = w.dtype
|
||||||
|
|
||||||
|
# Apply scaling
|
||||||
|
scale = self.alpha / self.rank
|
||||||
# Reconstruct the two matrices m1 and m2
|
# Reconstruct the two matrices m1 and m2
|
||||||
if self.use_tucker:
|
if self.use_tucker:
|
||||||
# CP/Tucker decomposition case
|
# CP/Tucker decomposition case
|
||||||
m1 = torch.einsum('i j k l, j r, i p -> p r k l', self.t1, self.w1b, self.w1a)
|
diff_weight = HadaWeightTucker.apply(self.t1, self.w1_a, self.w1_b, self.t2, self.w2_a, self.w2_b, scale)
|
||||||
m2 = torch.einsum('i j k l, j r, i p -> p r k l', self.t2, self.w2b, self.w2a)
|
|
||||||
else:
|
else:
|
||||||
# Standard Hadmard product case
|
diff_weight = HadaWeight.apply(self.w1_a, self.w1_b, self.w2_a, self.w2_b, scale)
|
||||||
m1 = self.w1a @ self.w1b
|
|
||||||
m2 = self.w2a @ self.w2b
|
|
||||||
|
|
||||||
# Calculate the final difference via element-wise product
|
|
||||||
diff = m1 * m2
|
|
||||||
|
|
||||||
# Apply scaling
|
|
||||||
scale = self.alpha / self.rank
|
|
||||||
|
|
||||||
# Add the scaled difference to the original weight
|
# Add the scaled difference to the original weight
|
||||||
weight = w + scale * diff.reshape(w.shape)
|
weight = w + diff_weight.reshape(w.shape).to(w.dtype)
|
||||||
|
|
||||||
return weight.to(org_dtype)
|
return weight.to(org_dtype)
|
||||||
|
|
||||||
@ -75,12 +136,12 @@ class LoHaAdapter(WeightAdapterBase):
|
|||||||
in_dim = weight.shape[1:].numel()
|
in_dim = weight.shape[1:].numel()
|
||||||
mat1 = torch.empty(out_dim, rank, device=weight.device, dtype=weight.dtype)
|
mat1 = torch.empty(out_dim, rank, device=weight.device, dtype=weight.dtype)
|
||||||
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=weight.dtype)
|
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=weight.dtype)
|
||||||
torch.nn.init.kaiming_uniform_(mat1, a=5**0.5)
|
torch.nn.init.normal_(mat1, 0.1)
|
||||||
torch.nn.init.constant_(mat2, 0.0)
|
torch.nn.init.constant_(mat2, 0.0)
|
||||||
mat3 = torch.empty(out_dim, rank, device=weight.device, dtype=weight.dtype)
|
mat3 = torch.empty(out_dim, rank, device=weight.device, dtype=weight.dtype)
|
||||||
mat4 = torch.empty(rank, in_dim, device=weight.device, dtype=weight.dtype)
|
mat4 = torch.empty(rank, in_dim, device=weight.device, dtype=weight.dtype)
|
||||||
torch.nn.init.kaiming_uniform_(mat1, a=5**0.5)
|
torch.nn.init.normal_(mat1, 1)
|
||||||
torch.nn.init.kaiming_uniform_(mat2, a=5**0.5)
|
torch.nn.init.normal_(mat2, 0.1)
|
||||||
return LohaDiff(
|
return LohaDiff(
|
||||||
(mat1, mat2, alpha, mat3, mat4, None, None, None)
|
(mat1, mat2, alpha, mat3, mat4, None, None, None)
|
||||||
)
|
)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user