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

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


                 的参数也会进行更新, 即客户端对本地模型的所有参数都会进行更新. 此外, 当客户端上传本地模型时, 并不上传
                 瓶颈层的参数.


                               模型块1            模型块2            模型块3            模型块4




                                       瓶颈层             瓶颈层             瓶颈层            全连接层


                                    q 1             q 2             q 3            q m
                                                                                           ζ L
                                                            ζ KL
                                                                                    y
                                          图 2 基于知识自蒸馏和块级多输出的模型训练

                    (4) 模型聚合. 在服务器收集到所有的         K  个客户端上传的本地模型之后, 服务器按照下面的模型聚合规则来进
                 行全局模型的更新:

                                                           1  ∑ K
                                                      x t+1  =  x t+1                                 (3)
                                                                 k
                                                           K   k=1
                                                                                                   t +1 个
                    注意, 客户端并不会上传瓶颈层的参数, 所以瓶颈层不会参与模型聚合. 最后, FedAlt 的模型训练在第
                 全局轮次继续进行模型广播并一直持续到全局模型收敛或者网络资源耗尽为止.
                  3.3   块级多输出正则化
                    客户端的本地更新过程会按照图            2  所示的结构进行更新. FedAvg     的本地模型训练只会得到         q m  这一条本地路
                 径的输出. 而在    FedAlt 中, 客户端会根据服务器决策出的全局模型块的数量来决定产生多少条输出, 客户端每接
                 收一个全局模型块就会多一条输出, 即对于客户端                k 而言共产生    m k +1 条输出. 产生输出的路径如图       2  所示, 我们
                 在每个额外的输出之前加上一个瓶颈层来统一输出的维度. 有了这多条额外输出, 客户端使用知识自蒸馏技术起
                 到了正则化的效果, 因此客户端可以吸取更多来自前面模型块的表征层信息, 从而使得本地模型向全局模型方向
                 靠拢. 值得注意的是, 当服务器发送给客户端的是整个的全局模型时, 客户端的本地更新过程是按照接收到                                 M −1
                 块全局模型块进行的, 即客户端直接按照图              2  所示结构进行本地更新.
                    为了更好地说明基于知识自蒸馏和块级多输出的正则化技术, 我们结合图                         2  给出了一个具体的例子. 假设我
                 们将一个模型分成了        4  个连续的模型块, 即    M = 4. 在图  2  中, 我们将模型块  4  和全连接层分开画出, 而在后面的
                 描述中, 服务器分发最后一块模型块时是指服务器将会把最后一块模型块以及全连接层一起发送. 此外, 瓶颈层并
                 不参与服务器与客户端之间的模型传输. 例如, 当服务器发送给某个客户端                       k 模型块  4  时, 它会将模型块    4  与其后
                 面的全连接层同时发送给客户端           k. 而当服务器发送给某个客户端          k 模型块  3  时, 它并不会将模型块     3  后面的瓶颈
                 层发送给客户端      k. 在某一个全局轮次      t 中, 服务器将前   2  块全局模型块    x t,1   和  x  发送给某个客户端  k, 即  m = 2.
                                                                                                    t
                                                                              t,2
                                                                                                    k
                 结合图   2, 模型块 1 与模型块 2 分别对应      x t,1   和  x , 而模型块  3  和模型块  4  分别对应客户端  k 在上一个全局轮次
                                                      t,2
                 t −1 的本地模型的后两块, 即      x t−1,3   和  x t−1,4 . 之后客户端  k 基于这个组合起来的模型进行本地更新. 如图    2  所示, 客
                                         k     k
                 户端  k 的输入经过本地正向传播会得到           3  条输出, 即   {q 1 ,q 2 ,q m } q m  是经过了所有的模型块得到的本地输出, 而   q i
                                                                  .
                 是经过   i 个模型块和一个最后的瓶颈层得到的输出. 例如,             q 2  是由输入经过模型块     1、模型块    2  和一个瓶颈层得到
                 的输出. 得到   3  条输出之后, 客户端     k 会将这  3  条输出放入到损失函数中. 其中, 本地输出和真实标签之间会采用
                 标准的交叉熵损失, 而块级多输出正则化中额外增加的输出与本地输出之间将使用                            KL  散度来进行额外知识的提
                 取. 损失函数将在第      4  节中具体说明. 注意, 为了清楚地表示所有情况, 图            2  展示的是服务器分发给客户端         k 整个
                 全局模型的场景. 而在图       1  所示的例子中, 由于此时客户端         1  只接收了两个全局模型块, 因此并不存在            q 3  这条输
   395   396   397   398   399   400   401   402   403   404   405