如何按client_id过滤TensorFlow的ZipDataset数据集
TensorFlow ZipDataset 过滤操作实现
问题复现背景
- 基于
image_dataset_from_directory加载4分类图像数据集,原始单条数据格式为(image_tensor, label)元组 - 生成与数据集长度匹配的
client_id序列并转换为tf.data.Dataset类型,通过以下代码完成两个数据集拼接:
dataset = tf.data.Dataset.zip((dataset, client_id))
- 拼接后数据集的元素签名为:
<ZipDataset element_spec=((TensorSpec(shape=(128, 128, 3), dtype=tf.float32, name=None), TensorSpec(shape=(), dtype=tf.int32, name=None)), TensorSpec(shape=(), dtype=tf.int64, name=None))>
此时单条样本的结构为嵌套元组:((图像张量, 类别标签), client_id)
- 尝试使用如下代码过滤
client_id=15的样本时触发报错:
dataset = dataset.filter( x : x[1]==15)
报错信息:
TypeError: 'ZipDataset' object is not subscriptable
- 直接遍历取单条样本时可以正常读取client_id值:
for x in dataset.take(1): print(x[1]) # 输出: tf.Tensor(15, shape=(), dtype=int64)
报错根因
两处写法错误导致异常:
- 过滤函数语法不合法:
x : x[1]==15不是合法的Python可调用对象,既没有加lambda关键字声明匿名函数,也没有用def定义正式函数,Python解析代码时会将ZipDataset实例本身作为参数传入,尝试对Dataset对象做下标取值[1],直接触发不可下标访问的错误。 - 图模式下Tensor值判断不规范:即使补全lambda语法,直接用Python原生
==判断Tensor值在静态图执行模式下可能出现兼容问题,应当使用TensorFlow内置的数值判断API。
正确实现方案
使用标准lambda声明过滤逻辑,搭配tf.equal完成Tensor值判断,代码如下:
import tensorflow as tf # 过滤client_id等于15的所有样本 dataset = dataset.filter(lambda sample: tf.equal(sample[1], 15))
如果需要把嵌套的元组结构展平为(图像张量, 类别标签, client_id)的平铺格式,方便后续处理,可以追加map转换:
dataset = dataset.filter(lambda sample: tf.equal(sample[1], 15)) .map(lambda sample: (sample[0][0], sample[0][1], sample[1]))
内容的提问来源于stack exchange,提问作者Los
相关产品推荐
相关产品推荐

