联邦学习

概览

本文档介绍了旨在促进联邦学习任务(例如使用现有的 TensorFlow 机器学习模型进行联邦训练或评估)的接口。在设计这些接口时,我们的首要目标是让用户能够在无需了解底层实现机制的情况下尝试联邦学习,并能够在各种现有模型和数据上评估所实现的联邦学习算法。我们鼓励您为平台做出贡献。TFF 的设计充分考虑了可扩展性和可组合性,我们欢迎您的贡献;我们非常期待看到您的创意!

该层提供的接口由以下三个关键部分组成:

  • 模型 (Models)。通过这些类和辅助函数,您可以封装现有的模型以便在 TFF 中使用。封装模型非常简单,只需调用单个封装函数(例如 tff.learning.models.from_keras_model),或者定义 tff.learning.models.VariableModel 接口的子类以获得完全的可定制性。

  • 联邦计算构建器 (Federated Computation Builders)。这些辅助函数使用您现有的模型构建用于训练或评估的联邦计算。

  • 数据集 (Datasets)。您可以下载并从 Python 中访问这些现成的数据集,用于模拟联邦学习场景。尽管联邦学习是为无法简单地在中心位置下载的去中心化数据而设计的,但在研究和开发阶段,使用可以在本地下载和操作的数据进行初步实验通常很方便,特别是对于刚接触该方法的开发者而言。

除了归入 tff.simulation 中的研究数据集和其他与模拟相关的功能外,这些接口主要定义在 tff.learning 命名空间中。该层是使用 联邦核心 (FC) 提供的底层接口实现的,FC 同时还提供了一个运行时环境。

在继续之前,我们建议您先阅读关于 图像分类文本生成 的教程,因为它们通过具体示例介绍了此处描述的大多数概念。如果您有兴趣进一步了解 TFF 的工作原理,可以粗略浏览 自定义算法 教程,作为对我们用于表达联邦计算逻辑的底层接口的介绍,并研究 tff.learning 接口的现有实现。

模型

架构假设

序列化

TFF 旨在支持各种分布式学习场景,在这些场景中,您编写的机器学习模型代码可能正在大量具有不同能力的异构客户端上执行。虽然在应用谱系的一端,这些客户端可能是强大的数据库服务器,但我们平台旨在支持的许多重要用例涉及资源受限的移动设备和嵌入式设备。我们不能假设这些设备具备托管 Python 运行时的能力;目前我们只能假设它们具备托管本地 TensorFlow 运行时的能力。因此,我们在 TFF 中做出的一个基本架构假设是:您的模型代码必须能够序列化为 TensorFlow 图。

您仍然可以(也应该)遵循最新的最佳实践(如使用即时执行模式/eager mode)来开发 TF 代码。然而,最终代码必须是可序列化的(例如,可以包装为即时执行代码的 tf.function)。这确保了执行时所需的任何 Python 状态或控制流都可以被序列化(可能需要 Autograph 的帮助)。

目前,TensorFlow 尚未完全支持序列化和反序列化即时执行模式的 TensorFlow。因此,TFF 中的序列化目前遵循 TF 1.0 模式,即所有代码必须在 TFF 控制的 tf.Graph 内构建。这意味着目前 TFF 无法直接使用已构建的模型;相反,模型定义逻辑被打包在一个不带参数的函数中,该函数返回一个 tff.learning.models.VariableModel。然后 TFF 调用此函数以确保模型的所有组件都被序列化。此外,作为一个强类型环境,TFF 将需要一些额外的元数据,例如模型输入类型的规范。

聚合

我们强烈建议大多数用户使用 Keras 构建模型,请参阅下方的“Keras 转换器”部分。这些包装器会自动处理模型更新的聚合以及为模型定义的任何指标。不过,了解如何为通用的 tff.learning.models.VariableModel 处理聚合仍然很有帮助。

联邦学习中始终至少有两层聚合:本地设备端聚合和跨设备(或联邦)聚合。

  • 本地聚合。此级别的聚合是指跨单个客户端拥有的多个样本批次进行聚合。它既适用于模型参数(变量,在模型本地训练时会持续演变),也适用于您计算的统计数据(例如平均损失、准确率和其他指标,这些指标会随模型迭代每个客户端的本地数据流而更新)。

    在此级别执行聚合是您的模型代码的职责,通过标准 TensorFlow 结构完成。

    处理的总体结构如下:

    • 模型首先构建 tf.Variable 来保存聚合值,例如批次数或已处理样本数、每批次或每个样本的损失总和等。

    • TFF 多次调用您的 Model 中的 forward_pass 方法,顺序遍历后续的客户端数据批次,这允许您将保存各种聚合结果的变量作为副作用进行更新。

    • 最后,TFF 调用您的 Model 中的 report_local_unfinalized_metrics 方法,允许您的模型将其收集的所有汇总统计数据编译成一组紧凑的指标,由客户端导出。例如,您的模型代码可以在此处将损失之和除以已处理的样本数量,以导出平均损失等。

  • 联邦聚合。此级别的聚合是指跨系统中的多个客户端(设备)进行聚合。同样,它既适用于跨客户端平均的模型参数(变量),也适用于您的模型作为本地聚合结果导出的指标。

    在此级别执行聚合是 TFF 的职责。然而,作为模型创建者,您可以控制此过程(详见下文)。

    处理的总体结构如下:

    • 初始模型以及训练所需的任何参数,由服务器分发给将参与一轮训练或评估的客户端子集。

    • 在每个客户端上,您的模型代码在本地数据批次流上独立且并行地重复调用,以生成一组新的模型参数(训练时)和如上所述的一组新的本地指标(这是本地聚合)。

    • TFF 运行分布式聚合协议,以在系统范围内累积和汇总模型参数以及本地导出的指标。此逻辑使用 TFF 自己的联邦计算语言(而非 TensorFlow)以声明方式表达。有关聚合 API 的更多信息,请参阅 自定义算法 教程。

抽象接口

这种基本的构造函数 + 元数据接口由接口 tff.learning.models.VariableModel 表示,如下所示:

  • 构造函数、forward_passreport_local_unfinalized_metrics 方法应分别构建模型变量、前向传播和您希望报告的统计数据。这些方法构建的 TensorFlow 代码必须是可序列化的,如上所述。

  • input_spec 属性,以及返回您的可训练、不可训练和本地变量子集的 3 个属性,共同构成了元数据。TFF 使用这些信息来确定如何将模型的各部分连接到联邦优化算法,并定义内部类型签名以帮助验证所构建系统的正确性(从而防止您的模型在不匹配的数据上被实例化)。

此外,抽象接口 tff.learning.models.VariableModel 公开了一个 metric_finalizers 属性,它接收指标的未终结值(由 report_local_unfinalized_metrics() 返回)并返回终结后的指标值。metric_finalizersreport_local_unfinalized_metrics() 方法将一起用于在定义联邦训练过程或评估计算时构建跨客户端指标聚合器。例如,简单的 tff.learning.metrics.sum_then_finalize 聚合器将首先对来自客户端的未终结指标值求和,然后在服务器上调用终结函数。

您可以在我们 图像分类 教程的第二部分,以及我们在 model_examples.py 中用于测试的示例模型中,找到关于如何定义您自己的自定义 tff.learning.models.VariableModel 的示例。

Keras 转换器

几乎所有 TFF 所需的信息都可以通过调用 tf.keras 接口导出,因此如果您有 Keras 模型,可以使用 tff.learning.models.from_keras_model 构建 tff.learning.models.VariableModel

请注意,TFF 仍然要求您提供一个构造函数——一个不带参数的模型函数,如下所示:

def model_fn():
  keras_model = ...
  return tff.learning.models.from_keras_model(keras_model, sample_batch, loss=...)

除了模型本身,您还需要提供一个样本数据批次,TFF 使用它来确定模型输入的类型和形状。这确保了 TFF 能够为客户端设备上实际存在的数据正确实例化模型(因为我们假设在构建要序列化的 TensorFlow 时,这些数据通常是不可用的)。

Keras 包装器的使用在我们的 图像分类文本生成 教程中进行了说明。

联邦计算构建器

tff.learning 包为执行与学习相关的任务的 federated_language.Computation 提供了几个构建器;我们预计这类计算的集合在未来会不断增加。

架构假设

执行

运行联邦计算有两个不同的阶段。

  • 编译:TFF 首先将联邦学习算法编译成整个分布式计算的抽象序列化表示。此时会进行 TensorFlow 序列化,但为了支持更高效的执行,也会进行其他转换。我们将编译器发出的序列化表示称为联邦计算

  • 执行:TFF 提供了执行这些计算的方法。目前,执行仅通过本地模拟支持(例如,在笔记本中使用模拟的去中心化数据)。

由 TFF 联邦学习 API 生成的联邦计算(例如使用 联邦模型平均 (federated model averaging) 的训练算法,或联邦评估)包含多个元素,最显著的是:

  • 您的模型代码的序列化形式,以及由联邦学习框架构建的额外 TensorFlow 代码,用于驱动模型的训练/评估循环(例如构建优化器、应用模型更新、遍历 tf.data.Dataset、计算指标,以及在服务器上应用聚合后的更新等等)。

  • 客户端服务器之间通信的声明性规范(通常是跨客户端设备的各种形式的聚合,以及从服务器到所有客户端的广播),以及这种分布式通信如何与 TensorFlow 代码的客户端本地或服务器本地执行交织在一起。

以这种序列化形式表示的联邦计算使用独立于平台的内部语言表达,该语言有别于 Python;但在使用联邦学习 API 时,您无需关注此表示的细节。这些计算在您的 Python 代码中表示为 federated_language.Computation 类型的对象,您在大多数情况下可以将它们视为不透明的 Python 可调用对象 (callable)

在教程中,您将像调用常规 Python 函数一样调用这些联邦计算,以便在本地执行。然而,TFF 的设计方式是表达联邦计算时与执行环境的大多数方面无关,以便它们可以潜在地部署到例如运行 Android 的设备群或数据中心内的集群中。同样,这一设计的主要后果是对 序列化 有严格的假设。特别是,当您调用下述任何 build_... 方法时,计算将被完全序列化。

建模状态

TFF 是一个函数式编程环境,但联邦学习中许多感兴趣的流程是有状态的。例如,涉及多轮联邦模型平均的训练循环就是我们可以归类为有状态流程的一个例子。在此流程中,从一轮到下一轮演进的状态包括正在训练的一组模型参数,以及可能与优化器关联的额外状态(例如动量向量)。

由于 TFF 是函数式的,有状态流程在 TFF 中被建模为接受当前状态作为输入,然后提供更新后的状态作为输出的计算。为了完整定义一个有状态流程,还需要指定初始状态来自何处(否则我们无法引导该流程)。这被捕获在辅助类 tff.templates.IterativeProcess 的定义中,其 2 个属性 initializenext 分别对应于初始化和迭代。

可用构建器

目前,TFF 提供了各种构建器函数,用于生成联邦训练和评估的联邦计算。两个显著的例子包括:

数据集

架构假设

客户端选择

在典型的联邦学习场景中,我们拥有一个庞大的总体,可能包含数亿个客户端设备,其中只有一小部分在任何特定时刻处于活跃状态并可用于训练(例如,这可能仅限于连接到电源、不在计量网络上且处于闲置状态的客户端)。通常,可参与训练或评估的客户端集合不在开发者的控制范围内。此外,由于协调数百万个客户端是不切实际的,典型的训练或评估轮次将仅包含可用客户端的一小部分,这些客户端可能是随机抽样的。

其关键后果是,联邦计算在设计上以忽略确切参与者集合的方式表达;所有处理都表示为针对一组抽象的匿名客户端的聚合操作,该组在训练的不同轮次之间可能会有所不同。因此,计算与具体参与者(从而与他们输入到计算中的具体数据)的实际绑定,被建模在计算本身之外。

为了模拟您的联邦学习代码的真实部署,您通常会编写一个如下所示的训练循环:

trainer = tff.learning.algorithms.build_weighted_fed_avg(...)
state = trainer.initialize()
federated_training_data = ...

def sample(federate_data):
  return ...

while True:
  data_for_this_round = sample(federated_training_data)
  result = trainer.next(state, data_for_this_round)
  state = result.state

为了促进这一点,在模拟中使用 TFF 时,联邦数据被接受为 Python list,每个参与的客户端设备对应一个元素,以表示该设备的本地 tf.data.Dataset

抽象接口

为了规范处理模拟的联邦数据集,TFF 提供了一个抽象接口 tff.simulation.datasets.ClientData,它允许枚举客户端集合,并构建一个包含特定客户端数据的 tf.data.Dataset。这些 tf.data.Dataset 可以直接作为输入提供给即时执行模式下的生成式联邦计算。

需要注意的是,访问客户端身份的能力仅由用于模拟的数据集提供,在这些模拟中可能需要针对特定客户端子集进行训练(例如,模拟不同类型客户端的昼夜可用性)。已编译的计算和底层运行时涉及任何客户端身份的概念。一旦来自特定客户端子集的数据被选择为输入(例如,在调用 tff.templates.IterativeProcess.next 时),客户端身份就不再出现在其中。

可用数据集

我们将 tff.simulation.datasets 命名空间专门用于实现 tff.simulation.datasets.ClientData 接口的数据集,以便在模拟中使用,并以此为基础提供了支持 图像分类文本生成 教程的数据集。我们希望鼓励您将自己的数据集贡献给平台。