小编Con*_*tih的帖子

访问 tf.keras.callbacks.Callback 中已弃用的属性“validation_data”

我决定从 keras 切换到 tf.keras(这里推荐)。因此,我安装tf.__version__=2.0.0和tf.keras.__version__=2.2.4-tf。在我的代码的较旧版本(使用一些较旧的 Tensorflow 版本tf.__version__=1.x.x)中,我使用回调来计算每个时期结束时整个验证数据的自定义指标。这样做的想法来自这里。但是,似乎不推荐使用“validation_data”属性,因此以下代码不再起作用。

class ValMetrics(Callback):

    def on_train_begin(self, logs={}):

        self.val_all_mse = []

    def on_epoch_end(self, epoch, logs):

        val_predict = np.asarray(self.model.predict(self.validation_data[0]))
        val_targ = self.validation_data[1]

        val_epoch_mse = mse_score(val_targ, val_predict)

        self.val_epoch_mse.append(val_epoch_mse)

        # Add custom metrics to the logs, so that we can use them with
        # EarlyStop and csvLogger callbacks
        logs["val_epoch_mse"] = val_epoch_mse

        print(f"\nEpoch: {epoch + 1}")
        print("-----------------")
        print("val_mse:     {:+.6f}".format(val_epoch_mse))

        return
Run Code Online (Sandbox Code Playgroud)

我目前的解决方法如下。我只是将validation_data作为ValMetrics类的参数:

class ValMetrics(Callback):

    def __init__(self, validation_data):
        super(Callback, self).__init__()
        self.X_val, …
Run Code Online (Sandbox Code Playgroud)

python callback keras tensorflow

6
推荐指数
1
解决办法
1296
查看次数

标签 统计

callback ×1

keras ×1

python ×1

tensorflow ×1