Update workflows, fix IS_CHANGED
This commit is contained in:
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -250,6 +250,7 @@ class ADMD_InitializeTraining:
|
|||||||
[
|
[
|
||||||
'Lion',
|
'Lion',
|
||||||
'AdamW',
|
'AdamW',
|
||||||
|
'prodigy'
|
||||||
], {
|
], {
|
||||||
"default": 'Lion'
|
"default": 'Lion'
|
||||||
}),
|
}),
|
||||||
@@ -336,6 +337,14 @@ class ADMD_InitializeTraining:
|
|||||||
if optimization_method == "AdamW":
|
if optimization_method == "AdamW":
|
||||||
print("Using AdamW optimizer for training")
|
print("Using AdamW optimizer for training")
|
||||||
optimizer = torch.optim.AdamW
|
optimizer = torch.optim.AdamW
|
||||||
|
elif optimization_method == "Prodigy":
|
||||||
|
try:
|
||||||
|
import prodigyopt
|
||||||
|
except ImportError:
|
||||||
|
raise ImportError("Prodigy not installed")
|
||||||
|
|
||||||
|
print(f"use Prodigy optimizer")
|
||||||
|
optimizer = prodigyopt.Prodigy
|
||||||
else:
|
else:
|
||||||
print("Using Lion optimizer for training")
|
print("Using Lion optimizer for training")
|
||||||
optimizer = Lion
|
optimizer = Lion
|
||||||
@@ -404,12 +413,7 @@ class ADMD_InitializeTraining:
|
|||||||
lr_scheduler_spatial_list.append(lr_scheduler_spatial)
|
lr_scheduler_spatial_list.append(lr_scheduler_spatial)
|
||||||
|
|
||||||
# Support mixed-precision training
|
# Support mixed-precision training
|
||||||
if 'scaler' not in globals():
|
|
||||||
scaler = torch.cuda.amp.GradScaler()
|
scaler = torch.cuda.amp.GradScaler()
|
||||||
print("initialize scaler")
|
|
||||||
else:
|
|
||||||
scaler.reset()
|
|
||||||
print("reset scaler")
|
|
||||||
|
|
||||||
admd_pipeline = {
|
admd_pipeline = {
|
||||||
"optimizer_temporal": optimizer_temporal,
|
"optimizer_temporal": optimizer_temporal,
|
||||||
@@ -584,7 +588,7 @@ class ADMD_DiffusersLoader:
|
|||||||
class ADMD_CheckpointLoader:
|
class ADMD_CheckpointLoader:
|
||||||
@classmethod
|
@classmethod
|
||||||
def IS_CHANGED(s):
|
def IS_CHANGED(s):
|
||||||
return ""
|
return float("nan")
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
|
|
||||||
@@ -713,7 +717,7 @@ class ADMD_CheckpointLoader:
|
|||||||
class ADMD_ComfyModelLoader:
|
class ADMD_ComfyModelLoader:
|
||||||
@classmethod
|
@classmethod
|
||||||
def IS_CHANGED(s):
|
def IS_CHANGED(s):
|
||||||
return ""
|
return float("nan")
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(cls):
|
def INPUT_TYPES(cls):
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user