如何制作单图目标数量不一致的标注张量并适配TF模型?
Hey there! I’ve tackled exactly this problem when fine-tuning RetinaNet for variable-count object detection (in my case, pedestrian datasets with 0-15 targets per image), so let’s walk through the key steps to make this work smoothly with the TF Object Detection API.
The core thing to remember: the API is designed to handle variable-length object annotations natively—you just need to set up your data pipeline and config correctly to avoid static graph conflicts.
1. Format Your Dataset as TFRecords with Variable-Length Annotations
First, convert your custom dataset into TFRecord format, where each sample stores its own variable number of bounding boxes and class labels (no need to pad all samples to a fixed count like 11).
Here’s a simplified snippet for writing annotations to TFRecords:
def create_tf_example(image_path, bboxes, classes): # Load image and get basic metadata with tf.io.gfile.GFile(image_path, 'rb') as fid: encoded_jpg = fid.read() image_format = b'jpg' height, width = tf.io.decode_jpeg(encoded_jpg).shape[:2] # Normalize bounding box coordinates to [0,1] range xmins = [bbox[0]/width for bbox in bboxes] ymins = [bbox[1]/height for bbox in bboxes] xmaxs = [bbox[2]/width for bbox in bboxes] ymaxs = [bbox[3]/height for bbox in bboxes] # Build TF Example with variable-length features tf_example = tf.train.Example(features=tf.train.Features(feature={ 'image/encoded': tf.train.Feature(bytes_list=tf.train.BytesList(value=[encoded_jpg])), 'image/format': tf.train.Feature(bytes_list=tf.train.BytesList(value=[image_format])), 'image/height': tf.train.Feature(int64_list=tf.train.Int64List(value=[height])), 'image/width': tf.train.Feature(int64_list=tf.train.Int64List(value=[width])), 'image/object/bbox/xmin': tf.train.Feature(float_list=tf.train.FloatList(value=xmins)), 'image/object/bbox/ymin': tf.train.Feature(float_list=tf.train.FloatList(value=ymins)), 'image/object/bbox/xmax': tf.train.Feature(float_list=tf.train.FloatList(value=xmaxs)), 'image/object/bbox/ymax': tf.train.Feature(float_list=tf.train.FloatList(value=ymaxs)), 'image/object/class/label': tf.train.Feature(int64_list=tf.train.Int64List(value=classes)), })) return tf_example
Each TFRecord sample only stores the exact number of bboxes/classes present in the image—no padding, no fixed length.
2. Build a Dynamic Input Pipeline with tf.data
When loading TFRecords, use tf.data.Dataset and avoid enforcing static shapes on annotation tensors. The API automatically handles variable-length tensors as RaggedTensors or dynamically shaped tensors in eager mode.
Here’s how to parse TFRecords correctly:
def parse_tfrecord_fn(example): feature_description = { 'image/encoded': tf.io.FixedLenFeature([], tf.string), 'image/format': tf.io.FixedLenFeature([], tf.string), 'image/height': tf.io.FixedLenFeature([], tf.int64), 'image/width': tf.io.FixedLenFeature([], tf.int64), 'image/object/bbox/xmin': tf.io.VarLenFeature(tf.float32), 'image/object/bbox/ymin': tf.io.VarLenFeature(tf.float32), 'image/object/bbox/xmax': tf.io.VarLenFeature(tf.float32), 'image/object/bbox/ymax': tf.io.VarLenFeature(tf.float32), 'image/object/class/label': tf.io.VarLenFeature(tf.int64), } example = tf.io.parse_single_example(example, feature_description) # Convert sparse variable-length features back to dense tensors xmin = tf.sparse.to_dense(example['image/object/bbox/xmin']) ymin = tf.sparse.to_dense(example['image/object/bbox/ymin']) xmax = tf.sparse.to_dense(example['image/object/bbox/xmax']) ymax = tf.sparse.to_dense(example['image/object/bbox/ymax']) bboxes = tf.stack([xmin, ymin, xmax, ymax], axis=-1) classes = tf.sparse.to_dense(example['image/object/class/label']) # Decode and preprocess image image = tf.io.decode_jpeg(example['image/encoded'], channels=3) image = tf.cast(image, tf.float32) # Return data in the API's expected format return { 'image': image, 'groundtruth_boxes': bboxes, 'groundtruth_classes': classes } # Build the training dataset train_dataset = tf.data.TFRecordDataset('train.tfrecord') train_dataset = train_dataset.map(parse_tfrecord_fn, num_parallel_calls=tf.data.AUTOTUNE) train_dataset = train_dataset.shuffle(1000).batch(8) # Batch size can be arbitrary; API handles variable lengths
We use tf.io.VarLenFeature to parse variable-length annotations, then convert them back to dense tensors with dynamic shapes. The batch method automatically handles varying bbox counts per sample in the batch.
3. Adjust Your RetinaNet Pipeline Config
Tweak your pipeline.config file to support variable-length annotations and your custom class count:
- Set
num_classesto 2 (matches your classification tensor shape: background + traffic light) - Point
label_map_pathto alabel_map.pbtxtwith exactly 2 classes:item { id: 1 name: 'traffic_light' } item { id: 0 name: 'background' } - In the
train_configsection, avoid static shape constraints (the API defaults to dynamic handling). Ensurefine_tune_checkpointpoints to your COCO-pre-trained RetinaNet checkpoint.
4. Avoid Graph Overwrites with TF2 Eager Mode or Proper tf.function Signatures
If you’re using TF2 (the recommended version for the latest Object Detection API), default eager execution eliminates static graph conflicts entirely—no need to worry about overwriting graphs since they’re built dynamically per step.
If you wrap your training loop in tf.function for performance, specify dynamic input signatures to allow variable shapes:
@tf.function(input_signature=[ tf.TensorSpec(shape=(None, None, None, 3), dtype=tf.float32), # (batch, height, width, 3) tf.TensorSpec(shape=(None, None, 4), dtype=tf.float32), # (batch, num_objects, 4) tf.TensorSpec(shape=(None, None), dtype=tf.int64) # (batch, num_objects) ]) def train_step(image, gt_boxes, gt_classes): with tf.GradientTape() as tape: predictions = model(image) loss = compute_loss(predictions, gt_boxes, gt_classes) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) return loss
The None values in tensor specs tell TensorFlow to accept any dimension size, so variable object counts won’t trigger graph rebuilds.
Key Pitfalls to Avoid
- Don’t pad annotations to a fixed length: This introduces unnecessary background labels and confuses the model. The API handles variable-length natively.
- Don’t hardcode tensor shapes: Always use dynamic shapes (e.g.,
tf.shape(bboxes)[1]instead of a fixed number like 11) when accessing annotation counts. - Use the latest API version: Older versions had better support for static shapes, but TF2+ versions are optimized for dynamic inputs.
内容的提问来源于stack exchange,提问作者Sergei

