【问题标题】:NotImplementedError: Learning rate schedule must override get_configNotImplementedError:学习率计划必须覆盖 get_config
【发布时间】:2020-08-16 19:13:50
【问题描述】:

我使用 tf.keras 创建了一个自定义计划,但在保存模型时遇到了这个错误:

NotImplementedError:学习率计划必须覆盖 get_config

类如下所示:

class CustomSchedule(tf.keras.optimizers.schedules.LearningRateSchedule):

    def __init__(self, d_model, warmup_steps=4000):
        super(CustomSchedule, self).__init__()

        self.d_model = d_model
        self.d_model = tf.cast(self.d_model, tf.float32)

        self.warmup_steps = warmup_steps

    def __call__(self, step):
        arg1 = tf.math.rsqrt(step)
        arg2 = step * (self.warmup_steps**-1.5)

        return tf.math.rsqrt(self.d_model) * tf.math.minimum(arg1, arg2)

    def get_config(self):
        config = {
            'd_model':self.d_model,
            'warmup_steps':self.warmup_steps

        }
        base_config = super(CustomSchedule, self).get_config()
        return dict(list(base_config.items()) + list(config.items()))

【问题讨论】:

    标签: python machine-learning keras tensorflow2.0 transformer


    【解决方案1】:

    当您使用自定义子类模型时,保存模型架构有点棘手。相反,使用 Model.save_weights() 只保存权重更容易。

    如果您将代码更改为此,您将不会看到该错误:

      def get_config(self):
        config = {
        'd_model': self.d_model,
        'warmup_steps': self.warmup_steps,
    
         }
        return config
    

    【讨论】:

      猜你喜欢
      • 2020-02-28
      • 2020-12-08
      • 2021-10-09
      • 2020-11-22
      • 2015-04-16
      • 1970-01-01
      • 1970-01-01
      • 2011-01-03
      • 1970-01-01
      相关资源
      最近更新 更多