如何将自定义AMSGrad.py转换为keras.optimizers.Optimizer?
将自定义AMSGrad优化器适配为Keras兼容类型(用于model.compile())
要让自定义的AMSGrad优化器能在TensorFlow的model.compile()中使用,需要将其重构为继承自tf.keras.optimizers.Optimizer的类,以下是具体改造步骤和代码示例:
核心改造要点
- 替换原类的继承父类:从原生TensorFlow优化器改为
tf.keras.optimizers.Optimizer - 实现Keras优化器要求的核心方法:
__init__、_create_slots、_resource_apply_dense、_resource_apply_sparse、get_config - 利用Keras优化器的内置参数管理机制,替代原生TensorFlow的变量创建方式
改造后的完整代码
import tensorflow as tf from tensorflow.keras.optimizers import Optimizer class AMSGrad(Optimizer): def __init__(self, learning_rate=0.001, beta_1=0.9, beta_2=0.999, epsilon=1e-8, name="AMSGrad", **kwargs): super().__init__(name, **kwargs) self._set_hyper("learning_rate", kwargs.get("lr", learning_rate)) self._set_hyper("beta_1", beta_1) self._set_hyper("beta_2", beta_2) self._set_hyper("epsilon", epsilon) self.epsilon = epsilon if epsilon is not None else tf.keras.backend.epsilon() def _create_slots(self, var_list): for var in var_list: self.add_slot(var, "m") self.add_slot(var, "v") self.add_slot(var, "vhat") def _resource_apply_dense(self, grad, var): var_dtype = var.dtype.base_dtype lr_t = self._decayed_lr(var_dtype) beta_1_t = self._get_hyper("beta_1", var_dtype) beta_2_t = self._get_hyper("beta_2", var_dtype) epsilon_t = self._get_hyper("epsilon", var_dtype) m = self.get_slot(var, "m") v = self.get_slot(var, "v") vhat = self.get_slot(var, "vhat") # 更新一阶矩估计 m_t = m.assign(beta_1_t * m + (1. - beta_1_t) * grad) # 更新二阶矩估计 v_t = v.assign(beta_2_t * v + (1. - beta_2_t) * tf.square(grad)) # 更新AMSGrad的max二阶矩 vhat_t = vhat.assign(tf.maximum(vhat, v_t)) # 计算参数更新 var_t = var - lr_t * m_t / (tf.sqrt(vhat_t) + epsilon_t) var.assign(var_t) return tf.group(*[var_t, m_t, v_t, vhat_t]) def _resource_apply_sparse(self, grad, var, indices): var_dtype = var.dtype.base_dtype lr_t = self._decayed_lr(var_dtype) beta_1_t = self._get_hyper("beta_1", var_dtype) beta_2_t = self._get_hyper("beta_2", var_dtype) epsilon_t = self._get_hyper("epsilon", var_dtype) m = self.get_slot(var, "m") v = self.get_slot(var, "v") vhat = self.get_slot(var, "vhat") # 稀疏更新一阶矩 m_scaled_g_values = grad * (1. - beta_1_t) m_t = m.assign(beta_1_t * m) with tf.control_dependencies([m_t]): m_t = self._resource_scatter_add(m, indices, m_scaled_g_values) # 稀疏更新二阶矩 v_scaled_g_values = tf.square(grad) * (1. - beta_2_t) v_t = v.assign(beta_2_t * v) with tf.control_dependencies([v_t]): v_t = self._resource_scatter_add(v, indices, v_scaled_g_values) # 更新max二阶矩 vhat_t = vhat.assign(tf.maximum(vhat, v_t)) # 计算稀疏参数更新 var_scaled_update = lr_t * m_t / (tf.sqrt(vhat_t) + epsilon_t) var_t = var.assign_sub(var_scaled_update, use_locking=self._use_locking) return tf.group(*[var_t, m_t, v_t, vhat_t]) def get_config(self): config = super().get_config() config.update({ "learning_rate": self._serialize_hyperparameter("learning_rate"), "beta_1": self._serialize_hyperparameter("beta_1"), "beta_2": self._serialize_hyperparameter("beta_2"), "epsilon": self._serialize_hyperparameter("epsilon"), }) return config
使用方式
在模型编译时直接实例化该优化器传入即可:
model.compile(optimizer=AMSGrad(learning_rate=0.001), loss='sparse_categorical_crossentropy', metrics=['accuracy'])
内容的提问来源于stack exchange,提问作者atcw108
相关产品推荐
相关产品推荐

