| 614 | |
| 615 | |
| 616 | class 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 |
no outgoing calls
searching dependent graphs…