如何在TensorFlow C API中实现张量转置、维度扩展与ArgMax操作?
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; }
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_Statusto 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::CreateTensorcan handle scalar tensors (empty shape{}) since we use those for axis definitions.
Content sourced from Stack Exchange, asked by TheDude

