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

TensorFlow自定义优化器报错:学习率未定义及_set_hyper属性问题

TensorFlow自定义优化器报错解决方案

错误原因分析

第一个错误:学习率未定义

直接给self.learning_rate赋值不符合TensorFlow优化器的规范,且你没有调用父类keras.optimizers.Optimizer的__init__方法,导致父类的超参数管理机制未初始化,系统无法识别你设置的学习率。

第二个错误:找不到_set_hyper方法

同样是因为未调用父类的__init__,父类的方法和属性没有被加载到子类实例中,自然找不到_set_hyper方法。

正确的自定义优化器代码

import numpy as np
import tensorflow as tf
from tensorflow import keras

class GradientDescent(keras.optimizers.Optimizer):
    def __init__(self, learning_rate=0.001, name="GradientDescent", **kwargs):
        # 必须先调用父类初始化,传入名称和其他参数
        super().__init__(name, **kwargs)
        # 用父类方法注册学习率,同时兼容lr作为参数别名
        self._set_hyper("learning_rate", kwargs.get("lr", learning_rate))

    def apply_gradients(self, grads_and_vars, name=None):
        # 遵循TensorFlow规范,参数为(梯度, 变量)对的列表
        for grad, var in grads_and_vars:
            if grad is None:
                continue
            # 获取当前学习率(支持动态调整场景)
            lr = self._get_hyper("learning_rate", tf.float32)
            # 更新变量
            var.assign_sub(lr * grad)
        return tf.no_op(name=name)

    # 重写该方法,让优化器参数能被正确序列化,支持模型保存加载
    def get_config(self):
        config = super().get_config()
        config.update({
            "learning_rate": self._serialize_hyperparameter("learning_rate"),
        })
        return config

# 测试代码(补全缺失的数据集示例)
keras.backend.clear_session()
np.random.seed(42)
tf.random.set_seed(42)
X_train_scaled = np.random.rand(100, 8)
y_train = np.random.rand(100, 1)

model = keras.models.Sequential([keras.layers.Dense(1, input_shape=[8])])
model.compile(loss="mse", optimizer=GradientDescent(learning_rate=0.001))
model.fit(X_train_scaled, y_train, epochs=5)

核心注意点

  • 必须调用父类的__init__方法,这是初始化优化器核心机制的前提
  • 用_set_hyper和_get_hyper管理超参数,这是TensorFlow优化器的标准做法,还能支持动态学习率调度
  • apply_gradients的参数要遵循父类约定的grads_and_vars格式,避免参数名冲突(比如不要用内置函数名vars)
  • 重写get_config方法,确保优化器参数能被正确保存和恢复

内容的提问来源于stack exchange,提问作者Tony Yu

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 05:54:57