Page 335 - 《软件学报》2026年第7期
P. 335

3020                                                       软件学报  2026  年第  37  卷第  7  期


                    从线程组织形式的角度来看, 上述           5  个操作主要呈现     3  种并行关系: 第  1  种为  n 并行的域扩张和域收缩; 第      2
                     N/2 并行的正向和逆向      NTT  变换; 第  3  种为受拓展环  NTT  分解影响的多项式点乘. 由于并行关系的不同, 使
                 种为
                 用多个核函数是自然的想法, 但是这显然会带来和全局内存过于频繁交互的效率降低, 因此本文使用单一核函数
                 实现上述    5  个操作. 鉴于上述   5  个操作中, 正向和逆向     NTT  变换占比最高且耗时最多, 线程组织形式的设计在优
                 先减少上述两个操作分支执行的基础上, 力求尽可能多地满足更多操作的并行特性.
                    从伪梅森数不完整       NTT  的并行特性出发, 由于      CTRU-Prime 的  3  组参数的拓展环分解中, 都含      3k (k = 2,3) 个
                 基  2-NTT, 因此本文中, 每个线程使用      8  个寄存器以完成层数为       3  的层融合, 这既是适应基      2-NTT  数量的考量, 也
                 是对  SM  理论占用率的核函数线程数量和寄存器使用量的平衡. 具体来说, 每个线程块包含                        N/8 个线程, 每个线程
                 拥有  8  个寄存器以存储不同位置的多项式系数. 为成功实现层融合, 需要正确设计在共享内存中的存取位置. 具体
                                                                                            i
                                                                               i       N/(8×2 ) 个连续线程的
                 来说, 假如当前位于第       i 层, 则根据其使用旋转因子的异同, 线程被均匀地分到               2  组中, 即
                                                                i
                 存取位置被限制在对应组内. 为实现层融合, 各线程以                N/(8×2 ) 的间隔从组内存取数据. 上述方式将         3  次存取共享
                 内存的操作减少为       1  次, 从而进一步提高了存储效率. 为在         CTRU-Prime-653  中应用层融合技术, 本文将伪梅森数
                 不完整   NTT  的  FFT-trick  计算顺序调整为先计算  6  层基-2 FFT trick, 再计算  1  层基-3 FFT trick.
                    在上述线程组织形式下, 每个线程块完成一次多项式乘法, 使用大小为                      3  个拓展环上多项式的共享内存空间.
                 对于除正向和逆向       NTT  以外的其他操作,    N/8 的线程束数量仅引起一次单个线程束的分支执行, 并未带来较大的
                 性能惩罚. 除点乘外, 域扩张和域收缩过程的线程下标对齐是简单且自然的. 多项式点乘的实现方式为在                                CRT  同
                 构图的叶子节点上完成低维的教科书式多项式乘法, 其下标对齐需考虑旋转因子表的访问. 例如, 在                              CTRU-Prime-
                                                             7
                 653  中, 经过正向  NTT  后, 需在  192  个形如  Z 16777153 [x]/(x −ζ 64k+br 6 (j)+1 ) (k = 0,1,2;0 ⩽ j ⩽ 63)  的域上完成多项式点
                 乘. CRT  同构树形图的第    i 个叶子节点对应旋转因子表中的第          64+2×i 个位置, 其对应的域为     Z 16777153 [x]/(x −ωζ  64+2i ),
                                                                                                7
                 其中  ω 为  3  次本原单位根. 本文使用    Karatsuba 技巧进一步加速点乘过程.
                    本文引入了循环展开技术, 以降低循环迭代次数、增强指令级并行性、提升数据缓存效率, 从而提高整体效
                 率. 同时, 本文实施延迟约减策略, 并非在每一层级的计算后立即执行约减, 而是基于具体计算结果, 仅约减那些可
                 能超出数据表示范围的情况, 从而有效减少不必要的约减次数以加快计算速度. 此外, 基于伪梅森数不完整                                  NTT
                 旋转因子中存在大量−1        和  1  的情况, 为  s 进一步减少约减次数, 对于旋转因子为          1  的蝴蝶操作, 省略乘法直接进
                 行加减, 对于旋转因子为−1       的蝴蝶操作, 将     a[i]+a[j]×twiddle (twiddle 代表旋转因子) 变为  a[i]–a[j], 将  a[i]–a[j]×
                 twiddle 变为  a[i]+a[j].
                  3.2   基于  Tensor Core 的教科书式多项式乘法
                                         Z q [x]/(x − x−1) 上多项式乘法矩阵表示的数学推导, 然后给出了基于            Tensor Core
                                               n
                    本节首先给出了素阶数域
                 的素阶数域上多项式乘法的          GPU  实现.
                  3.2.1    素阶数域上多项式乘法的矩阵转化
                    假设   s = f ·g (注意这里并未进行多项式模运算). 教科书式的多项式乘法, 在得到                s 的基础上, 利用素阶数域上
                           n               h = s mod (x − x−1). 图                   x ≡ x+1 对多项式模运算
                                                                                     n
                                                    n
                 的特殊性质    x ≡ x+1 进一步计算                       2  以  n = 3 为例, 展示了性质
                 的影响. 我们使用     p i  代表多项式   p 的第  i 个多项式系数. 图  2  中   s 3  所在虚线框发射出两条箭头分别到达     s 0  和  , 这
                                                                                                    s 1
                                x ≡ x+1 s 3  在多项式模运算中对第     0  位和第  1                           h 1  上. 进一
                                 3
                                       ,
                 两条箭头代表依据                                            位存在影响, 即    s 3  会被累加到  h 0  和
                 步推广得到如下规律: 当       i ⩾ n 时,   s i  会被累加到   h i mod n  和  h (i mod n)+1  上. 应用该规律, 得到矩阵  M 1 , 对其每一列进行求
                 和, 即可得到最终结果多项式         h 的各个系数. 进一步地, 将矩阵       M 1  拆分为矩阵和向量的乘积      (M 2 + M 3 )·v, 其中  v 为
                 多项式  g 各个系数构成的列向量,        M 2  和  M 3  由多项式   f  导出, 分别对应了   s i  对   h i mod n  和  h (i mod n)+1  的影响.
                                               ∑  n−1     ∑  n−1                         ∑ 2n−2
                                                                                                i
                                                      i
                                                                 i
                                                      ,
                    设   f和g 是  n−1 次多项式, 其中   f =   f i x g =  g i x , 二者的直接乘积可以写为    s =     s i x , 其中:
                                                  i=0        i=0                           i=0
                                                  ∑  i
                                                 
                                                       f i−k g k ,  0 ⩽ i ⩽ n−1
                                                 
                                                 
                                                     k=0
                                                                                                      (1)
                                                 
                                                     n−1
                                              s i =  ∑
                                                 
                                                 
                                                         f i−k g k ,  n ⩽ i ⩽ 2n−2
                                                 
                                                      k=i−n+1
   330   331   332   333   334   335   336   337   338   339   340