From f9d3824c8a07d4f7300e866e21eb4e5df69362f1 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 19 Sep 2024 16:33:56 +0300 Subject: [PATCH] Add schedulerfree optimizers --- flux_train_comfy.py | 7 +++++ library/train_util.py | 61 +++++++++++++++++++++++++++++++++++++------ nodes.py | 12 ++++++--- requirements.txt | 3 ++- train_network.py | 1 + 5 files changed, 72 insertions(+), 12 deletions(-) diff --git a/flux_train_comfy.py b/flux_train_comfy.py index b6fdc61..1412fba 100644 --- a/flux_train_comfy.py +++ b/flux_train_comfy.py @@ -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 diff --git a/library/train_util.py b/library/train_util.py index 4af4f5b..dfaf8e8 100644 --- a/library/train_util.py +++ b/library/train_util.py @@ -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 diff --git a/nodes.py b/nodes.py index d9d745c..3245262 100644 --- a/nodes.py +++ b/nodes.py @@ -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 = { diff --git a/requirements.txt b/requirements.txt index 3da67ec..3615260 100644 --- a/requirements.txt +++ b/requirements.txt @@ -20,4 +20,5 @@ came_pytorch matplotlib # for T5XXL tokenizer (SD3/FLUX) sentencepiece>=0.2.0 -protobuf \ No newline at end of file +protobuf +schedulefree>=1.2.7 \ No newline at end of file diff --git a/train_network.py b/train_network.py index 777edde..37e1caa 100644 --- a/train_network.py +++ b/train_network.py @@ -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