TensorFlow:如何基于真值类别索引提取目标检测框预测张量
Got it, let's break this down clearly. You're trying to pull the 6-dimensional box parameters for each ground truth class from your H×W×C×6 prediction tensor, resulting in a clean H×W×6 tensor. Let's walk through the most straightforward approaches, including demystifying the batch_dims parameter in tf.gather_nd.
Method 1: Use tf.gather_nd with batch_dims (Recommended)
This is the most elegant approach because it preserves your original tensor structure without needing to flatten or reshape dimensions repeatedly. The batch_dims parameter is exactly what you need here—it lets you define how many leading dimensions are treated as "independent batches" that don't require explicit indexing.
Here's how it works:
- Your prediction tensor has dimensions
[H, W, C, 6]. We want to keep the first two dimensions (H,W) intact, and only index into the third dimension (C) using the ground truth class indices. - We first reshape the
H×Wground truth tensor toH×W×1—this adds a trailing dimension so the indices match the rank expected bygather_nd. - Set
batch_dims=2to tell TensorFlow that the first 2 dimensions (H,W) are separate batches. For each position(i,j), it will use the class index atgt_classes[i,j]to pull the corresponding6-dimensional vector frompreds[i,j,:,:].
import tensorflow as tf # Simulate your input data H, W, C = 2, 3, 4 preds = tf.random.normal((H, W, C, 6)) # Shape: (2, 3, 4, 6) gt_classes = tf.random.uniform((H, W), minval=0, maxval=C, dtype=tf.int32) # Shape: (2, 3) # Step 1: Expand ground truth indices to match gather_nd's requirements indices = tf.expand_dims(gt_classes, axis=-1) # Shape: (2, 3, 1) # Step 2: Use gather_nd with batch_dims output = tf.gather_nd(preds, indices, batch_dims=2) print(output.shape) # Output: (2, 3, 6) → Exactly what you need!
What batch_dims Actually Does
Think of batch_dims as defining how many leading dimensions are "untouched" by the indexing. When you set batch_dims=2:
- TensorFlow treats every
(i,j)pair as its own mini-batch. - For each mini-batch, it uses the single index from
indices[i,j]to look up the corresponding slice inpreds[i,j,:,:](which is aC×6tensor). - It then stacks all these mini-batch results back together to form the final
H×W×6tensor.
Without batch_dims, you'd have to generate full 3D indices (including i and j for every position), which is more verbose (we'll show that below for context).
Method 2: Use tf.gather with Dimension Flattening
If you prefer working with tf.gather, you can flatten the first two dimensions, perform the gather, then reshape back. This is a valid approach but requires more dimension manipulation:
# Flatten the first two dimensions of predictions and ground truth preds_flat = tf.reshape(preds, (-1, C, 6)) # Shape: (H*W, C, 6) gt_classes_flat = tf.reshape(gt_classes, (-1,)) # Shape: (H*W,) # Gather along the C dimension output_flat = tf.gather(preds_flat, gt_classes_flat, axis=1) # Shape: (H*W, 6) # Reshape back to original spatial dimensions output = tf.reshape(output_flat, (H, W, 6)) print(output.shape) # Output: (2, 3, 6)
Method 3: tf.gather_nd Without batch_dims (For Context)
Just to show what batch_dims saves you from, here's how you'd do it by generating full indices for every (i,j,class) position:
# Generate indices for H and W dimensions i_indices = tf.tile(tf.range(H)[:, None], [1, W]) # Shape: (H, W) j_indices = tf.tile(tf.range(W)[None, :], [H, 1]) # Shape: (H, W) # Stack to form full 3D indices full_indices = tf.stack([i_indices, j_indices, gt_classes], axis=-1) # Shape: (H, W, 3) # Gather using full indices output = tf.gather_nd(preds, full_indices) print(output.shape) # Output: (2, 3, 6)
As you can see, this is more code and less efficient than using batch_dims, since you have to manually create the spatial indices.
Final Recommendation
Stick with Method 1 using tf.gather_nd and batch_dims=2. It's clean, efficient, and directly leverages TensorFlow's indexing capabilities without unnecessary reshaping.
内容的提问来源于stack exchange,提问作者A Tyshka

