Spaces:
Paused
Paused
| """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"] | |
| ) | |