MCPcopy Create free account
hub / github.com/pytorch/pytorch / AdagradOptimizer

Class AdagradOptimizer

caffe2/python/optimizer.py:616–1193  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

614
615
616class AdagradOptimizer(Optimizer):
617 def __init__(
618 self,
619 alpha=0.01,
620 epsilon=1e-4,
621 decay=1,
622 weight_decay=0.0,
623 policy="fixed",
624 sparse_dedup_aggregator=None,
625 rowWise=False,
626 engine="",
627 lars=None,
628 output_effective_lr=False,
629 output_effective_lr_and_update=False,
630 pruning_options=None,
631 swa_options=None,
632 ema_options=None,
633 weight_scale=None,
634 counter_halflife=-1,
635 use_dedicated_lr_iteration_counter=False,
636 **kwargs
637 ):
638 super().__init__()
639 self.alpha = alpha
640 self.epsilon = epsilon
641 self.decay = decay
642 self.weight_decay = float(weight_decay)
643 self.policy = policy
644 self.sparse_dedup_aggregator = sparse_dedup_aggregator
645 self.rowWise = rowWise
646 self.engine = engine
647 self.lars = lars
648 self.output_effective_lr = output_effective_lr
649 self.output_effective_lr_and_update = output_effective_lr_and_update
650 self.counter_halflife = counter_halflife
651 self.init_kwargs = kwargs
652 self.weight_scale = weight_scale
653 self.use_dedicated_lr_iteration_counter = use_dedicated_lr_iteration_counter
654
655 self._process_pruning_options(pruning_options)
656 self._process_swa_options(swa_options)
657 self._process_ema_options(ema_options)
658
659 def set_mapping_for_param2ema_teacher_param(self, param_mapping: Dict[str, Any]) -> None:
660 self.param2ema_teacher_param = param_mapping
661
662 def _process_swa_options(self, swa_options):
663 self.swa_enabled = True if swa_options else False
664 if self.swa_enabled:
665 self.swa_avg_start_it = swa_options.get("swa_avg_start_it", None)
666 self.swa_avg_end_it = swa_options.get("swa_avg_end_it", None)
667 self.swa_feedback_start_it = swa_options.get("swa_feedback_start_it", None)
668 self.swa_feedback_step = swa_options.get("swa_feedback_step", None)
669 self.swa_feedback_end_it = swa_options.get("swa_feedback_end_it", None)
670
671 def _process_ema_options(self, ema_options):
672 logger.info(f"ema_options: {str(ema_options)}")
673 self.ema_enabled = ema_options and ema_options.get("ema_alpha", None) is not None

Callers 2

build_adagradFunction · 0.85

Calls

no outgoing calls

Used in the wild real call sites across dependent graphs

searching dependent graphs…