【问题标题】:How to use a TensorFlow LinearClassifier in Java如何在 Java 中使用 TensorFlow 线性分类器
【发布时间】:2017-05-17 10:34:38
【问题描述】:

在 Python 中,我训练了一个 TensorFlow LinearClassifier 并将其保存为:

model = tf.contrib.learn.LinearClassifier(feature_columns=columns)
model.fit(input_fn=train_input_fn, steps=100)
model.export_savedmodel(export_dir, parsing_serving_input_fn)

通过使用 TensorFlow Java API,我可以使用 Java 加载这个模型:

model = SavedModelBundle.load(export_dir, "serve");

看来我应该能够使用类似的东西来运行图表

model.session().runner().feed(???, ???).fetch(???, ???).run()

但是我应该向图表提供/从图表中获取哪些变量名称/数据以提供其功能并获取类的概率?据我所知,Java 文档缺少此信息。

【问题讨论】:

    标签: java tensorflow


    【解决方案1】:

    要馈送的节点的名称取决于parsing_serving_input_fn 所做的事情,特别是它们应该是parsing_serving_input_fn 返回的Tensor 对象的名称。要获取的节点名称取决于您的预测(如果使用 Python 中的模型,则为 model.predict() 的参数)。

    也就是说,TensorFlow 保存的模型格式确实包含模型的“签名”(即,可以馈送或获取的所有张量的名称)作为可以提供提示的元数据。

    您可以从 Python 加载保存的模型并使用以下内容列出其签名:

    with tf.Session() as sess:
      md = tf.saved_model.loader.load(sess, ['serve'], export_dir)
      sig = md.signature_def[tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY]
      print(sig)
    

    将打印如下内容:

    inputs {
      key: "inputs"
      value {
        name: "input_example_tensor:0"
        dtype: DT_STRING
        tensor_shape {
          dim {
            size: -1
          }
        }
      }
    }
    outputs {
      key: "scores"
      value {
        name: "linear/binary_logistic_head/predictions/probabilities:0"
        dtype: DT_FLOAT
        tensor_shape {
          dim {
            size: -1
          }
          dim {
            size: 2
          }
        }
      }
    }
    method_name: "tensorflow/serving/classify"
    

    建议你想用 Java 做的是:

    Tensor t = /* Tensor object to be fed */
    model.session().runner().feed("input_example_tensor", t).fetch("linear/binary_logistic_head/predictions/probabilities").run()
    

    如果您的程序包含为 TensorFlow 协议缓冲区生成的 Java 代码(打包在 org.tensorflow:proto artifact 中),您也可以纯粹在 Java 中提取此信息,使用如下:

    // Same as tf.saved_model.signature_constants.DEFAULT_SERVING_SIGNATURE_DEF_KEY
    // in Python. Perhaps this should be an exported constant in TensorFlow's Java API.
    final String DEFAULT_SERVING_SIGNATURE_DEF_KEY = "serving_default"; 
    
    final SignatureDef sig =
          MetaGraphDef.parseFrom(model.metaGraphDef())
              .getSignatureDefOrThrow(DEFAULT_SERVING_SIGNATURE_DEF_KEY);
    

    您必须添加:

    import org.tensorflow.framework.MetaGraphDef;
    import org.tensorflow.framework.SignatureDef;
    

    由于 Java API 和保存模型格式有些新,文档还有很大的改进空间。

    希望对您有所帮助。

    【讨论】:

    • 感谢您的回答!这看起来很有希望。但是,我必须为 input_example_tensor 提供什么?例如,考虑TensorFlow Iris classification tutorial:导出该模型会产生与您提供的签名相同的签名(输入,dtype:DT_STRING),但我需要以某种方式为该模型提供 4 个数字。
    • 据我所知,模型需要一个序列化的 Example 协议缓冲区,但目前(1)协议缓冲区在 Java 中不可用,(2)使用 DataType String 创建张量(序列化示例需要)尚不支持。 :(
    • 仅供参考:协议缓冲区在 Java 中的 org.tensorflow:proto maven 工件 (javadoc) 中可用,标量(即单个字符串)支持 DataType.STRING 张量,但不支持多维数组然而 (github.com/tensorflow/tensorflow/issues/8531) 希望有所帮助。
    • 再次感谢您的反馈。很高兴知道这些原型也可以在 Java 中使用。关于带字符串的张量:我需要将字符串向量输入到 input_example_tensor,对吗?所以字符串标量目前没有帮助。或者我可以解决这个问题吗?
    猜你喜欢
    • 2018-03-19
    • 2023-03-09
    • 1970-01-01
    • 2015-10-03
    • 1970-01-01
    • 1970-01-01
    • 1970-01-01
    • 2018-07-06
    • 1970-01-01
    相关资源
    最近更新 更多