Skip to content

分块矩阵的乘法

矩阵乘法是线性代数中最重要的运算之一。在机器学习中,矩阵乘法也是经常用到的运算,最常见于 MLP 线性层。

而在实际的模型训练和推理系统中,模型参数和中间激活的张量可能非常大,而 GPU 显存空间有限。因此,我们需要将张量切分为多个块,以在 GPU 上实现并行计算。而这和分块矩阵的乘法有着紧密的联系。


分块矩阵的定义

一个 m×n 维的矩阵 W 定义如下

Wm×n=[w11w12w1nw21w22w2nwm1wm2wmn]

现在将这个矩阵在行方向上划分为 X 个块,在列方向上划分为 Y 个块。这样,每个块的维度为 x×y,其中

{x=mXy=nY

于是,分块后的矩阵可记为

WX×Y=[W11W12W1YW21W22W2YWX1WX2WXY]

其中,每个块 Wij (1i,jX,Y) 可以表示为

Wij=[w(i1)x+1,(j1)y+1w(i1)x+1,(j1)y+2w(i1)x+1,jyw(i1)x+2,(j1)y+1w(i1)x+2,(j1)y+2w(i1)x+2,jywix,(j1)y+1wix,(j1)y+2wix,jy]

向量与分块矩阵相乘

设一个 m 维的向量

λ=[k1k2km]

要使向量 λ 右乘分块矩阵 W,需要相应地将向量 λ 分成 X 个块,即

λ=[λ1λ2λX]

其中,每个块 λi (1iX) 可以表示为

λi=[k(i1)x+1k(i1)x+2kix]

因此,向量 λ 右乘分块矩阵 W 可以表示为

λW=[i=1XλiWi1i=1XλiWi2i=1XλiWiy]

矩阵与分块矩阵相乘

设一个 d×m 维的矩阵

Ad×m=[a11a12a1ma21a22a2mad1ad2adm]

仿照向量与分块矩阵相乘的方式,如果对矩阵 A 作 1D 切分,即

A=[A1A2AX]

其中,每个分块矩阵的维度是 d×x。那么,矩阵 A 右乘分块矩阵 W 可以表示为

AW=[i=1XAiWi1i=1XAiWi2i=1XAiWiy]

如果对矩阵 A 作 2D 切分,在行方向上切分成 T 个块,在列方向上切分成 X 个块。这样,每个块的维度就变成 t×x,其中

t=dT

那么矩阵 A 可表示为

AT×X=[A11A12A1XA21A22A2XAT1AT2ATX]

这样,矩阵 A 右乘分块矩阵 W 可以表示为

AW=[i=1XA1iWi1i=1XA1iWi2i=1XA1iWiyi=1XA2iWi1i=1XA2iWi2i=1XA2iWiyi=1XATiWi1i=1XATiWi2i=1XATiWiy]

分块矩阵乘法的一般规律

与普通的矩阵乘法相比,分块矩阵乘法多了一个步骤:分块累加。用规范化的数学语言描述,即

Am×p=[A11A12A1PA21A22A2PAM1AM2AMP]Bp×n=[B11B12B1NB21B22B2NBP1BP2BPN]AB=[i=1PA1iBi1i=1PA1iBi2i=1PA1iBiNi=1PA2iBi1i=1PA2iBi2i=1PA2iBiNi=1PAMiBi1i=1PAMiBi2i=1PAMiBiN]

其中,矩阵 A 被 2D 切分为 M×P 份,矩阵 B 被 2D 切分为 P×N 份。

而至于向量与分块矩阵的乘法,以及 1D 和 2D 切分的情况,都属于上述一般规律的特例。是否需要将分块累加,取决于是否对 p 维进行切分。

图解分块矩阵乘法

设想有 2 个矩阵都被切分成 2×2 的块,每个块在 4 张不同的 GPU 上进行矩阵乘法运算。则每个 GPU 上的计算分别如下面 4 张图所示:

分块矩阵乘法的意义

矩阵乘法是线性代数中最重要的运算之一。在机器学习中,矩阵乘法也是经常用到的运算,最常见于 MLP 线性层。

而在实际的模型训练和推理系统中,模型参数和中间激活的张量可能非常大,而 GPU 显存空间有限。因此,我们需要将张量切分为多个块,以在 GPU 上实现并行计算。

而分块矩阵乘法的意义在于,分块意味着一个完整的张量可以进行切分,在 GPU 上进行并行的矩阵乘法计算从理论上说是可行的。

不过,如果我们将矩阵在 p 维上切分,最终的计算结果需要沿着 p 维进行聚合。这样的操作会带来额外的通信开销。

此外,最终的结果从整体上看是多个小块矩阵的拼接,这意味着每个 GPU 可以存储相应的不同的结果。至于是否存在通信开销,取决于后面的算子是否需要完整的张量。

最近更新