You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在C++中获取TensorFlow计算图中所有张量的形状?

TensorFlow计算图张量形状提取:C++实现方案

当TensorFlow计算图中所有张量形状都已明确定义,且保存为protobuf文件后,需要提取所有张量的形状。以下是生成该protobuf文件的Python示例代码:

with tf.Graph().as_default() as graph:
    a = tf.compat.v1.placeholder(tf.int32, shape=(4, 3), name='a')
    b = tf.compat.v1.placeholder(tf.int32, shape=(4, 3), name='b')
    c = tf.add(a, b, name='c')

with tf.compat.v1.Session(graph=graph) as sess:
    graph_def = sess.graph.as_graph_def()
    with open('simple_graph.pb', 'wb') as f:
        f.write(graph_def.SerializeToString())

我尝试用以下C++代码提取形状,但shape()返回空向量:

void GetTensorShapes(const tensorflow::GraphDef &graph_def) {
    for (const auto &node : graph_def.node()) {
            const auto &shape_attr = node.attr().at("shape");
            const tensorflow::TensorShapeProto &shape = shape_attr.shape();
            std::cout << "Node name: " << node.name() << ", shape: ";
            for (const auto &dim : shape.dim()) {
                std::cout << dim.size() << " ";
            }
            std::cout << std::endl;
    }
}

对应的Python实现可以正常工作:

def get_shapes(path):
    graph_def = tf.GraphDef()
    with open(path, 'rb') as f:
        graph_def.ParseFromString(f.read())

    with tf.Graph().as_default() as graph:
        tf.import_graph_def(graph_def)

    with tf.Session() as sess:
        input_shapes = []
        for op in graph.get_operations():
            for output in op.outputs:
                shape = output.shape
                input_shapes.append([int(d.value) for d in shape.dims])
    
    return input_shapes

正确的C++实现方案

直接从GraphDef的节点属性中读取shape并不准确——只有占位符这类节点会在attr里存shape,而像Add这类运算节点的输出形状需要通过TensorFlow的形状推断才能得到,就像Python代码里导入图定义后自动完成的逻辑。

要在C++中实现等价功能,需要将GraphDef导入到Graph对象中,利用TensorFlow的形状推断机制获取每个张量的形状:

#include <tensorflow/core/graph/graph.h>
#include <tensorflow/core/graph/graph_def_builder.h>
#include <tensorflow/core/framework/tensor_shape.h>
#include <iostream>
#include <tensorflow/core/framework/op.h>

void GetTensorShapes(const tensorflow::GraphDef& graph_def) {
    tensorflow::Graph graph(tensorflow::OpRegistry::Global());
    // 将GraphDef导入到Graph对象,触发形状推断
    tensorflow::Status import_status = tensorflow::ImportGraphDef(graph_def, &graph);
    if (!import_status.ok()) {
        std::cerr << "导入图定义失败: " << import_status.ToString() << std::endl;
        return;
    }

    // 遍历所有运算节点的输出张量
    for (const auto* node : graph.nodes()) {
        for (int output_idx = 0; output_idx < node->num_outputs(); ++output_idx) {
            const tensorflow::TensorShape& tensor_shape = node->output_type(output_idx).shape();
            std::cout << "张量: " << node->name() << ":" << output_idx << ", 形状: ";
            for (int dim_idx = 0; dim_idx < tensor_shape.dims(); ++dim_idx) {
                std::cout << tensor_shape.dim_size(dim_idx) << " ";
            }
            std::cout << std::endl;
        }
    }
}

关键说明

  • GraphDef仅保存节点的原始配置,运算节点的输出形状需要导入到Graph中完成推断后才能获取。
  • graph.nodes()遍历所有运算节点,每个节点的output_type(output_idx)可获取对应输出张量的类型与形状信息。
  • 导入过程依赖TensorFlow的Op注册表,确保所有运算的定义能被正确识别。

内容的提问来源于stack exchange,提问作者Dan8757

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.26 13:45:24