MCPcopy Create free account
hub / github.com/tdrussell/diffusion-pipe / get_optimizer

Function get_optimizer

train.py:640–805  ·  view source on GitHub ↗
(model_parameters)

Source from the content-addressed store, hash-verified

638 print(f"Computed eval_every_n_steps = {config['eval_every_n_steps']}")
639
640 def get_optimizer(model_parameters):
641 if len(model_parameters) == 0:
642 return DummyOptimizer()
643
644 optim_config = config['optimizer']
645 optim_type = optim_config['type']
646 optim_type_lower = optim_type.lower()
647
648 if beta2_half_life := optim_config.pop('beta2_half_life', None):
649 betas = optim_config['betas']
650 assert len(betas) == 2
651 betas[1] = 0.5 ** (global_batch_size / beta2_half_life)
652 print(f'Computed beta2 = {betas[1]}')
653 optim_config['betas'] = betas
654
655 args = []
656 kwargs = {k: v for k, v in optim_config.items() if k not in ['type', 'gradient_release']}
657
658 if optim_type_lower == 'adamw':
659 # TODO: fix this. I'm getting "fatal error: cuda_runtime.h: No such file or directory"
660 # when Deepspeed tries to build the fused Adam extension.
661 # klass = deepspeed.ops.adam.FusedAdam
662 klass = torch.optim.AdamW
663 elif optim_type_lower == 'adamw8bit':
664 import bitsandbytes
665 klass = bitsandbytes.optim.AdamW8bit
666 elif optim_type_lower == 'adamw_optimi':
667 import optimi
668 klass = optimi.AdamW
669 elif optim_type_lower == 'stableadamw':
670 import optimi
671 klass = optimi.StableAdamW
672 elif optim_type_lower == 'sgd':
673 klass = torch.optim.SGD
674 elif optim_type_lower == 'adamw8bitkahan':
675 from optimizers import adamw_8bit
676 klass = adamw_8bit.AdamW8bitKahan
677 elif optim_type_lower == 'offload':
678 from torchao.prototype.low_bit_optim import CPUOffloadOptimizer
679 klass = CPUOffloadOptimizer
680 args.append(torch.optim.AdamW)
681 kwargs['fused'] = True
682 elif optim_type_lower == 'automagic':
683 from optimizers import automagic
684 klass = automagic.Automagic
685 elif optim_type_lower == 'genericoptim':
686 from optimizers import generic_optim
687 klass = generic_optim.GenericOptim
688 else:
689 import pytorch_optimizer
690 klass = getattr(pytorch_optimizer, optim_type)
691
692 if optim_config.get('gradient_release', False):
693 # Prevent deepspeed from logging every single param group lr
694 def _report_progress(self, step):
695 lr = self.get_lr()
696 mom = self.get_mom()
697 deepspeed.utils.logging.log_dist(f"step={step}, skipped={self.skipped_steps}, lr={lr[0]}, mom={mom[0]}", ranks=[0])

Callers

nothing calls this directly

Calls 3

DummyOptimizerClass · 0.85
getMethod · 0.80
get_param_groupsMethod · 0.45

Tested by

no test coverage detected