如何构建仅包含特定类别样本的TensorFlow数据集管道?
构建仅包含特定类别的TensorFlow数据集管道
当然有可行的实现方法,下面以CIFAR-10数据集为例,提供两种常用的方案:
方案一:先在Numpy层面筛选样本,再构建数据集管道
适合内存足够容纳筛选后数据的场景,操作简单直观:
import tensorflow as tf # 加载完整CIFAR-10数据集 (x_train, y_train), (x_test, y_test) = tf.keras.datasets.cifar10.load_data() # 指定要保留的类别索引,示例为前5个类别(0-4) target_classes = [0, 1, 2, 3, 4] # 生成训练集筛选掩码:找出标签属于目标类别的样本位置 train_mask = tf.reduce_any(tf.equal(y_train, tf.constant(target_classes, dtype=tf.uint8)[:, None]), axis=0) # 提取筛选后的训练集数据 x_train_filtered = x_train[train_mask.numpy().flatten()] y_train_filtered = y_train[train_mask.numpy().flatten()] # 对测试集执行同样的筛选操作 test_mask = tf.reduce_any(tf.equal(y_test, tf.constant(target_classes, dtype=tf.uint8)[:, None]), axis=0) x_test_filtered = x_test[test_mask.numpy().flatten()] y_test_filtered = y_test[test_mask.numpy().flatten()] # 构建最终的数据集管道 train_dataset = tf.data.Dataset.from_tensor_slices((x_train_filtered, y_train_filtered)).batch(64) test_dataset = tf.data.Dataset.from_tensor_slices((x_test_filtered, y_test_filtered)).batch(64)
方案二:使用tf.data.Dataset的filter方法直接过滤
适合大数据集场景,无需将所有筛选后的数据加载到内存,全程基于TensorFlow操作,更高效:
import tensorflow as tf # 加载完整数据集并转为原始Dataset (x_train, y_train), (x_test, y_test) = tf.keras.datasets.cifar10.load_data() train_dataset_full = tf.data.Dataset.from_tensor_slices((x_train, y_train)) test_dataset_full = tf.data.Dataset.from_tensor_slices((x_test, y_test)) # 指定要保留的类别索引 target_classes = tf.constant([0, 1, 2, 3, 4], dtype=tf.uint8) # 定义过滤函数:检查样本标签是否属于目标类别 def filter_fn(image, label): # CIFAR-10的标签是形状为(1,)的张量,先展平为标量 label = tf.squeeze(label) return tf.reduce_any(tf.equal(label, target_classes)) # 过滤数据集并设置批量大小 train_dataset = train_dataset_full.filter(filter_fn).batch(64) test_dataset = test_dataset_full.filter(filter_fn).batch(64)
两种方案都能得到仅包含指定类别的数据集管道,可直接用于Keras模型训练。
内容的提问来源于stack exchange,提问作者tasrif
相关产品推荐
相关产品推荐

