【问题标题】:How to save&restore DNNClassifier trained in TensorFlow python; iris example如何保存和恢复在 TensorFlow python 中训练的 DNNClassifier;虹膜示例
【发布时间】:2016-07-13 17:42:09
【问题描述】:

我是 TensorFlow 新手,前几天刚开始学习。我已经完成了本教程(https://www.tensorflow.org/versions/r0.9/tutorials/tflearn/index.html#tf-contrib-learn-quickstart)并将完全相同的想法应用于我自己的数据集。 (结果非常好!)

现在我想保存和恢复经过训练的 DNNClassifier 以供进一步使用。如果有人知道怎么做,请使用上面链接中的 iris 示例代码告诉我。提前感谢您的帮助!

【问题讨论】:

    标签: python tensorflow


    【解决方案1】:

    找到解决方案了吗?如果您没有这样做,您可以在创建 DNNClassifier 时在构造函数中指定 model_dir 参数,这将在此目录中创建所有检查点和文件(保存步骤)。当您想要执行恢复步骤时,您只需创建另一个 DNNClassifier 传递相同的 model_dir 参数(恢复阶段),这将从第一次创建的文件中恢复模型。

    希望对你有所帮助。

    【讨论】:

    • 很抱歉延迟回复您的第一条评论...非常感谢您的帮助!!!我今天会试试的!我只尝试在初始化中使用 model_dir,但调用了 restore 函数来加载训练有素的分类器。它抱怨 model.def 丢失,但您的建议不是调用 restore 函数,而是通过将 model_dir 的路径设置正确来使用其构造函数初始化 DNNClassifier?
    • 是的,没错,您在 DNNclassifier 的构造函数中使用了 model_dir 参数,这会保存模型,为了恢复,您只需创建另一个具有相同 model_dir 的 DNNClassifier,它将读取生成的文件
    • 您的解决方案真的有效吗?我收到以下错误:ValueError:应定义 linear_feature_columns 或 dnn_feature_columns。
    • 抱歉,我不确定如何在此处放置我的代码,因此将其发布为下面的新答案。请看一下!我将 v9.0 TensorFlow 与 Python 2.7 一起使用。非常感谢!
    • 哦,我想我看到了你的问题,tensorflow 中有一个错误,尽管你有这个 model_dir 参数来恢复你的模型,但它不能立即使用,所以你必须至少做一些“ dummy” 恢复后的训练,这意味着你必须再次运行“fit”至少一步:classifier.fit(x=x_train, y=y_train, steps=2);这将获取您保存的模型,并执行“虚拟”训练,以便再次使用它。
    【解决方案2】:

    下面是我的代码...

    import tensorflow as tf
    import numpy as np
    
    if __name__ == '__main__':
    # Data sets
    IRIS_TRAINING = "iris_training.csv"
    IRIS_TEST = "iris_test.csv"
    
    # Load datasets.
    training_set = tf.contrib.learn.datasets.base.load_csv(filename=IRIS_TRAINING, target_dtype=np.int)
    test_set = tf.contrib.learn.datasets.base.load_csv(filename=IRIS_TEST, target_dtype=np.int)
    
    x_train, x_test, y_train, y_test = training_set.data, test_set.data, training_set.target, test_set.target
    
    # Build 3 layer DNN with 10, 20, 10 units respectively.
    classifier = tf.contrib.learn.DNNClassifier(hidden_units=[10, 20, 10], n_classes=3, model_dir="path_to_my_local_dir")
    
    # print classifier.model_dir
    
    # Fit model.
    print "start fitting model..."
    classifier.fit(x=x_train, y=y_train, steps=200)
    print "finished fitting model!!!"
    
    # Evaluate accuracy.
    accuracy_score = classifier.evaluate(x=x_test, y=y_test)["accuracy"]
    print('Accuracy: {0:f}'.format(accuracy_score))
    
    #Classify two new flower samples.
    new_samples = np.array(
        [[6.4, 3.2, 4.5, 1.5], [5.8, 3.1, 5.0, 1.7]], dtype=float)
    y = classifier.predict_proba(new_samples)
    print ('Predictions: {}'.format(str(y)))
    
    #---------------------------------------------------------------------------------
    #model_dir below has to be the same as the previously specified path!
    new_classifier = tf.contrib.learn.DNNClassifier(hidden_units=[10, 20, 10], n_classes=3, model_dir="path_to_my_local_dir")
    accuracy_score = new_classifier.evaluate(x=x_test, y=y_test)["accuracy"]
    print('Accuracy: {0:f}'.format(accuracy_score))
    new_samples = np.array(
        [[6.4, 3.2, 4.5, 1.5], [5.8, 3.1, 5.0, 1.7]], dtype=float)
    y = classifier.predict_proba(new_samples)
    print ('Predictions: {}'.format(str(y)))
    

    【讨论】:

    • 这段代码真的是你自己问题的答案吗?如果它只是显示您的问题,那么请将其包含在问题中。
    猜你喜欢
    • 1970-01-01
    • 1970-01-01
    • 2018-02-08
    • 2018-04-25
    • 2016-02-14
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多