TensorFlow中如何用filter按标签过滤image_dataset_from_directory生成的数据集?
问题解决:TensorFlow数据集按标签过滤的正确方式
错误原因
你遇到的ValueError是因为tf.data.Dataset.filter()要求传入的predicate必须返回标量布尔张量(用来判断是否保留整个batch),但你的代码中y是形状为(32,)的batch级张量(对应默认batch_size=32),y==0会生成一个同样形状的布尔数组,不是标量,因此不符合要求。
解决方案
方法1:先拆分样本再过滤(推荐)
先通过unbatch()将批量数据集拆分为单个样本,此时每个样本的标签y是标量张量,y==0会返回标量布尔值,符合filter()的要求,最后再重新组合成批量:
import tensorflow as tf SIZE = 224 # 替换成你的图像尺寸 full_ds = tf.keras.utils.image_dataset_from_directory( 'the_path', image_size=(SIZE, SIZE), ) # 拆分单个样本 -> 过滤标签为0的样本 -> 重新组成批量 fibrosis_ds = full_ds.unbatch().filter(lambda x, y: y == 0).batch(32)
方法2:批量内过滤样本(保留批量结构)
如果不想拆分批量,可以通过map()配合tf.boolean_mask提取每个batch内标签为0的样本,之后再统一调整批量大小:
def filter_samples_in_batch(x, y): # 生成标签为0的掩码 mask = tf.equal(y, 0) # 提取符合条件的图像和标签 filtered_x = tf.boolean_mask(x, mask) filtered_y = tf.boolean_mask(y, mask) return filtered_x, filtered_y # 先过滤每个batch内的目标样本,再拆分重组批量 fibrosis_ds = full_ds.map(filter_samples_in_batch).unbatch().batch(32)
验证方法
可以通过循环打印过滤后的数据集标签,确认只保留了目标类别:
for x, y in fibrosis_ds: print(y) break
输出应该全为0的张量,比如:tf.Tensor([0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0], shape=(32,), dtype=int32)
内容的提问来源于stack exchange,提问作者honkskillet
相关产品推荐
相关产品推荐

