mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-15 04:16:39 +08:00
Fix adapter weight init
This commit is contained in:
parent
7d593baf91
commit
eea608e48f
@ -130,12 +130,12 @@ class LoHaAdapter(WeightAdapterBase):
|
|||||||
def create_train(cls, weight, rank=1, alpha=1.0):
|
def create_train(cls, weight, rank=1, alpha=1.0):
|
||||||
out_dim = weight.shape[0]
|
out_dim = weight.shape[0]
|
||||||
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=torch.float32)
|
||||||
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=weight.dtype)
|
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=torch.float32)
|
||||||
torch.nn.init.normal_(mat1, 0.1)
|
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=torch.float32)
|
||||||
mat4 = torch.empty(rank, in_dim, device=weight.device, dtype=weight.dtype)
|
mat4 = torch.empty(rank, in_dim, device=weight.device, dtype=torch.float32)
|
||||||
torch.nn.init.normal_(mat3, 0.1)
|
torch.nn.init.normal_(mat3, 0.1)
|
||||||
torch.nn.init.normal_(mat4, 0.01)
|
torch.nn.init.normal_(mat4, 0.01)
|
||||||
return LohaDiff(
|
return LohaDiff(
|
||||||
|
|||||||
@ -89,8 +89,8 @@ class LoKrAdapter(WeightAdapterBase):
|
|||||||
in_dim = weight.shape[1:].numel()
|
in_dim = weight.shape[1:].numel()
|
||||||
out1, out2 = factorization(out_dim, rank)
|
out1, out2 = factorization(out_dim, rank)
|
||||||
in1, in2 = factorization(in_dim, rank)
|
in1, in2 = factorization(in_dim, rank)
|
||||||
mat1 = torch.empty(out1, in1, device=weight.device, dtype=weight.dtype)
|
mat1 = torch.empty(out1, in1, device=weight.device, dtype=torch.float32)
|
||||||
mat2 = torch.empty(out2, in2, device=weight.device, dtype=weight.dtype)
|
mat2 = torch.empty(out2, in2, device=weight.device, dtype=torch.float32)
|
||||||
torch.nn.init.kaiming_uniform_(mat2, a=5**0.5)
|
torch.nn.init.kaiming_uniform_(mat2, a=5**0.5)
|
||||||
torch.nn.init.constant_(mat1, 0.0)
|
torch.nn.init.constant_(mat1, 0.0)
|
||||||
return LokrDiff(
|
return LokrDiff(
|
||||||
|
|||||||
@ -66,8 +66,8 @@ class LoRAAdapter(WeightAdapterBase):
|
|||||||
def create_train(cls, weight, rank=1, alpha=1.0):
|
def create_train(cls, weight, rank=1, alpha=1.0):
|
||||||
out_dim = weight.shape[0]
|
out_dim = weight.shape[0]
|
||||||
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=torch.float32)
|
||||||
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=weight.dtype)
|
mat2 = torch.empty(rank, in_dim, device=weight.device, dtype=torch.float32)
|
||||||
torch.nn.init.kaiming_uniform_(mat1, a=5**0.5)
|
torch.nn.init.kaiming_uniform_(mat1, a=5**0.5)
|
||||||
torch.nn.init.constant_(mat2, 0.0)
|
torch.nn.init.constant_(mat2, 0.0)
|
||||||
return LoraDiff(
|
return LoraDiff(
|
||||||
|
|||||||
@ -68,7 +68,7 @@ class OFTAdapter(WeightAdapterBase):
|
|||||||
def create_train(cls, weight, rank=1, alpha=1.0):
|
def create_train(cls, weight, rank=1, alpha=1.0):
|
||||||
out_dim = weight.shape[0]
|
out_dim = weight.shape[0]
|
||||||
block_size, block_num = factorization(out_dim, rank)
|
block_size, block_num = factorization(out_dim, rank)
|
||||||
block = torch.zeros(block_num, block_size, block_size, device=weight.device, dtype=weight.dtype)
|
block = torch.zeros(block_num, block_size, block_size, device=weight.device, dtype=torch.float32)
|
||||||
return OFTDiff(
|
return OFTDiff(
|
||||||
(block, None, alpha, None)
|
(block, None, alpha, None)
|
||||||
)
|
)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user