MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / get_updates

Method get_updates

tensorflow/python/keras/optimizers.py:652–697  ·  view source on GitHub ↗
(self, loss, params)

Source from the content-addressed store, hash-verified

650 self.schedule_decay = schedule_decay
651
652 def get_updates(self, loss, params):
653 grads = self.get_gradients(loss, params)
654 self.updates = []
655
656 with ops.control_dependencies([state_ops.assign_add(self.iterations, 1)]):
657 t = math_ops.cast(self.iterations, K.floatx())
658
659 # Due to the recommendations in [2], i.e. warming momentum schedule
660 momentum_cache_t = self.beta_1 * (
661 1. - 0.5 *
662 (math_ops.pow(K.cast_to_floatx(0.96), t * self.schedule_decay)))
663 momentum_cache_t_1 = self.beta_1 * (
664 1. - 0.5 *
665 (math_ops.pow(K.cast_to_floatx(0.96), (t + 1) * self.schedule_decay)))
666 m_schedule_new = self.m_schedule * momentum_cache_t
667 m_schedule_next = self.m_schedule * momentum_cache_t * momentum_cache_t_1
668 self.updates.append((self.m_schedule, m_schedule_new))
669
670 shapes = [K.int_shape(p) for p in params]
671 ms = [K.zeros(shape) for shape in shapes]
672 vs = [K.zeros(shape) for shape in shapes]
673
674 self.weights = [self.iterations, self.m_schedule] + ms + vs
675
676 for p, g, m, v in zip(params, grads, ms, vs):
677 # the following equations given in [1]
678 g_prime = g / (1. - m_schedule_new)
679 m_t = self.beta_1 * m + (1. - self.beta_1) * g
680 m_t_prime = m_t / (1. - m_schedule_next)
681 v_t = self.beta_2 * v + (1. - self.beta_2) * math_ops.square(g)
682 v_t_prime = v_t / (1. - math_ops.pow(self.beta_2, t))
683 m_t_bar = (1. -
684 momentum_cache_t) * g_prime + momentum_cache_t_1 * m_t_prime
685
686 self.updates.append(state_ops.assign(m, m_t))
687 self.updates.append(state_ops.assign(v, v_t))
688
689 p_t = p - self.lr * m_t_bar / (K.sqrt(v_t_prime) + self.epsilon)
690 new_p = p_t
691
692 # Apply constraints.
693 if getattr(p, 'constraint', None) is not None:
694 new_p = p.constraint(new_p)
695
696 self.updates.append(state_ops.assign(p, new_p))
697 return self.updates
698
699 def get_config(self):
700 config = {

Callers

nothing calls this directly

Calls 8

get_gradientsMethod · 0.45
control_dependenciesMethod · 0.45
assign_addMethod · 0.45
castMethod · 0.45
appendMethod · 0.45
squareMethod · 0.45
assignMethod · 0.45
constraintMethod · 0.45

Tested by

no test coverage detected