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

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图执行模式的实现逻辑直接相关,核心触发原因如下:

  1. tf.py_function 不兼容静态图下的批量数据集拼接逻辑:create_tf_dataset_from_all_clients() 的底层实现是遍历所有传入的client_id,分别生成每个客户端的子数据集后做静态图层面的拼接合并。但包裹在tf.py_function中的Python原生列表推导extract_file_paths不会随每个client_id的传入动态执行,只会在构图阶段捕获单次执行返回的路径值,导致所有客户端的子数据集路径列表异常,最终合并去重后仅保留1个有效文件。
  2. 单客户端联邦场景运行正常属于执行模式差异导致的特例:单独读取单个客户端数据集时运行在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 12:18:14