如何将NumPy数组格式的自定义数据集转为TFRecord用于TensorFlow目标检测
Got it, let's walk through how to convert your NumPy-based dataset into TFRecords for the TensorFlow Object Detection API. I’ve done this plenty of times with large datasets, so here’s a step-by-step guide tailored exactly to your setup:
First, make sure you have the required libraries installed. You’ll need tensorflow and numpy (which you probably already have, but just in case):
pip install tensorflow numpy
TFRecords store data as tf.train.Example protos, so we need a function to convert each image-label pair into this format. Since your images are flattened uint8 arrays, we’ll convert them to bytes and optionally store their shape to make reshaping easier later.
import tensorflow as tf import numpy as np def serialize_example(image_array, label): # Convert the flattened uint8 image array to raw bytes image_bytes = image_array.tobytes() # Store the shape of the flattened array (or your original image shape if you know it) # If you know your original image dimensions (e.g., 152x152), replace image_array.shape with (152,152) image_shape = image_array.shape # Build the feature dictionary feature = { 'image': tf.train.Feature(bytes_list=tf.train.BytesList(value=[image_bytes])), 'label': tf.train.Feature(int64_list=tf.train.Int64List(value=[label])), 'image_shape': tf.train.Feature(int64_list=tf.train.Int64List(value=image_shape)) } # Create and serialize the Example proto example_proto = tf.train.Example(features=tf.train.Features(feature=feature)) return example_proto.SerializeToString()
Now, we’ll loop through your entire dataset, serialize each sample, and write them to a TFRecord file. For 40k samples, this should run smoothly, but if you hit memory issues, you can split your NumPy arrays into chunks and write multiple TFRecord files (e.g., train_00.tfrecord, train_01.tfrecord).
def generate_tfrecord(images, labels, output_path): # Open a TFRecord writer with tf.io.TFRecordWriter(output_path) as writer: for idx in range(len(images)): image = images[idx] label = labels[idx] # Optional: Validate data to catch errors early assert image.dtype == np.uint8, f"Image at index {idx} is not uint8 (got {image.dtype})" assert 0 <= label <= 3, f"Label at index {idx} is out of range 0-3 (got {label})" # Serialize and write the sample serialized_sample = serialize_example(image, label) writer.write(serialized_sample) # Print progress every 1000 samples if (idx + 1) % 1000 == 0: print(f"Processed {idx+1}/{len(images)} samples") # Usage: Replace with your actual NumPy arrays # train_images = np.load("your_train_images.npy") # Shape [40000, 23456] # train_labels = np.load("your_train_labels.npy") # Shape [40000] # generate_tfrecord(train_images, train_labels, "train.tfrecord")
Don’t skip this step—confirm that your TFRecord is valid and contains the correct data by parsing it back:
def parse_tfrecord_sample(example_proto): # Define the feature schema to match what we wrote feature_schema = { 'image': tf.io.FixedLenFeature([], tf.string), 'label': tf.io.FixedLenFeature([], tf.int64), 'image_shape': tf.io.FixedLenFeature([2], tf.int64) } # Parse the proto parsed_features = tf.io.parse_single_example(example_proto, feature_schema) # Convert image bytes back to a uint8 array and reshape image = tf.io.decode_raw(parsed_features['image'], tf.uint8) image = tf.reshape(image, parsed_features['image_shape']) # Cast label to int32 (optional, matches most model input requirements) label = tf.cast(parsed_features['label'], tf.int32) return image, label # Test the parser on your TFRecord raw_dataset = tf.data.TFRecordDataset("train.tfrecord") parsed_dataset = raw_dataset.map(parse_tfrecord_sample) # Check the first 3 samples for image, label in parsed_dataset.take(3): print(f"Sample shape: {image.shape}, Label: {label.numpy()}, Dtype: {image.dtype}")
- Image Shape: You mentioned your images are flattened to 23456 elements—make sure you know the original 2D shape (e.g., 154x154 doesn’t add up, so double-check how you flattened them). Storing
image_shapein the TFExample lets you dynamically reshape later without hardcoding. - Object Detection vs. Classification: Wait, the Object Detection API is typically used for bounding box tasks, but your labels are 0-3 (classification). If you’re actually doing detection (and forgot to mention bounding boxes), you’ll need to add box coordinates to the
featuredictionary inserialize_example. If it’s classification, the API still works—just adjust your model config to use a classification head. - Performance: Using
tf.data.TFRecordDatasetwith parallel parsing will speed up training significantly compared to loading NumPy arrays directly, especially for 40k samples. You can addnum_parallel_calls=tf.data.AUTOTUNEto themapfunction for extra speed.
内容的提问来源于stack exchange,提问作者Govinda Malavipathirana

