如何用tf.data.Dataset从TFRecord目录实现指定类与样本采样
Hey there! Let's tackle your two sampling requirements using pure tf.data.Dataset operations, and get you that 25-image batch (5 per class) you need. First, let's align on your data structure: you've got one directory per class, each filled with .tfrecords files (each file is a single image). Great, let's break this down step by step.
First, Define a TFRecord Parsing Function
We'll start with a reusable function to parse your .tfrecords files. Adjust the feature description to match how you stored your image data:
import tensorflow as tf import os def parse_tfrecord_fn(example): # Customize this to match your TFRecord schema feature_description = { 'image': tf.io.FixedLenFeature([], tf.string), 'label': tf.io.FixedLenFeature([], tf.int64), # Add other features like height/width if you stored them } # Parse the example parsed = tf.io.parse_single_example(example, feature_description) # Decode and preprocess the image (resize to your target shape) image = tf.io.decode_jpeg(parsed['image'], channels=3) image = tf.image.resize(image, (224, 224)) # Adjust size as needed label = parsed['label'] return image, label
Requirement 1: Sample 5 Images from Every Class, Then Return Batches of 5 Classes × 5 Images
This approach pre-samples 5 images per class first, then creates batches by picking 5 random classes each time.
Step-by-Step Implementation
- List all class directories: Grab paths to each class folder.
- Create pre-sampled datasets per class: For each directory, load its TFRecords, parse them, and keep only 5 images.
- Build a dataset of class datasets: Treat each pre-sampled class dataset as a single element.
- Sample 5 classes and combine into a 25-image batch: Shuffle the class datasets, pick 5, concatenate their images, and batch them.
# 1. Define your dataset root and get class directories data_root = "/path/to/your/dataset_root" class_dirs = [ os.path.join(data_root, dir_name) for dir_name in os.listdir(data_root) if os.path.isdir(os.path.join(data_root, dir_name)) ] # 2. Create a pre-sampled dataset for each class (5 images per class) class_datasets = [] for dir_path in class_dirs: # List all TFRecords in the class directory tfrecord_files = tf.data.Dataset.list_files( os.path.join(dir_path, "*.tfrecords"), shuffle=False # Set to True if you want random 5 images each run ) # Load, parse, and take 5 images class_ds = tfrecord_files.interleave( lambda file_path: tf.data.TFRecordDataset(file_path).map(parse_tfrecord_fn), num_parallel_calls=tf.data.AUTOTUNE ).take(5) class_datasets.append(class_ds) # 3. Create a dataset that holds all pre-sampled class datasets dataset_of_classes = tf.data.Dataset.from_tensor_slices(class_datasets) # 4. Combine 5 random classes into a 25-image batch def combine_5_classes(class_ds_list): # Concatenate the 5 class datasets (each with 5 images) combined = class_ds_list[0] for ds in class_ds_list[1:]: combined = combined.concatenate(ds) return combined.batch(25) # Exact 25 images per batch final_ds = dataset_of_classes.shuffle(len(class_dirs)).batch(5).flat_map(combine_5_classes) # Test the iterator for batch_images, batch_labels in final_ds: print(f"Batch shape: {batch_images.shape}") # Output: (25, 224, 224, 3) print(f"Label shape: {batch_labels.shape}") # Output: (25,) # Verify 5 samples per class (optional) unique_labels, counts = tf.unique_with_counts(batch_labels) print(f"Class counts: {counts.numpy()}") # Should show [5,5,5,5,5] break
Requirement 2: First Sample 5 Random Classes, Then Sample 5 Images from Each
This approach dynamically picks 5 random classes first, then samples 5 images from each selected class (great if you want fresh samples from classes each batch).
Step-by-Step Implementation
- List all class directories: Same as before.
- Create a dataset of class directories: Treat each directory path as an element.
- Sample 5 directories and load 5 images from each: Shuffle the directories, pick 5, load and sample 5 images from each, then combine into a 25-image batch.
# 1. Reuse the class directories from above data_root = "/path/to/your/dataset_root" class_dirs = [ os.path.join(data_root, dir_name) for dir_name in os.listdir(data_root) if os.path.isdir(os.path.join(data_root, dir_name)) ] # 2. Create a dataset of class directory paths class_dirs_ds = tf.data.Dataset.from_tensor_slices(class_dirs) # 3. Function to load 5 random images from a single class directory def get_5_samples_from_class(dir_path): # List TFRecords with shuffle=True to get random samples each time tfrecord_files = tf.data.Dataset.list_files( os.path.join(dir_path, "*.tfrecords"), shuffle=True ) # Load, parse, and take 5 random images return tfrecord_files.interleave( lambda file_path: tf.data.TFRecordDataset(file_path).map(parse_tfrecord_fn), num_parallel_calls=tf.data.AUTOTUNE ).take(5) # 4. Combine 5 sampled classes into a 25-image batch final_ds = class_dirs_ds.shuffle(len(class_dirs)).batch(5).flat_map( lambda selected_dirs: tf.data.Dataset.concatenate( get_5_samples_from_class(selected_dirs[0]), get_5_samples_from_class(selected_dirs[1]) ).concatenate(get_5_samples_from_class(selected_dirs[2])) .concatenate(get_5_samples_from_class(selected_dirs[3])) .concatenate(get_5_samples_from_class(selected_dirs[4])) .batch(25) ) # Test the iterator for batch_images, batch_labels in final_ds: print(f"Batch shape: {batch_images.shape}") # Output: (25, 224, 224, 3) print(f"Label shape: {batch_labels.shape}") # Output: (25,) break
Key Notes
- Randomness: Adjust the
shuffleparameter inlist_filesto control whether you get fixed or random samples per class. - Performance: Use
num_parallel_calls=tf.data.AUTOTUNEto let TensorFlow optimize parallel loading/parsing. - Preprocessing: Add any image augmentation steps inside
parse_tfrecord_fnor after creating the dataset (usingmap).
内容的提问来源于stack exchange,提问作者Siavash

