You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

  1. List all class directories: Grab paths to each class folder.
  2. Create pre-sampled datasets per class: For each directory, load its TFRecords, parse them, and keep only 5 images.
  3. Build a dataset of class datasets: Treat each pre-sampled class dataset as a single element.
  4. 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

  1. List all class directories: Same as before.
  2. Create a dataset of class directories: Treat each directory path as an element.
  3. 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 shuffle parameter in list_files to control whether you get fixed or random samples per class.
  • Performance: Use num_parallel_calls=tf.data.AUTOTUNE to let TensorFlow optimize parallel loading/parsing.
  • Preprocessing: Add any image augmentation steps inside parse_tfrecord_fn or after creating the dataset (using map).

内容的提问来源于stack exchange,提问作者Siavash

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.27 10:09:32