TFF联邦学习场景下如何为各客户端加载本地数据集而非模拟数据
实现逻辑
TFF远程执行模式下,客户端侧的计算逻辑会自动分发到各个worker节点本地执行,你不需要在服务端提前加载所有客户端数据,只需要把数据加载逻辑封装为可在客户端执行的计算单元,由各个客户端本地完成数据读取即可。
具体实现步骤
1. 封装客户端本地数据加载逻辑
编写被@tff.tf_computation装饰的加载函数,该函数会在每个客户端节点本地运行,直接读取客户端本地的sqlite文件:
import tensorflow_federated as tff import tensorflow as tf # 该函数仅在客户端本地执行,不会在服务端运行 @tff.tf_computation def load_local_client_data(): # 读取当前客户端本地的sqlite文件,确保所有客户端的数据集都放在这个路径下 database_path = '/root/.tff/emnist_all.sqlite' client_data = tff.simulation.datasets.SqlClientData(database_path, 'digits_only_train') # 执行你自定义的预处理逻辑 preprocessed_data = client_data.preprocess(_add_proto_parsing) # 返回当前客户端需要用到的数据集,单客户端单用户可以直接取对应用户的数据集,多用户可以合并返回 return preprocessed_data.create_tf_dataset_from_all_clients()
2. 构造联邦数据集
通过federated_eval把加载逻辑分发到所有客户端执行,批量获取各客户端本地的数据集:
@tff.federated_computation def get_federated_dataset(): # 将加载逻辑分发到所有客户端执行,返回客户端放置的联邦数据集 return tff.federated_eval(load_local_client_data, tff.CLIENTS) # 调用即可得到所有客户端本地加载完成的预处理数据集 preprocessed_data_for_clients = get_federated_dataset()
3. 接入原有流程
得到的preprocessed_data_for_clients可以直接传入你原来的训练、评估逻辑,无需额外修改:
# 原有逻辑直接复用即可 print('### GET CHANNELS') # set up the remote executors channels = get_channels(list_host) tff.backends.native.set_remote_execution_context(channels) print('### TRAINER') trainer = tff.learning.build_federated_averaging_process(model_fn, client_optimizer_fn=lambda: tf.keras.optimizers.SGD(learning_rate=0.01)) print('### EVALUATE') evaluate(trainer, preprocessed_data_for_clients)
注意事项
- 所有客户端节点必须提前将自有数据集放在代码指定的路径,各节点的数据集仅存储自身需要处理的分区即可,不需要存储全量数据
- 数据预处理逻辑
_add_proto_parsing必须为纯TensorFlow实现,不能依赖服务端的本地变量、资源,否则分发到客户端执行时会出现序列化错误 - 如果需要按客户端标识加载对应分区,可以给加载函数增加
client_id入参,将客户端ID列表作为联邦输入传入federated_map,每个客户端会拿到自身对应的ID来加载匹配的数据集
内容的提问来源于stack exchange,提问作者David González
相关产品推荐
相关产品推荐

