mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-04 01:57:08 +08:00
Correct dtype handling and better default arg
This commit is contained in:
parent
218c3e3793
commit
31c8cc9aa3
@ -64,7 +64,8 @@ class BiasDiff(torch.nn.Module):
|
|||||||
self.bias = bias
|
self.bias = bias
|
||||||
|
|
||||||
def __call__(self, b):
|
def __call__(self, b):
|
||||||
return b + self.bias
|
org_dtype = b.dtype
|
||||||
|
return (b.to(self.bias) + self.bias).to(org_dtype)
|
||||||
|
|
||||||
def passive_memory_usage(self):
|
def passive_memory_usage(self):
|
||||||
return self.bias.nelement() * self.bias.element_size()
|
return self.bias.nelement() * self.bias.element_size()
|
||||||
@ -285,19 +286,19 @@ class TrainLoraNode:
|
|||||||
"steps": (
|
"steps": (
|
||||||
IO.INT,
|
IO.INT,
|
||||||
{
|
{
|
||||||
"default": 50,
|
"default": 16,
|
||||||
"min": 1,
|
"min": 1,
|
||||||
"max": 1000,
|
"max": 100000,
|
||||||
"tooltip": "The number of steps to train the LoRA for.",
|
"tooltip": "The number of steps to train the LoRA for.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
"learning_rate": (
|
"learning_rate": (
|
||||||
IO.FLOAT,
|
IO.FLOAT,
|
||||||
{
|
{
|
||||||
"default": 0.0003,
|
"default": 0.0005,
|
||||||
"min": 0.0000001,
|
"min": 0.0000001,
|
||||||
"max": 1.0,
|
"max": 1.0,
|
||||||
"step": 0.00001,
|
"step": 0.000001,
|
||||||
"tooltip": "The learning rate to use for training.",
|
"tooltip": "The learning rate to use for training.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
@ -311,9 +312,9 @@ class TrainLoraNode:
|
|||||||
},
|
},
|
||||||
),
|
),
|
||||||
"optimizer": (
|
"optimizer": (
|
||||||
["Adam", "AdamW", "SGD", "RMSprop"],
|
["AdamW", "Adam", "SGD", "RMSprop"],
|
||||||
{
|
{
|
||||||
"default": "Adam",
|
"default": "AdamW",
|
||||||
"tooltip": "The optimizer to use for training.",
|
"tooltip": "The optimizer to use for training.",
|
||||||
},
|
},
|
||||||
),
|
),
|
||||||
@ -434,7 +435,7 @@ class TrainLoraNode:
|
|||||||
else:
|
else:
|
||||||
diff = torch.nn.Parameter(
|
diff = torch.nn.Parameter(
|
||||||
torch.zeros(
|
torch.zeros(
|
||||||
m.weight.shape, dtype=dtype, requires_grad=True
|
m.weight.shape, dtype=lora_dtype, requires_grad=True
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
diff_module = BiasDiff(diff)
|
diff_module = BiasDiff(diff)
|
||||||
@ -444,7 +445,7 @@ class TrainLoraNode:
|
|||||||
if hasattr(m, "bias") and m.bias is not None:
|
if hasattr(m, "bias") and m.bias is not None:
|
||||||
key = "{}.bias".format(n)
|
key = "{}.bias".format(n)
|
||||||
bias = torch.nn.Parameter(
|
bias = torch.nn.Parameter(
|
||||||
torch.zeros(m.bias.shape, dtype=dtype, requires_grad=True)
|
torch.zeros(m.bias.shape, dtype=lora_dtype, requires_grad=True)
|
||||||
)
|
)
|
||||||
bias_module = BiasDiff(bias)
|
bias_module = BiasDiff(bias)
|
||||||
lora_sd["{}.diff_b".format(n)] = bias
|
lora_sd["{}.diff_b".format(n)] = bias
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user