代码之家  ›  专栏  ›  技术社区  ›  Mark Anthony Sulleza

一个Tensorflow会话中的多个图

  •  0
  • Mark Anthony Sulleza  · 技术社区  · 8 年前

    我目前正在尝试实现一个代码,允许我的无人机使用tensorflow在室内导航。我需要在一个会话中运行两个模型。

    一个是主导航——这是一个经过重新培训的Inception V3模型,负责对走廊图像进行分类,并执行向前、向左或向右移动的决策——第二个是对象跟踪模型,它将跟踪对象并计算到摄像机的相对距离。

    我不知道如何在一个会话中使用多个图,所以我尝试在循环中创建一个单独的会话,这会产生很大的开销,并导致脚本以0 FPS的速度运行。

    def inception_model():
    # Graph for the InceptionV3 Model
    graph = load_graph('inception_v3_frozen/inception_v3_2016_08_28_frozen.pb')
    
    with tf.Session(graph = graph) as sess:
        while camera.isOpened():
            ok, img = camera.read()
            cv.imwrite("frame_temp.jpeg", img)
            t = read_tensor_from_image('frame_temp.jpeg')
    
            input_layer = "input"
            output_layer = "InceptionV3/Predictions/Reshape_1"
    
            input_name = "import/" + input_layer
            output_name = "import/" + output_layer
    
            input_operation = graph.get_operation_by_name(input_name)
            output_operation = graph.get_operation_by_name(output_name)
    
            results = sess.run(output_operation.outputs[0], {
                input_operation.outputs[0] : t
            })
            results = np.squeeze(results)
    
            top_k = results.argsort()[-5:][::-1]
            for i in top_k:
                print(labels[i], results[i])
    
    # inception_model()
    with tf.Session(graph = object_detection_graph) as sess:
        while camera.isOpened():
            ok, img = camera.read()
            cv.imwrite("frame_temp.jpeg", img)
            img = np.array(img)
            rows = img.shape[0]
            cols = img.shape[1]
    
            inp = cv.resize(img, (299, 299))
    
            # inception_model()
            # # Graph for the InceptionV3 Model
            # graph = load_graph('inception_v3_frozen/inception_v3_2016_08_28_frozen.pb')
    
            # t = read_tensor_from_image('frame_temp.jpeg')
    
            # input_layer = "input"
            # output_layer = "InceptionV3/Predictions/Reshape_1"
    
            # input_name = "import/" + input_layer
            # output_name = "import/" + output_layer
    
            # input_operation = graph.get_operation_by_name(input_name)
            # output_operation = graph.get_operation_by_name(output_name)
    
            # with tf.Session(graph = graph) as sess:
            #     results = sess.run(output_operation.outputs[0], {
            #         input_operation.outputs[0] : t
            #     })
            # results = np.squeeze(results)
    
            # top_k = results.argsort()[-5:][::-1]
            # for i in top_k:
            #     print(labels[i], results[i])
    
    
            inp = inp[:, :, [2, 1, 0]]  # BGR2RGB
    
    
            # Run the model
            out = sess.run([object_detection_graph.get_tensor_by_name('num_detections:0'),
                            object_detection_graph.get_tensor_by_name('detection_scores:0'),
                            object_detection_graph.get_tensor_by_name('detection_boxes:0'),
                            object_detection_graph.get_tensor_by_name('detection_classes:0')],
                        feed_dict={'image_tensor:0': inp.reshape(1, inp.shape[0], inp.shape[1], 3)})
    
    1 回复  |  直到 8 年前
        1
  •  3
  •   iga    8 年前

    您不必在每次迭代中创建新会话。创建它们一次并继续调用它们的run方法。Tensorflow支持多个活动会话。

    另一种选择是 Graph 对象和单个 Session 。该图可以将两个模型都包含为断开连接的子图。当你在 Session.run() Tensorflow将只运行计算所需张量所需的内容。因此,另一个子图将不会运行(尽管需要一些时间,可能很短,才能将其删除)