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

