(self, loss, params)
| 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 = { |
nothing calls this directly
no test coverage detected