mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-04 08:07:05 +08:00
Add scheme of TrainBase class
This commit is contained in:
parent
e8f3bc5ab7
commit
14c20852c5
@ -33,10 +33,22 @@ class WeightAdapterBase:
|
|||||||
|
|
||||||
|
|
||||||
class WeightAdapterTrainBase(nn.Module):
|
class WeightAdapterTrainBase(nn.Module):
|
||||||
|
# We follow the scheme of PR #7032
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
|
||||||
# [TODO] Collaborate with LoRA training PR #7032
|
def __call__(self, w):
|
||||||
|
"""
|
||||||
|
w: The original weight tensor to be modified.
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def passive_memory_usage(self):
|
||||||
|
raise NotImplementedError("passive_memory_usage is not implemented")
|
||||||
|
|
||||||
|
def move_to(self, device):
|
||||||
|
self.to(device)
|
||||||
|
return self.passive_memory_usage()
|
||||||
|
|
||||||
|
|
||||||
def weight_decompose(dora_scale, weight, lora_diff, alpha, strength, intermediate_dtype, function):
|
def weight_decompose(dora_scale, weight, lora_diff, alpha, strength, intermediate_dtype, function):
|
||||||
@ -92,3 +104,14 @@ def pad_tensor_to_shape(tensor: torch.Tensor, new_shape: list[int]) -> torch.Ten
|
|||||||
padded_tensor[new_slices] = tensor[orig_slices]
|
padded_tensor[new_slices] = tensor[orig_slices]
|
||||||
|
|
||||||
return padded_tensor
|
return padded_tensor
|
||||||
|
|
||||||
|
|
||||||
|
def tucker_weight_from_conv(up, down, mid):
|
||||||
|
up = up.reshape(up.size(0), up.size(1))
|
||||||
|
down = down.reshape(down.size(0), down.size(1))
|
||||||
|
return torch.einsum("m n ..., i m, n j -> i j ...", mid, up, down)
|
||||||
|
|
||||||
|
|
||||||
|
def tucker_weight(wa, wb, t):
|
||||||
|
temp = torch.einsum("i j ..., j r -> i r ...", t, wb)
|
||||||
|
return torch.einsum("i j ..., i r -> r j ...", temp, wa)
|
||||||
Loading…
x
Reference in New Issue
Block a user