【infra之路】Tiled Matrix Multiplication详解

先回顾:Naive 矩阵乘法长什么样

假设我们要算 C = A × B,矩阵都是 N×N。C 的每个元素 C[row][col] 是 A 的第 row 行和 B 的第 col 列做点积:

C[row][col] = A[row][0]*B[0][col] + A[row][1]*B[1][col] + ... + A[row][N-1]*B[N-1][col]
                                   └─── 沿 K 维度求和 ──────────────────────────────────┘

Naive 版本每个线程算 C 的一个元素,从 Global Memory 读 A 的一整行和 B 的一整列。问题是 A 和 B 的数据被大量重复读取——计算 C[0][0] 读了 A 的第 0 行,计算 C[0][1] 又要再读一遍 A 的第 0 行。

Tiled 的核心思想:分块加载到 Shared Memory

把大矩阵切成小块(tile),每次只处理一小块,这样数据可以被 Block 内的线程复用

画个图最直观。假设 N=8,TILE_SIZE=4:

矩阵 A (8×8)          矩阵 B (8×8)          矩阵 C (8×8)
┌────┬────┐           ┌────┬────┐           ┌────┬────┐
│A00 │A01 │           │B00 │B01 │           │C00 │C01 │
│4×4 │4×4 │           │4×4 │4×4 │           │4×4 │4×4 │
├────┼────┤           ├────┼────┤           ├────┼────┤
│A10 │A11 │           │B10 │B11 │           │C10 │C11 │
│4×4 │4×4 │           │4×4 │4×4 │           │4×4 │4×4 │
└────┴────┘           └────┴────┘           └────┴────┘

要算 C00(左上角的 4×4 子矩阵),需要:

C00 = A00 × B00 + A01 × B10
      ─────────   ─────────
       第 1 个 tile   第 2 个 tile

关键洞察:C00 这块 4×4 = 16 个元素,每个元素都需要读 A 的第 0-3 行和 B 的第 0-3 列。如果不用 Shared Memory,16 个线程各读各的,A 的同一行被读 4 次。用了 Shared Memory,整个 Block 协作把 A00 这块 4×4 一次性搬到 Shared Memory,然后 16 个线程都从 Shared Memory 读——Global Memory 访问量减少到 1/4

逐行解读代码

__global__ void matMulShared(float *A, float *B, float *C, int N) {
    // ① 声明 Shared Memory:每个 Block 分配两块 TILE_SIZE × TILE_SIZE 的共享空间
    __shared__ float sA[TILE_SIZE][TILE_SIZE];
    __shared__ float sB[TILE_SIZE][TILE_SIZE];

__shared__ 关键字告诉编译器这块内存在 SM 的 Shared Memory 上,Block 内所有线程共享。

    // ② 计算当前线程负责 C 的哪个元素
    int row = threadIdx.y + blockIdx.y * TILE_SIZE;
    int col = threadIdx.x + blockIdx.x * TILE_SIZE;

假设 TILE_SIZE=4,blockIdx = (1, 2),threadIdx = (3, 1):

  • row = 1 + 2 × 4 = 9 → 第 9 行
  • col = 3 + 1 × 4 = 7 → 第 7 列
  • 这个线程负责计算 C[9][7]

这里 threadIdx.x 对应列、threadIdx.y 对应行——正好满足上一课讲的合并访问原则threadIdx.x 相邻 → col 相邻 → 内存连续)。

    float sum = 0.0f;

    // ③ 沿 K 维度(公共维度)分块遍历
    for (int t = 0; t < N / TILE_SIZE; t++) {

t 是 tile 的编号。N=8, TILE_SIZE=4 时,t 从 0 到 1,循环两次——对应前面图里的 A00×B00A01×B10

        // ④ 协作加载:每个线程从 Global Memory 读一个元素,存入 Shared Memory
        sA[threadIdx.y][threadIdx.x] = A[row * N + (t * TILE_SIZE + threadIdx.x)];
        sB[threadIdx.y][threadIdx.x] = B[(t * TILE_SIZE + threadIdx.y) * N + col];

这是最关键的一步。Block 内有 TILE_SIZE × TILE_SIZE 个线程(比如 4×4 = 16 个),每个线程负责搬一个元素:

线程 (ty=0, tx=0) 搬 A[row][t*4 + 0] → sA[0][0]
线程 (ty=0, tx=1) 搬 A[row][t*4 + 1] → sA[0][1]
...
线程 (ty=3, tx=3) 搬 A[row][t*4 + 3] → sA[3][3]

16 个线程各搬 1 个 → 一次协作就把 4×4 的 tile 从 Global Memory 搬到了 Shared Memory

注意 sA 和 sB 的索引方式:

  • sA[ty][tx]:ty 是行,tx 是列。每个线程搬 A 中当前行、第 t 块的一个元素
  • sB[ty][tx]:ty 是行,tx 是列。每个线程搬 B 中第 t 块、当前列的一个元素
        __syncthreads();  // ⑤ 屏障同步:等所有线程都搬完了再继续

为什么必须有这个? 如果没有,可能线程 0 已经搬完了开始计算,但线程 15 还没搬完——线程 0 就会读到 sA 中的垃圾数据。__syncthreads() 确保 Block 内所有线程都走到这里之后才一起继续。

        // ⑥ 在 Shared Memory 内做计算
        for (int k = 0; k < TILE_SIZE; k++) {
            sum += sA[threadIdx.y][k] * sB[k][threadIdx.x];
        }

现在数据在 Shared Memory 里了,做点积。对于计算 C[row][col] 的线程:

sum += sA[ty][0] * sB[0][tx]    ← 从 Shared Memory 读,~20 cycles
sum += sA[ty][1] * sB[1][tx]
sum += sA[ty][2] * sB[2][tx]
sum += sA[ty][3] * sB[3][tx]

对比 naive 版直接从 Global Memory 读(~300+ cycles/次),每次访问快了 15 倍以上

        __syncthreads();  // ⑦ 确保所有线程算完了再加载下一个 tile(覆盖 sA/sB)
    }

    // ⑧ 写回结果
    C[row * N + col] = sum;
}

第二个 __syncthreads() 也不能省:下一轮循环会用新的数据覆盖 sA 和 sB,如果还有线程没算完就被覆盖了,结果就错了。

一图总结整个流程

                K 维度分块遍历
                ┌──────────┐
                │  t = 0   │
                └────┬─────┘
                     ▼
    ┌────────────────────────────────────┐
    │ 16 个线程协作加载 A_tile, B_tile    │  ← Global → Shared
    │         __syncthreads()             │
    │ 每个线程在 Shared Memory 内做部分求和 │  ← 快速读取
    │         __syncthreads()             │
    └────────────────┬───────────────────┘
                     ▼
                ┌──────────┐
                │  t = 1   │  ← 加载下一块 tile,覆盖 sA/sB
                └────┬─────┘
                     ▼
                  ... 重复 ...
                     ▼
               所有 tile 算完
                     ▼
              写 sum 到 Global Memory

核心收益:Naive 版本中 A 的每一行被读 N 次(C 的每一列都要用),tiled 版本中只从 Global Memory 读 N/TILE_SIZE 次(每次搬一个 tile,Block 内复用 TILE_SIZE 次)。当 TILE_SIZE=32 时,Global Memory 带宽需求降低到 1/32

评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值