| 473 | self.amsgrad = amsgrad |
| 474 | |
| 475 | def get_updates(self, loss, params): |
| 476 | grads = self.get_gradients(loss, params) |
| 477 | self.updates = [] |
| 478 | |
| 479 | lr = self.lr |
| 480 | if self.initial_decay > 0: |
| 481 | lr = lr * ( # pylint: disable=g-no-augmented-assignment |
| 482 | 1. / |
| 483 | (1. + |
| 484 | self.decay * math_ops.cast(self.iterations, K.dtype(self.decay)))) |
| 485 | |
| 486 | with ops.control_dependencies([state_ops.assign_add(self.iterations, 1)]): |
| 487 | t = math_ops.cast(self.iterations, K.floatx()) |
| 488 | lr_t = lr * ( |
| 489 | K.sqrt(1. - math_ops.pow(self.beta_2, t)) / |
| 490 | (1. - math_ops.pow(self.beta_1, t))) |
| 491 | |
| 492 | ms = [K.zeros(K.int_shape(p), dtype=K.dtype(p)) for p in params] |
| 493 | vs = [K.zeros(K.int_shape(p), dtype=K.dtype(p)) for p in params] |
| 494 | if self.amsgrad: |
| 495 | vhats = [K.zeros(K.int_shape(p), dtype=K.dtype(p)) for p in params] |
| 496 | else: |
| 497 | vhats = [K.zeros(1) for _ in params] |
| 498 | self.weights = [self.iterations] + ms + vs + vhats |
| 499 | |
| 500 | for p, g, m, v, vhat in zip(params, grads, ms, vs, vhats): |
| 501 | m_t = (self.beta_1 * m) + (1. - self.beta_1) * g |
| 502 | v_t = (self.beta_2 * v) + (1. - self.beta_2) * math_ops.square(g) |
| 503 | if self.amsgrad: |
| 504 | vhat_t = math_ops.maximum(vhat, v_t) |
| 505 | p_t = p - lr_t * m_t / (K.sqrt(vhat_t) + self.epsilon) |
| 506 | self.updates.append(state_ops.assign(vhat, vhat_t)) |
| 507 | else: |
| 508 | p_t = p - lr_t * m_t / (K.sqrt(v_t) + self.epsilon) |
| 509 | |
| 510 | self.updates.append(state_ops.assign(m, m_t)) |
| 511 | self.updates.append(state_ops.assign(v, v_t)) |
| 512 | new_p = p_t |
| 513 | |
| 514 | # Apply constraints. |
| 515 | if getattr(p, 'constraint', None) is not None: |
| 516 | new_p = p.constraint(new_p) |
| 517 | |
| 518 | self.updates.append(state_ops.assign(p, new_p)) |
| 519 | return self.updates |
| 520 | |
| 521 | def get_config(self): |
| 522 | config = { |