Add schedulerfree optimizers
This commit is contained in:
@@ -362,8 +362,14 @@ class FluxTrainer:
|
||||
|
||||
logger.info(f"using {len(optimizers)} optimizers for blockwise fused optimizers")
|
||||
|
||||
if train_util.is_schedulefree_optimizer(optimizers[0], args):
|
||||
raise ValueError("Schedule-free optimizer is not supported with blockwise fused optimizers")
|
||||
self.optimizer_train_fn = lambda: None # dummy function
|
||||
self.optimizer_eval_fn = lambda: None # dummy function
|
||||
|
||||
else:
|
||||
_, _, optimizer = train_util.get_optimizer(args, trainable_params=params_to_optimize)
|
||||
self.optimizer_train_fn, self.optimizer_eval_fn = train_util.get_optimizer_train_eval_fn(optimizer, args)
|
||||
|
||||
# prepare dataloader
|
||||
# strategies are set here because they cannot be referenced in another process. Copy them with the dataset
|
||||
@@ -783,6 +789,7 @@ class FluxTrainer:
|
||||
if accelerator.sync_gradients:
|
||||
progress_bar.update(1)
|
||||
self.global_step += 1
|
||||
|
||||
|
||||
# flux_train_utils.sample_images(
|
||||
# accelerator, args, None, global_step, flux, ae, [clip_l, t5xxl], sample_prompts_te_outputs
|
||||
|
||||
+53
-8
@@ -13,6 +13,7 @@ import shutil
|
||||
import time
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
Dict,
|
||||
List,
|
||||
NamedTuple,
|
||||
@@ -2417,7 +2418,7 @@ def is_disk_cached_latents_is_expected(reso, npz_path: str, flip_aug: bool, alph
|
||||
if alpha_mask:
|
||||
if "alpha_mask" not in npz:
|
||||
return False
|
||||
if npz["alpha_mask"].shape[0:2] != reso: # HxW
|
||||
if (npz["alpha_mask"].shape[1], npz["alpha_mask"].shape[0]) != reso: # HxW => WxH != reso
|
||||
return False
|
||||
else:
|
||||
if "alpha_mask" in npz:
|
||||
@@ -4561,6 +4562,23 @@ def get_optimizer(args, trainable_params):
|
||||
optimizer_class = torch.optim.AdamW
|
||||
optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs)
|
||||
|
||||
elif optimizer_type.endswith("schedulefree".lower()):
|
||||
try:
|
||||
import schedulefree as sf
|
||||
except ImportError:
|
||||
raise ImportError("No schedulefree / schedulefreeがインストールされていないようです")
|
||||
if optimizer_type == "AdamWScheduleFree".lower():
|
||||
optimizer_class = sf.AdamWScheduleFree
|
||||
logger.info(f"use AdamWScheduleFree optimizer | {optimizer_kwargs}")
|
||||
elif optimizer_type == "SGDScheduleFree".lower():
|
||||
optimizer_class = sf.SGDScheduleFree
|
||||
logger.info(f"use SGDScheduleFree optimizer | {optimizer_kwargs}")
|
||||
else:
|
||||
raise ValueError(f"Unknown optimizer type: {optimizer_type}")
|
||||
optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs)
|
||||
# make optimizer as train mode: we don't need to call train again, because eval will not be called in training loop
|
||||
optimizer.train()
|
||||
|
||||
elif optimizer_type == "CAME".lower():
|
||||
logger.info(f"use CAME optimizer | {optimizer_kwargs}")
|
||||
try:
|
||||
@@ -4577,16 +4595,16 @@ def get_optimizer(args, trainable_params):
|
||||
|
||||
if optimizer is None:
|
||||
# 任意のoptimizerを使う
|
||||
optimizer_type = args.optimizer_type # lowerでないやつ(微妙)
|
||||
logger.info(f"use {optimizer_type} | {optimizer_kwargs}")
|
||||
if "." not in optimizer_type:
|
||||
case_sensitive_optimizer_type = args.optimizer_type # not lower
|
||||
logger.info(f"use {case_sensitive_optimizer_type} | {optimizer_kwargs}")
|
||||
if "." not in case_sensitive_optimizer_type: # from torch.optim
|
||||
optimizer_module = torch.optim
|
||||
else:
|
||||
values = optimizer_type.split(".")
|
||||
else: # from other library
|
||||
values = case_sensitive_optimizer_type.split(".")
|
||||
optimizer_module = importlib.import_module(".".join(values[:-1]))
|
||||
optimizer_type = values[-1]
|
||||
case_sensitive_optimizer_type = values[-1]
|
||||
|
||||
optimizer_class = getattr(optimizer_module, optimizer_type)
|
||||
optimizer_class = getattr(optimizer_module, case_sensitive_optimizer_type)
|
||||
optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs)
|
||||
|
||||
optimizer_name = optimizer_class.__module__ + "." + optimizer_class.__name__
|
||||
@@ -4594,6 +4612,30 @@ def get_optimizer(args, trainable_params):
|
||||
|
||||
return optimizer_name, optimizer_args, optimizer
|
||||
|
||||
def get_optimizer_train_eval_fn(optimizer: Optimizer, args: argparse.Namespace) -> Tuple[Callable, Callable]:
|
||||
if not is_schedulefree_optimizer(optimizer, args):
|
||||
# return dummy func
|
||||
return lambda: None, lambda: None
|
||||
# get train and eval functions from optimizer
|
||||
train_fn = optimizer.train
|
||||
eval_fn = optimizer.eval
|
||||
return train_fn, eval_fn
|
||||
|
||||
def is_schedulefree_optimizer(optimizer: Optimizer, args: argparse.Namespace) -> bool:
|
||||
return args.optimizer_type.lower().endswith("schedulefree".lower()) # or args.optimizer_schedulefree_wrapper
|
||||
|
||||
def get_dummy_scheduler(optimizer: Optimizer) -> Any:
|
||||
# dummy scheduler for schedulefree optimizer. supports only empty step(), get_last_lr() and optimizers.
|
||||
# this scheduler is used for logging only.
|
||||
# this isn't be wrapped by accelerator because of this class is not a subclass of torch.optim.lr_scheduler._LRScheduler
|
||||
class DummyScheduler:
|
||||
def __init__(self, optimizer: Optimizer):
|
||||
self.optimizer = optimizer
|
||||
def step(self):
|
||||
pass
|
||||
def get_last_lr(self):
|
||||
return [group["lr"] for group in self.optimizer.param_groups]
|
||||
return DummyScheduler(optimizer)
|
||||
|
||||
# Modified version of get_scheduler() function from diffusers.optimizer.get_scheduler
|
||||
# Add some checking and features to the original function.
|
||||
@@ -4603,6 +4645,9 @@ def get_scheduler_fix(args, optimizer: Optimizer, num_processes: int):
|
||||
"""
|
||||
Unified API to get any scheduler from its name.
|
||||
"""
|
||||
# if schedulefree optimizer, return dummy scheduler
|
||||
if is_schedulefree_optimizer(optimizer, args):
|
||||
return get_dummy_scheduler(optimizer)
|
||||
name = args.lr_scheduler
|
||||
num_warmup_steps: Optional[int] = args.lr_warmup_steps
|
||||
num_training_steps = args.max_train_steps * num_processes # * args.gradient_accumulation_steps
|
||||
|
||||
@@ -228,7 +228,7 @@ class OptimizerConfig:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"optimizer_type": (["adamw8bit", "adamw","prodigy", "CAME", "Lion8bit", "Lion"], {"default": "adamw8bit", "tooltip": "optimizer type"}),
|
||||
"optimizer_type": (["adamw8bit", "adamw","prodigy", "CAME", "Lion8bit", "Lion", "adamwschedulefree", "sgdschedulefree"], {"default": "adamw8bit", "tooltip": "optimizer type"}),
|
||||
"max_grad_norm": ("FLOAT",{"default": 1.0, "min": 0.0, "tooltip": "gradient clipping"}),
|
||||
"lr_scheduler": (["constant", "cosine", "cosine_with_restarts", "polynomial", "constant_with_warmup"], {"default": "constant", "tooltip": "learning rate scheduler"}),
|
||||
"lr_warmup_steps": ("INT",{"default": 0, "min": 0, "tooltip": "learning rate warmup steps"}),
|
||||
@@ -295,7 +295,7 @@ class OptimizerConfigProdigy:
|
||||
"lr_warmup_steps": ("INT",{"default": 0, "min": 0, "tooltip": "learning rate warmup steps"}),
|
||||
"lr_scheduler_num_cycles": ("INT",{"default": 1, "min": 1, "tooltip": "learning rate scheduler num cycles"}),
|
||||
"lr_scheduler_power": ("FLOAT",{"default": 1.0, "min": 0.0, "tooltip": "learning rate scheduler power"}),
|
||||
"weight_decay": ("FLOAT",{"default": 0.0, "tooltip": "weight decay (L2 penalty)"}),
|
||||
"weight_decay": ("FLOAT",{"default": 0.0, "step": 0.0001, "tooltip": "weight decay (L2 penalty)"}),
|
||||
"decouple": ("BOOLEAN",{"default": True, "tooltip": "use AdamW style weight decay"}),
|
||||
"use_bias_correction": ("BOOLEAN",{"default": False, "tooltip": "turn on Adam's bias correction"}),
|
||||
"min_snr_gamma": ("FLOAT",{"default": 5.0, "min": 0.0, "step": 0.01, "tooltip": "gamma for reducing the weight of high loss timesteps. Lower numbers have stronger effect. 5 is recommended by the paper"}),
|
||||
@@ -791,6 +791,9 @@ class FluxTrainLoop:
|
||||
target_global_step = network_trainer.global_step + steps
|
||||
comfy_pbar = comfy.utils.ProgressBar(steps)
|
||||
network_trainer.comfy_pbar = comfy_pbar
|
||||
|
||||
network_trainer.optimizer_train_fn()
|
||||
|
||||
while network_trainer.global_step < target_global_step:
|
||||
steps_done = training_loop(
|
||||
break_at_steps = target_global_step,
|
||||
@@ -869,6 +872,8 @@ class FluxTrainSaveModel:
|
||||
with torch.inference_mode(False):
|
||||
trainer = network_trainer["network_trainer"]
|
||||
global_step = trainer.global_step
|
||||
|
||||
trainer.optimizer_eval_fn()
|
||||
|
||||
ckpt_name = train_util.get_step_ckpt_name(trainer.args, "." + trainer.args.save_model_as, global_step)
|
||||
flux_train_utils.save_flux_model_on_epoch_end_or_stepwise(
|
||||
@@ -916,6 +921,7 @@ class FluxTrainEnd:
|
||||
network = network_trainer.accelerator.unwrap_model(network_trainer.network)
|
||||
|
||||
network_trainer.accelerator.end_training()
|
||||
network_trainer.optimizer_eval_fn()
|
||||
|
||||
if save_state:
|
||||
train_util.save_state_on_train_end(network_trainer.args, network_trainer.accelerator)
|
||||
@@ -1070,7 +1076,7 @@ class FluxTrainValidate:
|
||||
network_trainer.sample_prompts_te_outputs,
|
||||
validation_settings
|
||||
)
|
||||
|
||||
network_trainer.optimizer_eval_fn()
|
||||
image_tensors = network_trainer.sample_images(*params)
|
||||
|
||||
trainer = {
|
||||
|
||||
+2
-1
@@ -20,4 +20,5 @@ came_pytorch
|
||||
matplotlib
|
||||
# for T5XXL tokenizer (SD3/FLUX)
|
||||
sentencepiece>=0.2.0
|
||||
protobuf
|
||||
protobuf
|
||||
schedulefree>=1.2.7
|
||||
@@ -504,6 +504,7 @@ class NetworkTrainer:
|
||||
# accelerator.print(f"trainable_params: {k} = {v}")
|
||||
|
||||
optimizer_name, optimizer_args, optimizer = train_util.get_optimizer(args, trainable_params)
|
||||
self.optimizer_train_fn, self.optimizer_eval_fn = train_util.get_optimizer_train_eval_fn(optimizer, args)
|
||||
|
||||
# prepare dataloader
|
||||
# strategies are set here because they cannot be referenced in another process. Copy them with the dataset
|
||||
|
||||
Reference in New Issue
Block a user