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

如何在TensorFlow C API中实现张量转置、维度扩展与ArgMax操作?

Question

I’ve built code based on the hello_tf_c_api repo, and now I need to perform transpose, dimension expansion, and ArgMax operations on input/output tensors using the TensorFlow C API. The official documentation is really lacking, and I’m stuck figuring out how to implement these operations. Here’s my current implementation code:

auto graph = tf_utils::LoadGraph("D:\\projects\\vive\\Vive_Model_4_classes_frozen\\frozen_graph.pb");
SCOPE_EXIT{ tf_utils::DeleteGraph(graph); }; // Auto-delete on scope exit.
if (graph == nullptr) {
    std::cout << "Can't load graph" << std::endl;
    return 1;
}
const std::vector<TF_Output> input_ops = { {TF_GraphOperationByName(graph, "x"), 0} };
const std::vector<TF_Tensor*> input_tensors = { tf_utils::CreateTensor(TF_FLOAT, shape, floatv) };
SCOPE_EXIT{ tf_utils::DeleteTensors(input_tensors); }; // Auto-delete on scope exit.
const std::vector<TF_Output> out_ops = { {TF_GraphOperationByName(graph, "Identity"), 0} };
std::vector<TF_Tensor*> output_tensors = { nullptr };
SCOPE_EXIT{ tf_utils::DeleteTensors(output_tensors); }; // Auto-delete on scope exit.
auto session = tf_utils::CreateSession(graph);
SCOPE_EXIT{ tf_utils::DeleteSession(session); }; // Auto-delete on scope exit.
if (session == nullptr) {
    std::cout << "Can't create session" << std::endl;
    return 2;
}
auto code = tf_utils::RunSession(session, input_ops, input_tensors, out_ops, output_tensors);
if (code == TF_OK) {
    auto result = tf_utils::GetTensorData<float>(output_tensors[0]);
    outImage = avl::Image(width, height, avl::PlainType::UInt8, 3, atl::NIL);
    //auto values = output.get_data<int64_t>();
    for (int w = 0; w < width; ++w) {
        for (int h = 0; h < height; ++h) {
            //int data = result[width * height * width * h + w];
            int idx = getMaxIndex(result, 5, width, height, w, h);
            uint8_t* pixel = outImage.Ptr<uint8_t>(w, h);
            *pixel = (uint8_t)inColors[idx].X();
            *(pixel + 1) = (uint8_t)inColors[idx].Y();
            *(pixel + 2) = (uint8_t)inColors[idx].Z();
        }
    }
    //std::cout << "Output vals: " << result[0] << ", " << result[1] << ", " << result[2] << ", " << result[3] << std::endl;
} else {
    std::cout << "Error run session TF_CODE: " << code;
    return code;
}
Answer

Got it, let's break down how to implement each of these operations with the TensorFlow C API. Unlike the Python API, you’ll need to manually build new graph operations (nodes) for each transformation since the C API is far more low-level.

1. Dimension Expansion (Adding a New Axis)

To expand a tensor's dimensions (like adding a batch axis or channel dimension), use the ExpandDims operation. You’ll need to create both the ExpandDims node and a constant tensor to define which axis to add:

TF_Status* status = TF_NewStatus();

// Step 1: Create a constant tensor for the axis we want to add (e.g., axis=0 for batch dim)
int64_t axis_val = 0;
TF_Tensor* axis_tensor = tf_utils::CreateTensor(TF_INT64, {}, &axis_val);

// Step 2: Build the Const operation for the axis
TF_OperationDescription* axis_desc = TF_NewOperation(graph, "Const", "ExpandAxis");
TF_SetAttrTensor(axis_desc, "value", axis_tensor, status);
TF_SetAttrType(axis_desc, "dtype", TF_INT64);
TF_Operation* axis_const_op = TF_FinishOperation(axis_desc, status);
TF_DeleteTensor(axis_tensor); // Clean up the temp tensor

// Step 3: Build the ExpandDims operation
TF_OperationDescription* expand_desc = TF_NewOperation(graph, "ExpandDims", "InputExpand");
TF_AddInput(expand_desc, input_ops[0]); // Link to your original input tensor
TF_AddInput(expand_desc, {axis_const_op, 0}); // Link to the axis constant
TF_Operation* expand_op = TF_FinishOperation(expand_desc, status);

// Check for errors
if (TF_GetCode(status) != TF_OK) {
    std::cout << "Failed to create ExpandDims op: " << TF_Message(status) << std::endl;
    TF_DeleteStatus(status);
    return 1;
}

// Use this expanded tensor as input for subsequent ops
TF_Output expanded_input = {expand_op, 0};
TF_DeleteStatus(status);

2. Transpose Operation

For transposing, use the Transpose op. You need to define a permutation order (e.g., [0,2,1,3] swaps height and width in an NHWC tensor):

TF_Status* status = TF_NewStatus();

// Step 1: Create permutation tensor (adjust the values to match your shape needs)
int64_t perm[] = {0, 2, 1, 3}; // Example for NHWC -> NHCW
TF_Tensor* perm_tensor = tf_utils::CreateTensor(TF_INT64, {4}, perm);

// Step 2: Build Const op for permutation
TF_OperationDescription* perm_desc = TF_NewOperation(graph, "Const", "TransposePerm");
TF_SetAttrTensor(perm_desc, "value", perm_tensor, status);
TF_SetAttrType(perm_desc, "dtype", TF_INT64);
TF_Operation* perm_const_op = TF_FinishOperation(perm_desc, status);
TF_DeleteTensor(perm_tensor);

// Step 3: Build Transpose op
TF_OperationDescription* transpose_desc = TF_NewOperation(graph, "Transpose", "InputTranspose");
TF_AddInput(transpose_desc, expanded_input); // Use expanded input or original input
TF_AddInput(transpose_desc, {perm_const_op, 0}); // Link to permutation constant
TF_Operation* transpose_op = TF_FinishOperation(transpose_desc, status);

if (TF_GetCode(status) != TF_OK) {
    std::cout << "Failed to create Transpose op: " << TF_Message(status) << std::endl;
    TF_DeleteStatus(status);
    return 1;
}

// Use this transposed tensor for further processing
TF_Output transposed_input = {transpose_op, 0};
TF_DeleteStatus(status);

3. ArgMax Operation

To replace your manual getMaxIndex function, use the ArgMax op. This will compute the index of the maximum value along a specified axis (usually the class axis for segmentation models):

TF_Status* status = TF_NewStatus();

// Step 1: Create constant tensor for the axis to compute ArgMax on (e.g., axis=3 for NHWC class axis)
int64_t argmax_axis = 3;
TF_Tensor* axis_tensor = tf_utils::CreateTensor(TF_INT64, {}, &argmax_axis);

// Step 2: Build Const op for the axis
TF_OperationDescription* argmax_axis_desc = TF_NewOperation(graph, "Const", "ArgMaxAxis");
TF_SetAttrTensor(argmax_axis_desc, "value", axis_tensor, status);
TF_SetAttrType(argmax_axis_desc, "dtype", TF_INT64);
TF_Operation* argmax_axis_op = TF_FinishOperation(argmax_axis_desc, status);
TF_DeleteTensor(axis_tensor);

// Step 3: Build ArgMax op linked to your model's output
TF_OperationDescription* argmax_desc = TF_NewOperation(graph, "ArgMax", "OutputArgMax");
TF_AddInput(argmax_desc, out_ops[0]); // Link to your original model output
TF_AddInput(argmax_desc, {argmax_axis_op, 0}); // Link to the ArgMax axis constant
TF_SetAttrType(argmax_desc, "Tidx", TF_INT64); // Set output index type to int64
TF_Operation* argmax_op = TF_FinishOperation(argmax_desc, status);

if (TF_GetCode(status) != TF_OK) {
    std::cout << "Failed to create ArgMax op: " << TF_Message(status) << std::endl;
    TF_DeleteStatus(status);
    return 1;
}

// Update your output ops to use ArgMax result instead of raw model output
const std::vector<TF_Output> argmax_out_ops = { {argmax_op, 0} };
std::vector<TF_Tensor*> argmax_output_tensors = { nullptr };
SCOPE_EXIT{ tf_utils::DeleteTensors(argmax_output_tensors); };

// Run the session with the new output
auto code = tf_utils::RunSession(session, input_ops, input_tensors, argmax_out_ops, argmax_output_tensors);
if (code == TF_OK) {
    // Get the ArgMax indices (already int64, no need for manual max calculation)
    auto argmax_results = tf_utils::GetTensorData<int64_t>(argmax_output_tensors[0]);
    
    // Update your pixel coloring loop to use the precomputed indices
    for (int w = 0; w < width; ++w) {
        for (int h = 0; h < height; ++h) {
            // Adjust the index calculation to match your tensor's shape layout
            int idx = argmax_results[h * width + w];
            uint8_t* pixel = outImage.Ptr<uint8_t>(w, h);
            *pixel = (uint8_t)inColors[idx].X();
            *(pixel + 1) = (uint8_t)inColors[idx].Y();
            *(pixel + 2) = (uint8_t)inColors[idx].Z();
        }
    }
}
TF_DeleteStatus(status);

Critical Notes to Avoid Headaches:

  • Graph Modification Order: Add all new operations before creating your TensorFlow session. If you’ve already created a session, you’ll need to delete it and create a new one after modifying the graph.
  • Error Checking: Always use TF_Status to validate every operation creation and session run—this is the only way to get meaningful error messages from the C API.
  • Shape Alignment: Double-check your permutation order (for transpose) and ArgMax axis to match your tensor’s actual shape (e.g., NHWC vs NCHW formats).
  • Helper Functions: Ensure your tf_utils::CreateTensor can handle scalar tensors (empty shape {}) since we use those for axis definitions.

Content sourced from Stack Exchange, asked by TheDude

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 15:42:41