读取TF Object Detection API训练TFRecord文件时Reshape错误如何解决?
Hey there, let's work through this shape mismatch error step by step—this is a common gotcha when working with TFRecords for object detection, so we can pin it down quickly. The error tells us your tensor has 91090 values, but the code expects 921600 (which is 640x480x3, for reference). Here's how to diagnose and fix it:
1. First, Check Your TFRecord Reading Pipeline
The most likely culprit is hardcoding an image shape in your parsing logic that doesn't match the actual images stored in the TFRecord. Here's what to do:
Stop hardcoding
reshapedimensions: If your parsing function looks like this (where you force a fixed shape), that's probably the issue:def parse_tfrecord(example_proto): features = tf.io.parse_single_example(example_proto, { 'image': tf.io.FixedLenFeature([], tf.string), # Other features (bboxes, labels, etc.) }) image = tf.io.decode_jpeg(features['image'], channels=3) # ❌ This hardcoded shape will fail if images aren't exactly 640x480x3 image = tf.reshape(image, (640, 480, 3)) return imageStore and use image metadata in your TFRecords: When creating your TFRecords, you should save the original height, width, and channel count of each image. Then use those values to reshape dynamically during reading:
# Updated parsing function with metadata def parse_tfrecord(example_proto): features = tf.io.parse_single_example(example_proto, { 'image': tf.io.FixedLenFeature([], tf.string), 'height': tf.io.FixedLenFeature([], tf.int64), 'width': tf.io.FixedLenFeature([], tf.int64), 'channels': tf.io.FixedLenFeature([], tf.int64), # Other features }) image = tf.io.decode_jpeg(features['image'], channels=features['channels']) # ✅ Use stored metadata to reshape correctly image = tf.reshape(image, (features['height'], features['width'], features['channels'])) # If you need a fixed size for detection, use tf.image.resize instead of reshape! # image = tf.image.resize(image, (640, 480)) return image
2. Verify Your TFRecord Creation Logic
If adjusting the reading pipeline doesn't fix it, the error might be in how you built the TFRecords in the first place:
- Double-check image metadata during creation: When writing images to TFRecords, make sure you're capturing the correct height, width, and channels. For example:
Ensure you're not accidentally writing incorrect dimensions (e.g., swapping height and width) or using corrupted images that don't decode properly.def create_tf_example(image_path): # Read image bytes with open(image_path, 'rb') as f: image_bytes = f.read() # Get actual image dimensions image = tf.io.decode_jpeg(image_bytes) height, width, channels = image.shape # Build feature dict feature = { 'image': tf.train.Feature(bytes_list=tf.train.BytesList(value=[image_bytes])), 'height': tf.train.Feature(int64_list=tf.train.Int64List(value=[height])), 'width': tf.train.Feature(int64_list=tf.train.Int64List(value=[width])), 'channels': tf.train.Feature(int64_list=tf.train.Int64List(value=[channels])), # Add bbox/labels here } return tf.train.Example(features=tf.train.Features(feature=feature))
3. Debug a Single TFRecord Sample
To get concrete data, write a quick script to inspect one sample from your TFRecord. This will tell you exactly what's stored vs. what's expected:
import tensorflow as tf # Load one record dataset = tf.data.TFRecordDataset('your_training_record.tfrecord') for raw_record in dataset.take(1): example = tf.train.Example() example.ParseFromString(raw_record.numpy()) # Print stored metadata (if available) if 'height' in example.features.feature: print(f"Stored height: {example.features.feature['height'].int64_list.value[0]}") print(f"Stored width: {example.features.feature['width'].int64_list.value[0]}") print(f"Stored channels: {example.features.feature['channels'].int64_list.value[0]}") # Decode and check actual image shape image_bytes = example.features.feature['image'].bytes_list.value[0] image = tf.io.decode_jpeg(image_bytes) print(f"\nActual image shape: {image.shape}") print(f"Total elements in tensor: {tf.size(image).numpy()}")
Run this—you'll see if the actual image dimensions don't match the ones your code is trying to reshape to, or if the stored metadata is wrong.
4. Check TF Object Detection API Config
If you're using the official TFOD API pipeline config, make sure the image_resizer setting aligns with your data:
# In your pipeline.config image_resizer { fixed_shape_resizer { height: 640 width: 480 } }
The API uses this to resize images to a fixed size for detection. If your images are corrupted or can't be decoded, this step might throw a reshape error. Ensure the resizer dimensions make sense for your dataset.
内容的提问来源于stack exchange,提问作者Thomas Fauskanger

