create_tf_dataset_from_all_clients()生成集中式数据集仅含单文件问题
联邦数据集转集中式数据集仅返回单文件问题排查与修复
问题背景
- 待处理原始数据集包含三类字段:
path(文件存储路径)、client_id(客户端归属ID)、label(样本标签) - 实现目标:基于分客户端的联邦数据构建TFF
ClientData对象,再合并生成覆盖所有客户端样本的集中式数据集 - 异常表现:联邦场景下单独读取单个客户端的数据集逻辑正常,调用
create_tf_dataset_from_all_clients()生成集中式数据集时,最终输出仅包含1个文件,不符合多文件、多标签的预期。
原实现代码
客户端数据集构建函数:
def extract_file_paths(dataset): return [item["path"] for item in dataset] @tf.function def create_dataset(client_id): new_datset = tf.data.Dataset.from_tensor_slices(dict(df_aml)) client_id = int(client_id) client_id = tf.cast(client_id, dtype=tf.int64) files = new_datset.filter(lambda x: x['client_id'] == client_id) list_ds = tf.data.Dataset.list_files(tf.py_function(func=extract_file_paths,inp=[files], Tout = tf.string )) images_ds = list_ds.map(parse_image) return images_ds
ClientData与集中式数据集构建逻辑:
client_ids = ['0', '1', '2'] client_data = tff.simulation.datasets.ClientData.from_clients_and_tf_fn(client_ids, create_dataset) centralized = client_data.create_tf_dataset_from_all_clients()
问题根因
该问题确实与TensorFlow图执行模式的实现逻辑直接相关,核心触发原因如下:
tf.py_function不兼容静态图下的批量数据集拼接逻辑:create_tf_dataset_from_all_clients()的底层实现是遍历所有传入的client_id,分别生成每个客户端的子数据集后做静态图层面的拼接合并。但包裹在tf.py_function中的Python原生列表推导extract_file_paths不会随每个client_id的传入动态执行,只会在构图阶段捕获单次执行返回的路径值,导致所有客户端的子数据集路径列表异常,最终合并去重后仅保留1个有效文件。- 单客户端联邦场景运行正常属于执行模式差异导致的特例:单独读取单个客户端数据集时运行在eager即时执行模式下,
tf.py_function可以即时执行拉取对应客户端的全量路径,不会触发静态图下的逻辑失效问题。
修复方案
移除tf.py_function嵌套Python遍历的写法,全程使用TensorFlow原生数据集算子实现路径过滤逻辑,保证代码可以在静态图模式下正确按客户端ID过滤数据,修改后的代码如下:
def create_dataset(client_id): # 加载全量数据的路径与客户端ID映射 full_ds = tf.data.Dataset.from_tensor_slices( (df_aml["path"].values, df_aml["client_id"].values) ) # 统一client_id类型,用TF原生filter算子过滤当前客户端的路径 client_id = tf.cast( tf.strings.to_number(client_id, out_type=tf.int64), dtype=tf.int64 ) client_path_ds = full_ds.filter( lambda path, cid: cid == client_id ).map(lambda path, cid: path) # 读取路径对应文件并做预处理 images_ds = client_path_ds.interleave( lambda path: tf.data.Dataset.list_files(path), num_parallel_calls=tf.data.AUTOTUNE ).map(parse_image, num_parallel_calls=tf.data.AUTOTUNE) return images_ds
验证步骤
- 先逐一遍历
client_data下每个客户端ID对应的子数据集,统计样本数量,确认单客户端数据量符合预期 - 再调用
create_tf_dataset_from_all_clients()生成集中式数据集,统计总样本量,确认总样本数等于所有客户端样本数之和 - 调试阶段可暂时移除
create_dataset上的@tf.function装饰器,确认逻辑跑通后再加回装饰器做性能优化
内容的提问来源于stack exchange,提问作者ana
相关产品推荐
相关产品推荐

