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

1364                                                       软件学报  2026  年第  37  卷第  3  期



                 出, 同时也不存在产生      q 3  这条输出的计算路径.
                    本文的主要目的是在提高全局模型的泛化性能的同时减少通信开销, 框架所采用的块级多输出正则化可以实
                 现这一目的. 利用多条额外的输出可以使客户端本地模型吸收更多来自前面模型表征层的信息, 从而让本地训练
                 更加稳定并且向全局模型靠拢, 减轻了客户端之间本地非独立同分布数据带来的负面影响. 此外, 服务器采取分发
                 部分模型块的方式减少了通信开销. 服务器在决策发送客户端全局模型块的过程中充分考虑到客户端的本地数据
                 分布情况、计算能力以及通信能力. 直观上来看, 当客户端的本地数据分布偏离全局数据分布时, 客户端需要更多
                 的全局模型块来吸收更多的表征层信息. 而当客户端接收更多的全局模型块时, 则会消耗更多的通信资源. 因此,
                 如何在数据异构场景下决策出客户端在每一个全局轮次需要接收的全局模型块的数量, 从而保证全局模型性能是
                 一个关键的问题.
                  3.4   问题定义
                    在本节中, 我们给出在联邦学习框架下为每个客户端发送最适合模型块问题的形式化定义. 假设在网络中一
                 共有  N  个客户端, 其中有    K  个客户端参与训练. 我们主要考虑网络中的计算资源和通信资源, 使用                    C  和  B  来表示
                 计算资源和通信资源的总预算, 即训练过程中消耗的计算资源和通信资源不能超过                            C  和  B. 此外, 我们使用  c 和  b
                 来分别表示使用一个批次的数据对模型的梯度更新所需的计算资源和传递完整模型所需要的通信资源. 为了方便
                 进行表示, 我们假设模型被均分为          M  个块且每个模型块的参数数量都是相同的, 则每传递一个模型块产生                      b/M  的
                 通信开销. 由于全连接层参数数量相对于整个模型而言很少, 因此全连接层带来的计算开销和通信开销可以忽略
                                                             c b , 以及客户端        n k  个批次的数据, 并且在某一个
                 不计. 此外, 我们假设输入通过瓶颈层需要的计算开销为                             k 本地有
                 全局轮次            m  个模型块, 则客户端     k 在第  t 个全局轮次产生的计算开销为:
                                  t
                        t 中接收到    k
                                                      t
                                                          t
                                                     c = (m ·c b +c)·n k ·e                           (4)
                                                          k
                                                      k
                 其中, e 表示连续两个全局轮次之间的本地迭代数量. 公式                (4) 中  m ·c b  表示客户端  k 每接收一个模型块则会多产
                                                                    t
                                                                    k
                                        τ 轮会传送整个的全局模型, 此时客户端           k 的计算开销为:
                 生  c b  的计算开销. 服务器每隔

                                                    t
                                                   c = [( M −1)·c b +c]·n k ·e                        (5)
                                                    k
                    结合公式    (4) 和公式  (5), 可以得到客户端   k 的每个全局轮次的平均计算开销为:

                                                    (              )
                                                     M −1+τ·m t
                                                 t
                                                 c =          k  ·c b +c ·n k ·e                      (6)
                                                 k
                                                        τ+1
                       t
                 其中,  m ∈ {0,..., M −1}. 注意, 当服务器传送整个模型时,    m = 0. 对于通信开销, 假设在某一轮全局通信轮次             t 中,
                                                              t
                                                              k
                       k
                 客户端   k 下载和上传   m  个模型块, 则客户端     k 在第  t 个全局轮次产生的通信开销为:
                                  t
                                  k

                                                          (    t  )
                                                              m
                                                        t
                                                       b = 1+  k  ·b                                  (7)
                                                        k
                                                              M
                 其中, 1  表示客户端上传整个全局模型. 同理, 服务器每隔              τ 轮传送完整的全局模型, 此时客户端           k 的通信开销为
                 2·b. 综上所述, 可以得到客户端       k 的每个全局轮次的平均通信开销为:

                                                    (     (    t  )  )
                                                       τ     m     2
                                                  t
                                                 b =     · 1+  k  +   ·b                              (8)
                                                  k
                                                     τ+1     M    τ+1
                    本文的优化目标是最小化全局损失函数, 同时为每一个客户端找到最优的模型块数量, 从而定制最适合的本
                 地模型. 因此, 该问题可以被定义成如下形式:

                                                              ( )
                                                         min f x T                                    (9)
                                                        T∈{1,2,3,...}

                                                        T
                                                    ∑ ∑    K
                                                              t
                                                            c ⩽ C
                                                   
                                                             k
                                                       t=1  k=1
                                                   
                                                   
                                                   
                                                       T   K
                                                    ∑ ∑
                                                 s.t.        t                                      (10)
                                                            b ⩽ B
                                                             k
                                                       t=1  k=1
                                                   
                                                   
                                                   
                                                   
                                                      t
                                                     m ∈ {0,..., M −1}, ∀k,t
                                                       k
                         t
                 其中, 当  m = 0 时表示服务器传送给客户端         k 整个的全局模型. 第     1  组不等式保证了网络中的计算资源约束, 第            2
                         k
   396   397   398   399   400   401   402   403   404   405   406