"""Optimizer and learning-rate scheduler factory. Hosts the unified optimizer dispatcher (``get_optimizer``) covering AdamW / 8-bit / Lion / DAdaptation / Prodigy / Adafactor / schedule-free / arbitrary ``module.Class`` forms; the schedule-free helpers (``is_schedulefree_optimizer``, ``get_optimizer_train_eval_fn``, ``get_dummy_scheduler``); the LR scheduler factory (``get_scheduler_fix``); and the LR-logging helpers (``append_lr_to_logs``, ``append_lr_to_logs_with_names``). Extracted from ``library.train_util`` and re-exported there for backward compatibility. """ import argparse import ast import importlib import logging from typing import Any, Callable, Optional, Tuple import torch import transformers from diffusers.optimization import ( SchedulerType as DiffusersSchedulerType, TYPE_TO_SCHEDULER_FUNCTION as DIFFUSERS_TYPE_TO_SCHEDULER_FUNCTION, ) from torch.optim import Optimizer from transformers.optimization import SchedulerType, TYPE_TO_SCHEDULER_FUNCTION from library.utils import setup_logging setup_logging() logger = logging.getLogger(__name__) def get_optimizer(args, trainable_params) -> tuple[str, str, object]: # "Optimizer to use: AdamW, AdamW8bit, Lion, SGDNesterov, SGDNesterov8bit, PagedAdamW, PagedAdamW8bit, PagedAdamW32bit, Lion8bit, PagedLion8bit, AdEMAMix8bit, PagedAdEMAMix8bit, DAdaptation(DAdaptAdamPreprint), DAdaptAdaGrad, DAdaptAdam, DAdaptAdan, DAdaptAdanIP, DAdaptLion, DAdaptSGD, Adafactor" optimizer_type = args.optimizer_type if args.use_8bit_adam: assert ( not args.use_lion_optimizer ), "both option use_8bit_adam and use_lion_optimizer are specified / use_8bit_adamとuse_lion_optimizerの両方のオプションが指定されています" assert ( optimizer_type is None or optimizer_type == "" ), "both option use_8bit_adam and optimizer_type are specified / use_8bit_adamとoptimizer_typeの両方のオプションが指定されています" optimizer_type = "AdamW8bit" elif args.use_lion_optimizer: assert ( optimizer_type is None or optimizer_type == "" ), "both option use_lion_optimizer and optimizer_type are specified / use_lion_optimizerとoptimizer_typeの両方のオプションが指定されています" optimizer_type = "Lion" if optimizer_type is None or optimizer_type == "": optimizer_type = "AdamW" optimizer_type = optimizer_type.lower() if args.fused_backward_pass: assert ( optimizer_type == "Adafactor".lower() ), "fused_backward_pass currently only works with optimizer_type Adafactor / fused_backward_passは現在optimizer_type Adafactorでのみ機能します" assert ( args.gradient_accumulation_steps == 1 ), "fused_backward_pass does not work with gradient_accumulation_steps > 1 / fused_backward_passはgradient_accumulation_steps>1では機能しません" # 引数を分解する optimizer_kwargs = {} if args.optimizer_args is not None and len(args.optimizer_args) > 0: for arg in args.optimizer_args: key, value = arg.split("=") value = ast.literal_eval(value) # value = value.split(",") # for i in range(len(value)): # if value[i].lower() == "true" or value[i].lower() == "false": # value[i] = value[i].lower() == "true" # else: # value[i] = ast.float(value[i]) # if len(value) == 1: # value = value[0] # else: # value = tuple(value) optimizer_kwargs[key] = value # logger.info(f"optkwargs {optimizer}_{kwargs}") lr = args.learning_rate optimizer = None optimizer_class = None if optimizer_type == "Lion".lower(): try: import lion_pytorch except ImportError: raise ImportError("No lion_pytorch / lion_pytorch がインストールされていないようです") logger.info(f"use Lion optimizer | {optimizer_kwargs}") optimizer_class = lion_pytorch.Lion optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) elif optimizer_type.endswith("8bit".lower()): try: import bitsandbytes as bnb except ImportError: raise ImportError("No bitsandbytes / bitsandbytesがインストールされていないようです") if optimizer_type == "AdamW8bit".lower(): logger.info(f"use 8-bit AdamW optimizer | {optimizer_kwargs}") optimizer_class = bnb.optim.AdamW8bit optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) elif optimizer_type == "SGDNesterov8bit".lower(): logger.info(f"use 8-bit SGD with Nesterov optimizer | {optimizer_kwargs}") if "momentum" not in optimizer_kwargs: logger.warning( f"8-bit SGD with Nesterov must be with momentum, set momentum to 0.9 / 8-bit SGD with Nesterovはmomentum指定が必須のため0.9に設定します" ) optimizer_kwargs["momentum"] = 0.9 optimizer_class = bnb.optim.SGD8bit optimizer = optimizer_class(trainable_params, lr=lr, nesterov=True, **optimizer_kwargs) elif optimizer_type == "Lion8bit".lower(): logger.info(f"use 8-bit Lion optimizer | {optimizer_kwargs}") try: optimizer_class = bnb.optim.Lion8bit except AttributeError: raise AttributeError( "No Lion8bit. The version of bitsandbytes installed seems to be old. Please install 0.38.0 or later. / Lion8bitが定義されていません。インストールされているbitsandbytesのバージョンが古いようです。0.38.0以上をインストールしてください" ) elif optimizer_type == "PagedAdamW8bit".lower(): logger.info(f"use 8-bit PagedAdamW optimizer | {optimizer_kwargs}") try: optimizer_class = bnb.optim.PagedAdamW8bit except AttributeError: raise AttributeError( "No PagedAdamW8bit. The version of bitsandbytes installed seems to be old. Please install 0.39.0 or later. / PagedAdamW8bitが定義されていません。インストールされているbitsandbytesのバージョンが古いようです。0.39.0以上をインストールしてください" ) elif optimizer_type == "PagedLion8bit".lower(): logger.info(f"use 8-bit Paged Lion optimizer | {optimizer_kwargs}") try: optimizer_class = bnb.optim.PagedLion8bit except AttributeError: raise AttributeError( "No PagedLion8bit. The version of bitsandbytes installed seems to be old. Please install 0.39.0 or later. / PagedLion8bitが定義されていません。インストールされているbitsandbytesのバージョンが古いようです。0.39.0以上をインストールしてください" ) if optimizer_class is not None: optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) elif optimizer_type == "PagedAdamW".lower(): logger.info(f"use PagedAdamW optimizer | {optimizer_kwargs}") try: import bitsandbytes as bnb except ImportError: raise ImportError("No bitsandbytes / bitsandbytesがインストールされていないようです") try: optimizer_class = bnb.optim.PagedAdamW except AttributeError: raise AttributeError( "No PagedAdamW. The version of bitsandbytes installed seems to be old. Please install 0.39.0 or later. / PagedAdamWが定義されていません。インストールされているbitsandbytesのバージョンが古いようです。0.39.0以上をインストールしてください" ) optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) elif optimizer_type == "PagedAdamW32bit".lower(): logger.info(f"use 32-bit PagedAdamW optimizer | {optimizer_kwargs}") try: import bitsandbytes as bnb except ImportError: raise ImportError("No bitsandbytes / bitsandbytesがインストールされていないようです") try: optimizer_class = bnb.optim.PagedAdamW32bit except AttributeError: raise AttributeError( "No PagedAdamW32bit. The version of bitsandbytes installed seems to be old. Please install 0.39.0 or later. / PagedAdamW32bitが定義されていません。インストールされているbitsandbytesのバージョンが古いようです。0.39.0以上をインストールしてください" ) optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) elif optimizer_type == "SGDNesterov".lower(): logger.info(f"use SGD with Nesterov optimizer | {optimizer_kwargs}") if "momentum" not in optimizer_kwargs: logger.info( f"SGD with Nesterov must be with momentum, set momentum to 0.9 / SGD with Nesterovはmomentum指定が必須のため0.9に設定します" ) optimizer_kwargs["momentum"] = 0.9 optimizer_class = torch.optim.SGD optimizer = optimizer_class(trainable_params, lr=lr, nesterov=True, **optimizer_kwargs) elif optimizer_type.startswith("DAdapt".lower()) or optimizer_type == "Prodigy".lower(): # check lr and lr_count, and logger.info warning actual_lr = lr lr_count = 1 if type(trainable_params) == list and type(trainable_params[0]) == dict: lrs = set() actual_lr = trainable_params[0].get("lr", actual_lr) for group in trainable_params: lrs.add(group.get("lr", actual_lr)) lr_count = len(lrs) if actual_lr <= 0.1: logger.warning( f"learning rate is too low. If using D-Adaptation or Prodigy, set learning rate around 1.0 / 学習率が低すぎるようです。D-AdaptationまたはProdigyの使用時は1.0前後の値を指定してください: lr={actual_lr}" ) logger.warning("recommend option: lr=1.0 / 推奨は1.0です") if lr_count > 1: logger.warning( f"when multiple learning rates are specified with dadaptation (e.g. for Text Encoder and U-Net), only the first one will take effect / D-AdaptationまたはProdigyで複数の学習率を指定した場合(Text EncoderとU-Netなど)、最初の学習率のみが有効になります: lr={actual_lr}" ) if optimizer_type.startswith("DAdapt".lower()): # DAdaptation family # check dadaptation is installed try: import dadaptation import dadaptation.experimental as experimental except ImportError: raise ImportError("No dadaptation / dadaptation がインストールされていないようです") # set optimizer if optimizer_type == "DAdaptation".lower() or optimizer_type == "DAdaptAdamPreprint".lower(): optimizer_class = experimental.DAdaptAdamPreprint logger.info(f"use D-Adaptation AdamPreprint optimizer | {optimizer_kwargs}") elif optimizer_type == "DAdaptAdaGrad".lower(): optimizer_class = dadaptation.DAdaptAdaGrad logger.info(f"use D-Adaptation AdaGrad optimizer | {optimizer_kwargs}") elif optimizer_type == "DAdaptAdam".lower(): optimizer_class = dadaptation.DAdaptAdam logger.info(f"use D-Adaptation Adam optimizer | {optimizer_kwargs}") elif optimizer_type == "DAdaptAdan".lower(): optimizer_class = dadaptation.DAdaptAdan logger.info(f"use D-Adaptation Adan optimizer | {optimizer_kwargs}") elif optimizer_type == "DAdaptAdanIP".lower(): optimizer_class = experimental.DAdaptAdanIP logger.info(f"use D-Adaptation AdanIP optimizer | {optimizer_kwargs}") elif optimizer_type == "DAdaptLion".lower(): optimizer_class = dadaptation.DAdaptLion logger.info(f"use D-Adaptation Lion optimizer | {optimizer_kwargs}") elif optimizer_type == "DAdaptSGD".lower(): optimizer_class = dadaptation.DAdaptSGD logger.info(f"use D-Adaptation SGD optimizer | {optimizer_kwargs}") else: raise ValueError(f"Unknown optimizer type: {optimizer_type}") optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) else: # Prodigy # check Prodigy is installed try: import prodigyopt except ImportError: raise ImportError("No Prodigy / Prodigy がインストールされていないようです") logger.info(f"use Prodigy optimizer | {optimizer_kwargs}") optimizer_class = prodigyopt.Prodigy optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) elif optimizer_type == "Adafactor".lower(): # 引数を確認して適宜補正する if "relative_step" not in optimizer_kwargs: optimizer_kwargs["relative_step"] = True # default if not optimizer_kwargs["relative_step"] and optimizer_kwargs.get("warmup_init", False): logger.info( f"set relative_step to True because warmup_init is True / warmup_initがTrueのためrelative_stepをTrueにします" ) optimizer_kwargs["relative_step"] = True logger.info(f"use Adafactor optimizer | {optimizer_kwargs}") if optimizer_kwargs["relative_step"]: logger.info(f"relative_step is true / relative_stepがtrueです") if lr != 0.0: logger.warning(f"learning rate is used as initial_lr / 指定したlearning rateはinitial_lrとして使用されます") args.learning_rate = None # trainable_paramsがgroupだった時の処理:lrを削除する if type(trainable_params) == list and type(trainable_params[0]) == dict: has_group_lr = False for group in trainable_params: p = group.pop("lr", None) has_group_lr = has_group_lr or (p is not None) if has_group_lr: # 一応argsを無効にしておく TODO 依存関係が逆転してるのであまり望ましくない logger.warning(f"unet_lr and text_encoder_lr are ignored / unet_lrとtext_encoder_lrは無視されます") args.unet_lr = None args.text_encoder_lr = None if args.lr_scheduler != "adafactor": logger.info(f"use adafactor_scheduler / スケジューラにadafactor_schedulerを使用します") args.lr_scheduler = f"adafactor:{lr}" # ちょっと微妙だけど lr = None else: if args.max_grad_norm != 0.0: logger.warning( f"because max_grad_norm is set, clip_grad_norm is enabled. consider set to 0 / max_grad_normが設定されているためclip_grad_normが有効になります。0に設定して無効にしたほうがいいかもしれません" ) if args.lr_scheduler != "constant_with_warmup": logger.warning(f"constant_with_warmup will be good / スケジューラはconstant_with_warmupが良いかもしれません") if optimizer_kwargs.get("clip_threshold", 1.0) != 1.0: logger.warning(f"clip_threshold=1.0 will be good / clip_thresholdは1.0が良いかもしれません") optimizer_class = transformers.optimization.Adafactor optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) elif optimizer_type == "AdamW".lower(): logger.info(f"use AdamW optimizer | {optimizer_kwargs}") 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 == "RAdamScheduleFree".lower(): optimizer_class = sf.RAdamScheduleFree logger.info(f"use RAdamScheduleFree optimizer | {optimizer_kwargs}") elif 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: optimizer_class = None if optimizer_class is not None: optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) if optimizer is None: # 任意のoptimizerを使う 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: # from other library values = case_sensitive_optimizer_type.split(".") optimizer_module = importlib.import_module(".".join(values[:-1])) case_sensitive_optimizer_type = values[-1] optimizer_class = getattr(optimizer_module, case_sensitive_optimizer_type) optimizer = optimizer_class(trainable_params, lr=lr, **optimizer_kwargs) """ # wrap any of above optimizer with schedulefree, if optimizer is not schedulefree if args.optimizer_schedulefree_wrapper and not optimizer_type.endswith("schedulefree".lower()): try: import schedulefree as sf except ImportError: raise ImportError("No schedulefree / schedulefreeがインストールされていないようです") schedulefree_wrapper_kwargs = {} if args.schedulefree_wrapper_args is not None and len(args.schedulefree_wrapper_args) > 0: for arg in args.schedulefree_wrapper_args: key, value = arg.split("=") value = ast.literal_eval(value) schedulefree_wrapper_kwargs[key] = value sf_wrapper = sf.ScheduleFreeWrapper(optimizer, **schedulefree_wrapper_kwargs) sf_wrapper.train() # make optimizer as train mode # we need to make optimizer as a subclass of torch.optim.Optimizer, we make another Proxy class over SFWrapper class OptimizerProxy(torch.optim.Optimizer): def __init__(self, sf_wrapper): self._sf_wrapper = sf_wrapper def __getattr__(self, name): return getattr(self._sf_wrapper, name) # override properties @property def state(self): return self._sf_wrapper.state @state.setter def state(self, state): self._sf_wrapper.state = state @property def param_groups(self): return self._sf_wrapper.param_groups @param_groups.setter def param_groups(self, param_groups): self._sf_wrapper.param_groups = param_groups @property def defaults(self): return self._sf_wrapper.defaults @defaults.setter def defaults(self, defaults): self._sf_wrapper.defaults = defaults def add_param_group(self, param_group): self._sf_wrapper.add_param_group(param_group) def load_state_dict(self, state_dict): self._sf_wrapper.load_state_dict(state_dict) def state_dict(self): return self._sf_wrapper.state_dict() def zero_grad(self): self._sf_wrapper.zero_grad() def step(self, closure=None): self._sf_wrapper.step(closure) def train(self): self._sf_wrapper.train() def eval(self): self._sf_wrapper.eval() # isinstance チェックをパスするためのメソッド def __instancecheck__(self, instance): return isinstance(instance, (type(self), Optimizer)) optimizer = OptimizerProxy(sf_wrapper) logger.info(f"wrap optimizer with ScheduleFreeWrapper | {schedulefree_wrapper_kwargs}") """ # for logging optimizer_name = optimizer_class.__module__ + "." + optimizer_class.__name__ optimizer_args = ",".join([f"{k}={v}" for k, v in optimizer_kwargs.items()]) if hasattr(optimizer, "train") and callable(optimizer.train): # make optimizer as train mode before training for schedulefree optimizer. the optimizer will be in eval mode in sampling and saving. optimizer.train() 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. 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_training_steps = args.max_train_steps * num_processes # * args.gradient_accumulation_steps num_warmup_steps: Optional[int] = ( int(args.lr_warmup_steps * num_training_steps) if isinstance(args.lr_warmup_steps, float) else args.lr_warmup_steps ) num_decay_steps: Optional[int] = ( int(args.lr_decay_steps * num_training_steps) if isinstance(args.lr_decay_steps, float) else args.lr_decay_steps ) num_stable_steps = num_training_steps - num_warmup_steps - num_decay_steps num_cycles = args.lr_scheduler_num_cycles power = args.lr_scheduler_power timescale = args.lr_scheduler_timescale min_lr_ratio = args.lr_scheduler_min_lr_ratio lr_scheduler_kwargs = {} # get custom lr_scheduler kwargs if args.lr_scheduler_args is not None and len(args.lr_scheduler_args) > 0: for arg in args.lr_scheduler_args: key, value = arg.split("=") value = ast.literal_eval(value) lr_scheduler_kwargs[key] = value def wrap_check_needless_num_warmup_steps(return_vals): if num_warmup_steps is not None and num_warmup_steps != 0: raise ValueError(f"{name} does not require `num_warmup_steps`. Set None or 0.") return return_vals # using any lr_scheduler from other library if args.lr_scheduler_type: lr_scheduler_type = args.lr_scheduler_type logger.info(f"use {lr_scheduler_type} | {lr_scheduler_kwargs} as lr_scheduler") if "." not in lr_scheduler_type: # default to use torch.optim lr_scheduler_module = torch.optim.lr_scheduler else: values = lr_scheduler_type.split(".") lr_scheduler_module = importlib.import_module(".".join(values[:-1])) lr_scheduler_type = values[-1] lr_scheduler_class = getattr(lr_scheduler_module, lr_scheduler_type) lr_scheduler = lr_scheduler_class(optimizer, **lr_scheduler_kwargs) return wrap_check_needless_num_warmup_steps(lr_scheduler) if name.startswith("adafactor"): assert ( type(optimizer) == transformers.optimization.Adafactor ), f"adafactor scheduler must be used with Adafactor optimizer / adafactor schedulerはAdafactorオプティマイザと同時に使ってください" initial_lr = float(name.split(":")[1]) # logger.info(f"adafactor scheduler init lr {initial_lr}") return wrap_check_needless_num_warmup_steps(transformers.optimization.AdafactorSchedule(optimizer, initial_lr)) if name == DiffusersSchedulerType.PIECEWISE_CONSTANT.value: name = DiffusersSchedulerType(name) schedule_func = DIFFUSERS_TYPE_TO_SCHEDULER_FUNCTION[name] return schedule_func(optimizer, **lr_scheduler_kwargs) # step_rules and last_epoch are given as kwargs name = SchedulerType(name) schedule_func = TYPE_TO_SCHEDULER_FUNCTION[name] if name == SchedulerType.CONSTANT: return wrap_check_needless_num_warmup_steps(schedule_func(optimizer, **lr_scheduler_kwargs)) # All other schedulers require `num_warmup_steps` if num_warmup_steps is None: raise ValueError(f"{name} requires `num_warmup_steps`, please provide that argument.") if name == SchedulerType.CONSTANT_WITH_WARMUP: return schedule_func(optimizer, num_warmup_steps=num_warmup_steps, **lr_scheduler_kwargs) if name == SchedulerType.INVERSE_SQRT: return schedule_func(optimizer, num_warmup_steps=num_warmup_steps, timescale=timescale, **lr_scheduler_kwargs) # All other schedulers require `num_training_steps` if num_training_steps is None: raise ValueError(f"{name} requires `num_training_steps`, please provide that argument.") if name == SchedulerType.COSINE_WITH_RESTARTS: return schedule_func( optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps, num_cycles=num_cycles, **lr_scheduler_kwargs, ) if name == SchedulerType.POLYNOMIAL: return schedule_func( optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps, power=power, **lr_scheduler_kwargs ) if name == SchedulerType.COSINE_WITH_MIN_LR: return schedule_func( optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps, num_cycles=num_cycles / 2, min_lr_rate=min_lr_ratio, **lr_scheduler_kwargs, ) # these schedulers do not require `num_decay_steps` if name == SchedulerType.LINEAR or name == SchedulerType.COSINE: return schedule_func( optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps, **lr_scheduler_kwargs, ) # All other schedulers require `num_decay_steps` if num_decay_steps is None: raise ValueError(f"{name} requires `num_decay_steps`, please provide that argument.") if name == SchedulerType.WARMUP_STABLE_DECAY: return schedule_func( optimizer, num_warmup_steps=num_warmup_steps, num_stable_steps=num_stable_steps, num_decay_steps=num_decay_steps, num_cycles=num_cycles / 2, min_lr_ratio=min_lr_ratio if min_lr_ratio is not None else 0.0, **lr_scheduler_kwargs, ) return schedule_func( optimizer, num_warmup_steps=num_warmup_steps, num_training_steps=num_training_steps, num_decay_steps=num_decay_steps, **lr_scheduler_kwargs, ) def append_lr_to_logs(logs, lr_scheduler, optimizer_type, including_unet=True): names = [] if including_unet: names.append("unet") names.append("text_encoder1") names.append("text_encoder2") names.append("text_encoder3") # SD3 append_lr_to_logs_with_names(logs, lr_scheduler, optimizer_type, names) def append_lr_to_logs_with_names(logs, lr_scheduler, optimizer_type, names): lrs = lr_scheduler.get_last_lr() for lr_index in range(len(lrs)): name = names[lr_index] logs["lr/" + name] = float(lrs[lr_index]) if optimizer_type.lower().startswith("DAdapt".lower()) or optimizer_type.lower().startswith("Prodigy".lower()): logs["lr/d*lr/" + name] = ( lr_scheduler.optimizers[-1].param_groups[lr_index]["d"] * lr_scheduler.optimizers[-1].param_groups[lr_index]["lr"] ) if "effective_lr" in lr_scheduler.optimizers[-1].param_groups[lr_index]: logs["lr/d*eff_lr/" + name] = ( lr_scheduler.optimizers[-1].param_groups[lr_index]["d"] * lr_scheduler.optimizers[-1].param_groups[lr_index]["effective_lr"] )