Page 402 - 《软件学报》2026年第3期
P. 402

刘建春 等: 基于块级多输出和知识自蒸馏的高效联邦学习框架                                                   1365


                 组不等式保证了网络中的通信资源约束, 第              3  组等式表示优化的变量      m  应该是一个整数.
                    由于上述问题中优化变量          m  是一个整数, 因此这是一个典型的整数规划问题. 而由于整数规划问题是                      NP  难
                 的问题  [35] , 所以直接解决该问题是一项很困难的任务. 为了解决这一问题, 本文使用能够反映客户端实时变化状
                 态的反馈变量, 并基于贪心策略来使服务器为每个客户端决策出最优的模型块数量. 算法的详细内容将在第                                   4  节
                 展开说明.

                  4   算法设计

                    本节中, 我们在服务器和客户端分别提出了一个算法来解决问题                      (9). 首先基于客户端的本地模型更新结构为
                 客户端定制了一个本地损失函数, 之后设计了一个反馈变量来表示客户端的状态信息, 最后计算出决策变量用于
                 确定客户端需要接收的模型块的数量.
                  4.1   准备工作
                    (1) 本地损失函数. 由于输入数据经过本文设计的块级多输出正则化结构之后会有多条输出, 且本节在多条输
                 出之间应用了知识自蒸馏技术来汲取额外的模型表征层信息, 所以本节将本地损失函数分为了两部分, 一部分是
                 本地损失项而另一部分是知识蒸馏项. 其中, 本节使用交叉熵损失函数来表示本地损失项, 该函数是深度学习领域
                 中一种最常见的损失函数         [36] . 具体的, 我们使用  CrossEntropy(a,b) 来表示两个相同维度变量      a  和  b  的交叉熵损
                 失, 则本地损失函数中的损失项可以表示为:

                                                   L ce = CrossEntropy(q m ,y)                       (11)
                 其中,  q m  表示数据经过完整的本地路径得到的输出, y 表示数据的真实标签. 此外, 我们使用                    Kullback-Leibler (KL)
                 散度  [37] 来表示本地损失函数中的知识蒸馏项.         KL(a,b) 表示两个相同维度向量       a  和  b  之间的  KL  散度, 则本地函数
                 的知识蒸馏项可以表示如下:

                                                        1  ∑ m−1
                                                  L KL =       KL(ˆq j , ˆq m )                      (12)
                                                       m−1    j=1
                                                                          δ
                 其中,   ˆ q 表示知识蒸馏过程中使用了超参数温度          δ 之后得到的输出, 而温度   在公式       (12) 中的作用是影响     Softmax
                 函数输出的平滑性. 将本地损失项和知识蒸馏项相结合, 我们给出了客户端                       k 如下的本地损失函数:

                                                                                                     (13)
                                                       L k = L ce +α· L KL
                 其中,  α (α ⩾ 0) 表示客户端  k 的本地损失函数中知识蒸馏项的权重. 客户端             k 根据公式   (13) 中的本地损失函数进
                 行本地模型更新.
                    (2) 反馈变量. 解决问题     (9) 的关键是在每一轮全局训练开始时, 为每个客户端确定要接收的最适合的全局模型
                 块的数量, 同时提高全局模型的性能. 根据问题             (9) 的约束条件, 客户端    k 需要接收的全局模型块的数量         m k  应与客户
                 端  k 的通信能力、计算能力以及本地数据分布情况有关. 对于数据分布, 我们采用本地模型和聚合后的全局模型之
                 间的差异来衡量客户端本地数据分布和全局数据分布的差异. 根据之前的研究表明, 本地数据分布和全局数据分布
                 之间的差异与本地模型和全局模型之间的差异成正比                  [38] . 我们使用  θ k  来表示客户端  k 的本地模型和聚合后的全局
                                            t  t 2
                 模型之间的差异, 则     θ k  可以表示为   || x − x || . 为了便于客户端之间相互衡量, 需要进行归一化, 因此     θ k  最终被表示如下:
                                               k

                                                             t  t 2
                                                           || x − x ||
                                                                k
                                                     θ k = ∑                                         (14)
                                                           K
                                                                  t
                                                               t
                                                             || x − x ′ || 2
                                                           ′
                                                          k =1    k
                    直观上来看, 当客户端       k 的本地数据分布与全局数据分布差异较大时, 客户端需要接收更多的全局模型块通
                 过知识自蒸馏技术吸收更多的模型表征层信息, 因此                 m k  应与  θ k  成正比. 对于通信能力和计算能力, 我们分别使用
                 通信时间和计算时间来量化它们. 我们使用              H k,b  和  H k,c  来分别表示客户端  k 在某一个全局轮次的通信时间和计算
                 时间. 为了便于比较, 我们使用了如下归一化的形式来表示客户端                    k 的通信时间和计算时间:

                                                       H k,b        H k,c
                                                 ′            ′
                                                 k,b  K       k,c   K
                                               H = ∑        , H = ∑                                  (15)
                                                         H k ′ ,b     H k ′ ,c
                                                      k ′ =1        k ′ =1
   397   398   399   400   401   402   403   404   405   406   407