首页
学习
活动
专区
圈层
工具
发布
社区首页 >问答首页 >Pipeline转换器与Tensorflow扩展管道(芝加哥出租车管道实例)

Pipeline转换器与Tensorflow扩展管道(芝加哥出租车管道实例)
EN

Stack Overflow用户
提问于 2019-03-26 17:22:21
回答 1查看 309关注 0票数 0

目标: TFX -> TF Lite转换器->在移动/物联网设备上的部署模型

我目前正在学习Tensorflow扩展及其芝加哥出租车管道实例。管道已经运行(尽管经历了很多困难),Pusher组件已经发出了一个Tensorflow SavedModel文件(.pb)。

但是,这里遇到了一个新的问题:通过Tensorflow夜/1.13.1(尝试了这两种方法)和Python2.7.6,我可以用一些简单的Python代码(如saved_model.loader.load )生成、保存和加载SavedModel(用于测试实用程序的mnist数字数据模型),但是在应用TFX Pusher发布的模型时经常会遇到错误,如下所示。

(也许我对TFX管道做错了什么?)

我使用的代码:

代码语言:javascript
复制
import tensorflow as tf
with tf.Session(graph=tf.Graph()) as sess:
    tf.compat.v1.saved_model.loader.load(sess, ["serve"], "/home/tigerpaws/taxi/serving_model/taxi_simple/1553187887")#"/home/tigerpaws/saved_model_example/model")
    graph=tf.get_default_graph()

错误:

代码语言:javascript
复制
KeyError                                  Traceback (most recent call last)
<ipython-input-11-a6978b82c3d2> in <module>()
      1 with tf.Session(graph=tf.Graph()) as sess:
----> 2     tf.compat.v1.saved_model.loader.load(sess, ["serve"], "/home/tigerpaws/taxi/serving_model/taxi_simple/1553187887")#"/home/tigerpaws/saved_model_example/model")
      3     graph=tf.get_default_graph()

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/util/deprecation.pyc in new_func(*args, **kwargs)
    322               'in a future version' if date is None else ('after %s' % date),
    323               instructions)
--> 324       return func(*args, **kwargs)
    325     return tf_decorator.make_decorator(
    326         func, new_func, 'deprecated',

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/saved_model/loader_impl.pyc in load(sess, tags, export_dir, import_scope, **saver_kwargs)
    267   """
    268   loader = SavedModelLoader(export_dir)
--> 269   return loader.load(sess, tags, import_scope, **saver_kwargs)
    270 
    271 

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/saved_model/loader_impl.pyc in load(self, sess, tags, import_scope, **saver_kwargs)
    418     with sess.graph.as_default():
    419       saver, _ = self.load_graph(sess.graph, tags, import_scope,
--> 420                                  **saver_kwargs)
    421       self.restore_variables(sess, saver, import_scope)
    422       self.run_init_ops(sess, tags, import_scope)

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/saved_model/loader_impl.pyc in load_graph(self, graph, tags, import_scope, **saver_kwargs)
    348     with graph.as_default():
    349       return tf_saver._import_meta_graph_with_return_elements(  # pylint: disable=protected-access
--> 350           meta_graph_def, import_scope=import_scope, **saver_kwargs)
    351 
    352   def restore_variables(self, sess, saver, import_scope=None):

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/training/saver.pyc in _import_meta_graph_with_return_elements(meta_graph_or_file, clear_devices, import_scope, return_elements, **kwargs)
   1455           import_scope=import_scope,
   1456           return_elements=return_elements,
-> 1457           **kwargs))
   1458 
   1459   saver = _create_saver_from_imported_meta_graph(

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/framework/meta_graph.pyc in import_scoped_meta_graph_with_return_elements(meta_graph_or_file, clear_devices, graph, import_scope, input_map, unbound_inputs_col_name, restore_collections_predicate, return_elements)
    804         input_map=input_map,
    805         producer_op_list=producer_op_list,
--> 806         return_elements=return_elements)
    807 
    808     # Restores all the other collections.

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/util/deprecation.pyc in new_func(*args, **kwargs)
    505                 'in a future version' if date is None else ('after %s' % date),
    506                 instructions)
--> 507       return func(*args, **kwargs)
    508 
    509     doc = _add_deprecated_arg_notice_to_docstring(

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/framework/importer.pyc in import_graph_def(graph_def, input_map, return_elements, name, op_dict, producer_op_list)
    397   if producer_op_list is not None:
    398     # TODO(skyewm): make a copy of graph_def so we're not mutating the argument?
--> 399     _RemoveDefaultAttrs(op_dict, producer_op_list, graph_def)
    400 
    401   graph = ops.get_default_graph()

/home/tigerpaws/anaconda/lib/python2.7/site-packages/tensorflow/python/framework/importer.pyc in _RemoveDefaultAttrs(op_dict, producer_op_list, graph_def)
    157     # Remove any default attr values that aren't in op_def.
    158     if node.op in producer_op_dict:
--> 159       op_def = op_dict[node.op]
    160       producer_op_def = producer_op_dict[node.op]
    161       # We make a copy of node.attr to iterate through since we may modify

KeyError: u'BucketizeWithInputBoundaries'

还有另一次尝试,我尝试将SavedModel转换成GraphDef (冻结图),这样我就可以再尝试一次。转换需要一个output_node_names,我不知道;我也找不到模型保存在代码中的位置(所以我可能可以在某个地方发现输出节点的名称)。

对这个问题有什么想法或者其他的方法吗?提前谢谢。

编辑:能帮助创建标签吗?我还没有达到1500个声誉,但这个问题实际上是关于tfx / tensorflow-extended的。

EN

回答 1

Stack Overflow用户

发布于 2019-05-25 10:09:23

很抱歉造成了混乱;问题实际上是由于读取SavedModel文件造成的。

在SavedModel中,有一个未在op_dict中定义的操作BucketizeWithInputBoundaries

这仍然在谷歌的TODO列表中,在他们的两个脚本中有评论。

这里这里.(Github链接):

代码语言:javascript
复制
# TODO(jyzhao): BucketizeWithInputBoundaries error without this.

导入指定的脚本后,这个问题就解决了。

代码语言:javascript
复制
from tensorflow.contrib.boosted_trees.python.ops import quantile_ops  # pylint: disable=unused-import
票数 0
EN
页面原文内容由Stack Overflow提供。腾讯云小微IT领域专用引擎提供翻译支持
原文链接:

https://stackoverflow.com/questions/55362917

复制
相关文章

相似问题

领券
问题归档专栏文章快讯文章归档关键词归档开发者手册归档开发者手册 Section 归档