任何tff.Computation(如next)将始终运行整个指定的计算。例如,如果您的tff.templates.IterativeProcess 是tff.learning.build_federated_averaging_process 的结果,则其next 函数将代表一轮联合平均算法。
联合平均算法在每个本地数据集上运行固定数量的 epochs(为简单起见,假设为 1)的训练,并按顺序在服务器上以数据加权的方式对模型更新进行平均完成一轮 - 请参阅 Algorithm 1 in the original federated averaging paper 了解算法规范。
现在,关于 TFF 如何表示和执行该算法。在build_federated_averaging_process 的文档中,next 函数具有类型签名:
(<S@SERVER, {B*}@CLIENTS> -> <S@SERVER, T@SERVER>)
TFF 的类型系统将数据集表示为tff.SequenceType(这是上面* 的含义),因此类型签名的参数中的第二个元素表示具有元素的数据集的集合(技术上是多重集) B 类型的,放置在客户端。
这在您的示例中的含义如下。您有一个tf.data.Datasets 列表,每个列表代表每个客户端上的本地数据——您可以将列表视为代表联合放置。在这种情况下,TFF 执行整个指定的计算意味着:TFF 将列表中的每个项目视为要在本轮中训练的客户端。根据上面链接的算法,您的数据集列表表示集合 S_t。
TFF 将忠实地执行一轮联合平均算法,列表中的Dataset 元素代表为这一轮选择的客户。训练将在每个客户端上运行一个 epoch(并行);如果数据集具有不同数量的数据,那么每个客户端的训练可能在不同时间完成是正确的。然而,这是单轮联合平均算法的正确语义,而不是像 Reptile 这样的类似算法的参数化,它为每个客户端运行固定数量的步骤。
如果您希望选择一个客户端子集来运行一轮训练,这应该在调用 TFF 之前在 Python 中完成,例如:
state = iterative_process.initialize()
# ls is list of datasets
sampled_clients = random.sample(ls, N_CLIENTS)
state = iterative_process.next(state, sampled_clients)
通常,您可以将 Python 运行时视为“实验驱动程序”层——例如,任何客户端选择都应该发生在这一层。有关更多详细信息,请参阅this answer 的开头。