Update workflows, fix IS_CHANGED

This commit is contained in:
kijai
2024-11-07 09:20:20 +02:00
parent 1ab2f3f4d6
commit 3c3a303668
3 changed files with 2316 additions and 1981 deletions
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+11 -7
View File
@@ -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):