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

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


                 述两个操作中, 连续内存空间和矩阵           (本质上是各个线程的寄存器) 间的映射关系对编程者是透明的. Zhou 等人                   [21]
                 通过逆向的方式获取该种映射关系. wmma::mma_sync 操作完成线程同步的                  MMA  操作. 对于不同的精度, Tensor
                 Core 的性能是不同的. 概括来说, 精度越低, 运算速度越快. 根据文献               [20], 当使用  8  位整型作为计算精度时, MMA
                 操作表现出最短的延迟.
                    基于此, 在   RTX 3060  平台上本文选择    SM86  架构支持的    m=16、n=16、k=16, 左矩阵和右矩阵精度为        8  位整
                 型, 累加器精度为     32  位整型的  MMA  操作来完成     CTRU-Prime 中的多项式乘法. 在     CTRU-Prime 中, 需要计算的
                 多项式乘法包含      2  种, 第  1  种是   R q  中的多项式乘法  h = g· f inv mod q 和  σ = hr mod q; 第  2  种是   R q 2  中的多项式乘
                    ′         ′      ±       g、r、 f  由参数为   2  或  3                       g、r、 f  只需一个
                                                   ′
                 法   m = c f = c(2f +1) mod q 2 . 其中             的中心二项分布采样得到, 因此
                 8  位整型表示.   f inv 、h、c 是  R q  上的多项式. 如表  1  所述, CTRU-Prime 的  3  组  q 都为  13  比特, 因此本文使用  2  个
                 8  位整型来表示   R q  上的多项式系数.
                    图  3  展示了使用   Tensor Core 完成多项式乘法    h = g· f ∈ R q , 其中  g 由两个  8  位整型表示,   f  由一个  8  位整型
                                f  依据公式                        B x ×B y  维矩阵  . 需要注意的是, 由于矩阵乘法操作会
                                                                          F
                 表示. 首先, 将  g 和          (4) 转化为   A x ×A y  维矩阵  G  和
                 对矩阵进行分块, 因此需要满足条件           m|A x 、n|A y 、k|B y 、A y = B x . 本文使用在空缺位置填充  0  的方式来满足上述条
                 件. 例如, 对于  CTRU-Prime-653  首先依据公式   (4) 转化为  653×653  维矩阵和  653×1  维矩阵的乘法, 在填充   0  后, 变
                 为  656×656  维矩阵和  653×16  维矩阵的乘法. 进一步地, 将矩阵      G  分为矩阵  G h  和  G l , 其中  G h  对应  G  的高  7  比特,
                       G 的低  7  比特. 在完成多项式到矩阵的转化后, 基于分块矩阵乘法的思想, 使用                  m×n×k 维的  wmma 操作完
                 G l  对应
                 成矩阵乘法    H h = G h · F  和  H l = G l · F. 具体来说, 每一个线程束负责一个  m×k 维子矩阵的计算, 为此其需以    m×n  维
                 子矩阵为单元遍历矩阵         A  的一行分块矩阵, 以     n×k  维子矩阵遍历矩阵      B  的一列分块矩阵, 在遍历过程中使用
                                                                                             H
                 wmma 操作完成一次矩阵乘法并累加计算结果到矩阵                 C  中. 最后, 合并矩阵   H h  和矩阵   H l  得到矩阵  , 并将矩阵  H
                 转化为多项式     h. 若  g 由一个  8  位整型表示,  f  由两个  8  位整型表示, 大致过程和上述类似, 只是        f  而非  g 需要拆分
                 为高、低比特两个矩阵, 在此不再赘述.

                                     f:
                                                  F                    SHL7
                                                                   H h
                                                                                        h:
                           g:
                                                                                H
                                                   G h
                                         G
                                                                   H l
                                                          Tensor Core
                                                   G l


                                                  A y          B y          B y

                                                            B x
                                           A x                           A x

                                                  A            B            C

                                         图 3 基于 Tensor Core 的素阶数域多项式乘法计算

                    在实现过程中, 核函数       initMA  完成多项式   g 到  G h  和  G l  的转变; 核函数  initMB  完成多项式   f  到  F  的转变; 核
                 函数  Twmma  同时完成  G h · F  和  G l · F; 核函数  Merge 完成  H h  和   H l  到  h 的转变. 内存模式上, 批量处理时, 以相邻的
                 高位矩阵和低位矩阵为存储单元. 如算法             7  所示, 假设线程位于第     i 个线程束组, 其中每个线程束组处理一个多项
                 式乘法, 含   WRAP_NUMS  个线程束. 在    Twmma  中, 如算法   7  的第  14、15  行所示, 该线程基于所在线程束编号
   332   333   334   335   336   337   338   339   340   341   342