You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将自定义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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.23 11:54:14