CUDA 算子笔记 2

§ 参考资料 共 6 条

Softmax Attention

输入 $\boldsymbol Q_{M\times d}$、$\boldsymbol K_{N\times d}$ 和 $\boldsymbol V_{N\times d}$,计算 Attention:

$$ \mathrm{Attention}(\boldsymbol Q,\boldsymbol K,\boldsymbol V)=\mathrm{softmax}\left(\frac{\boldsymbol Q\boldsymbol K^T}{\sqrt d}\right)\boldsymbol V $$

其中,Softmax 为逐行计算

例子:

Input:
Q (2x4): [[1.0, 0.0, 0.0, 0.0],
          [0.0, 1.0, 0.0, 0.0]]
K (3x4): [[1.0, 0.0, 0.0, 0.0],
          [0.0, 1.0, 0.0, 0.0],]
          [0.0, 0.0, 1.0, 0.0]]
V (3x4): [[1.0, 2.0, 3.0, 4.0],
          [5.0, 6.0, 7.0, 8.0],]
          [9.0, 10.0, 11.0, 12.0]]

Output:
O (2x4): [[4.29, 5.29, 6.29, 7.29],
          [5.0, 6.0, 7.0, 8.0]]

最简单的做法就是三个 Kernel:

  • 转置 + 矩阵乘法,计算 $\boldsymbol Q\boldsymbol K^T/\sqrt d$
  • Softmax,每一行作为一个 Block,用不到 Global Memory
  • 矩阵乘法

代码:

#include <cuda_runtime.h>

__device__ __forceinline__ constexpr int compute_offset(int row, int col,
                                                        int cols) {
  return row * cols + col;
}

// A: M x K
// B: K x N
// => output: AB, M x N
// A: M x K
// B: K x N
// => output: AB^T, M x N
template <const int block_size, const bool transpose, const bool norm>
__global__ void multiple_kernel(const float* A, const float* B, float* output,
                                int M, int N, int K) {
  __shared__ float tile_A[block_size][block_size];
  __shared__ float tile_B[block_size][block_size];

  const auto col = threadIdx.x + blockIdx.x * blockDim.x;
  const auto row = threadIdx.y + blockIdx.y * blockDim.y;

  const auto local_row = threadIdx.y;
  const auto local_col = threadIdx.x;

  float sum = 0.0f;

#pragma unroll
  for (int i = 0; i < K; i += block_size) {
    if (row < M && local_col + i < K) {
      tile_A[local_row][local_col] = A[compute_offset(row, local_col + i, K)];
    } else {
      tile_A[local_row][local_col] = 0.0f;
    }

    if constexpr (transpose) {
      if (col < N && local_row + i < K) {
        tile_B[local_row][local_col] = B[compute_offset(col, local_row + i, K)];
      } else {
        tile_B[local_row][local_col] = 0.0f;
      }
    } else {
      if (local_row + i < K && col < N) {
        tile_B[local_row][local_col] = B[compute_offset(local_row + i, col, N)];
      } else {
        tile_B[local_row][local_col] = 0.0f;
      }
    }

    __syncthreads();

#pragma unroll
    for (int j = 0; j < block_size; j++) {
      sum = fmaf(tile_A[local_row][j], tile_B[j][local_col], sum);
    }

    __syncthreads();
  }

  if (row < M && col < N) {
    if constexpr (norm) {
      output[compute_offset(row, col, N)] = sum * rsqrtf(static_cast<float>(K));
    } else {
      output[compute_offset(row, col, N)] = sum;
    }
  }
}

struct SoftmaxState {
  float max;
  float sum;
};

__device__ __forceinline__ SoftmaxState init_state() {
  return {-INFINITY, 0.0f};
}

__device__ __forceinline__ SoftmaxState mono_state(float value) {
  return {value, 1.0f};
}

__device__ __forceinline__ SoftmaxState combine(SoftmaxState a,
                                                SoftmaxState b) {
  if (a.max == -INFINITY) {
    return b;
  }
  if (b.max == -INFINITY) {
    return a;
  }

  SoftmaxState out;
  out.max = fmaxf(a.max, b.max);
  out.sum = a.sum * __expf(a.max - out.max) + b.sum * __expf(b.max - out.max);
  return out;
}

__device__ __forceinline__ SoftmaxState warp_reduce(SoftmaxState state) {
#pragma unroll
  for (int offset = 16; offset > 0; offset >>= 1) {
    SoftmaxState other;
    other.max = __shfl_down_sync(0xffffffff, state.max, offset);
    other.sum = __shfl_down_sync(0xffffffff, state.sum, offset);
    state = combine(state, other);
  }
  return state;
}

template <const int block_size>
__device__ __forceinline__ SoftmaxState block_reduce(SoftmaxState state) {
  __shared__ SoftmaxState warp_results[block_size];

  const auto lane = threadIdx.x & 31;
  const auto warp = threadIdx.x >> 5;
  const auto num_warps = (blockDim.x + 31) >> 5;

  state = warp_reduce(state);

  if (lane == 0) {
    warp_results[warp] = state;
  }

  __syncthreads();

  SoftmaxState block_state = init_state();
  if (warp == 0) {
    if (lane < num_warps) {
      block_state = warp_results[lane];
    }

    block_state = warp_reduce(block_state);
  }

  return block_state;
}

// softmax(input)
template <const int block_size>
__global__ void softmax_kernel(float* input, int M, int N) {
  const auto row = blockIdx.x;

  if (row >= M) {
    return;
  }

  auto state = init_state();

#pragma unroll
  for (int col = threadIdx.x; col < N; col += blockDim.x) {
    const int idx = compute_offset(row, col, N);
    state = combine(state, mono_state(input[idx]));
  }

  state = block_reduce<block_size>(state);

  __shared__ SoftmaxState final_state;
  if (threadIdx.x == 0) {
    final_state = state;
  }
  __syncthreads();

#pragma unroll
  for (int col = threadIdx.x; col < N; col += blockDim.x) {
    const int idx = compute_offset(row, col, N);
    input[idx] = __expf(input[idx] - final_state.max) / final_state.sum;
  }
}

// Q, K, V, output are device pointers
// Q: M x d
// K: N x d
// QK^T: M x N
// V: N x d
// output: M x d
extern "C" void solve(const float* Q, const float* K, const float* V,
                      float* output, int M, int N, int d) {
  auto total_items = M * N;
  float* attention;
  cudaMalloc(&attention, total_items * sizeof(float));
  cudaMemset(attention, 0, total_items * sizeof(float));

  // QK^T / sqrt(d)
  {
    constexpr auto block_size = 32;
    const auto grid_dim = dim3((N + block_size - 1) / block_size,
                               (M + block_size - 1) / block_size);
    const auto block_dim = dim3(block_size, block_size);

    multiple_kernel<block_size, true, true>
        <<<grid_dim, block_dim>>>(Q, K, attention, M, N, d);
  }

  // row-independent softmax(attention)
  {
    constexpr auto block_size = 256;
    const auto grid_dim = dim3(M);
    const auto block_dim = dim3(block_size);

    softmax_kernel<block_size><<<grid_dim, block_dim>>>(attention, M, N);
  }

  // softmax(attention) * V
  {
    constexpr auto block_size = 32;
    const auto grid_dim = dim3((d + 31) / 32, (M + 31) / 32);
    const auto block_dim = dim3(32, 32);

    multiple_kernel<block_size, false, false>
        <<<grid_dim, block_dim>>>(attention, V, output, M, d, N);
  }

  cudaDeviceSynchronize();
}

然后 multiple_kernel 其实有点问题,这里对于 transpose=false 情况下 tile_B 的读取不合并(但是没有 Bank Conflict),不能直接用 col,改一下:

template <const int block_size, const bool transpose, const bool norm>
__global__ void multiple_kernel(const float* A, const float* B, float* output,
                                int M, int N, int K) {
  __shared__ float tile_A[block_size][block_size];
  __shared__ float tile_B[block_size][block_size];

  const auto col = threadIdx.x + blockIdx.x * blockDim.x;
  const auto row = threadIdx.y + blockIdx.y * blockDim.y;

  const auto local_row = threadIdx.y;
  const auto local_col = threadIdx.x;

  float sum = 0.0f;

#pragma unroll
  for (int i = 0; i < K; i += block_size) {
    if (row < M && local_col + i < K) {
      tile_A[local_row][local_col] = A[compute_offset(row, local_col + i, K)];
    } else {
      tile_A[local_row][local_col] = 0.0f;
    }

    if constexpr (transpose) {
      // 左上角起点为 row=blockIdx.x * blockDim.x, col=i
      // threadIdx.x 是最快的,所以用它来作为 col 的自变量,相对的,用
      // threadIdx.y 作为 row 的自变量 最终得到的矩阵应该进行转置,来适配下面的
      // local 矩阵乘法
      const auto key = threadIdx.y + blockIdx.x * blockDim.x;
      if (key < N && local_col + i < K) {
        tile_B[local_col][local_row] = B[compute_offset(key, local_col + i, K)];
      } else {
        tile_B[local_col][local_row] = 0.0f;
      }
    } else {
      if (local_row + i < K && col < N) {
        tile_B[local_row][local_col] = B[compute_offset(local_row + i, col, N)];
      } else {
        tile_B[local_row][local_col] = 0.0f;
      }
    }

    __syncthreads();

#pragma unroll
    for (int j = 0; j < block_size; j++) {
      sum = fmaf(tile_A[local_row][j], tile_B[j][local_col], sum);
    }

    __syncthreads();
  }

  if (row < M && col < N) {
    if constexpr (norm) {
      output[compute_offset(row, col, N)] = sum * rsqrtf(static_cast<float>(K));
    } else {
      output[compute_offset(row, col, N)] = sum;
    }
  }
}

更高效的矩阵乘法 Kernel:

  • 对于读取矩阵部分,可以单个线程读多个元素,可以增加并行性
  • 对于计算部分,也可以一个线程负责多个元素的计算,这样也可以从 Shared Memory 中先把数据缓存到寄存器中

所以把原先的 block_size 参数拆开,先把整个矩阵分成若干 BM x BN 个子矩阵,每个 Block 负责一个子矩阵。每个子矩阵再分成若干 TM x TN 的小子矩阵,每个线程负责一个小子矩阵。这样增加了并行度和数据缓存度

然后每个 Block 单次缓存 BK 长度的子矩阵

  • 在读取的时候,每个线程读取若干个元素到 Shared Memory 中,共同得到 tile_Atile_B
  • 在计算的时候,优先枚举 k,这样可以把当前 A 矩阵的第 k 列和 B 矩阵的第 k 行先放到寄存器中,然后再处理当前线程 TM x TN 小子矩阵计算时的所有 k 的部分,累加到一个累加器 acc[TM][TN]
template <const int BM, const int BN, const int BK, const int TM, const int TN>
__global__ void multiple_kernel(const float* A, const float* B, float* output,
                                int M, int N, int K) {
  __shared__ float tile_A[BM][BK + 1];
  __shared__ float tile_B[BK][BN + 1];

  const int tx = threadIdx.x;
  const int ty = threadIdx.y;

  const int tid = threadIdx.y * blockDim.x + threadIdx.x;
  const int num_threads = blockDim.x * blockDim.y;

  // Block 负责的子矩阵的左上角坐标
  const auto block_row = blockIdx.y * BM;
  const auto block_col = blockIdx.x * BN;

  // 当前线程负责的子矩阵的左上角坐标(相对于 Block 负责的矩阵左上角)
  const int thread_row = ty * TM;
  const int thread_col = tx * TN;

  // 累加器
  float acc[TM][TN] = {};

#pragma unroll
  for (int k0 = 0; k0 < K; k0 += BK) {
    // 加载 tile_A 的数据
    for (int i = tid; i < BM * BK; i += num_threads) {
      const int local_row = i / BK;
      const int local_col = i % BK;

      const int global_row = block_row + local_row;
      const int global_col = k0 + local_col;

      if (global_row < M && global_col < K) {
        tile_A[local_row][local_col] = A[global_row * K + global_col];
      } else {
        tile_A[local_row][local_col] = 0.0f;
      }
    }

    // 加载 tile_B 的数据
    for (int i = tid; i < BK * BN; i += num_threads) {
      const int local_row = i / BN;
      const int local_col = i % BN;

      const int global_row = k0 + local_row;
      const int global_col = block_col + local_col;

      if (global_row < K && global_col < N) {
        tile_B[local_row][local_col] = B[global_row * N + global_col];
      } else {
        tile_B[local_row][local_col] = 0.0f;
      }
    }

    __syncthreads();

    // Register-tiled 计算,每个线程计算 TM x TN 块
#pragma unroll
    for (int k = 0; k < BK; ++k) {
      float a_reg[TM];
      float b_reg[TN];

#pragma unroll
      for (int i = 0; i < TM; ++i) {
        a_reg[i] = tile_A[thread_row + i][k];
      }

#pragma unroll
      for (int j = 0; j < TN; ++j) {
        b_reg[j] = tile_B[k][thread_col + j];
      }

#pragma unroll
      for (int i = 0; i < TM; ++i) {
#pragma unroll
        for (int j = 0; j < TN; ++j) {
          acc[i][j] = fmaf(a_reg[i], b_reg[j], acc[i][j]);
        }
      }
    }

    __syncthreads();
  }

#pragma unroll
  for (int i = 0; i < TM; ++i) {
    const int row = block_row + thread_row + i;

#pragma unroll
    for (int j = 0; j < TN; ++j) {
      const int col = block_col + thread_col + j;

      if (row < M && col < N) {
        output[row * N + col] = acc[i][j];
      }
    }
  }
}

{
  constexpr int BM = 64;
  constexpr int BN = 64;
  constexpr int BK = 32;
  constexpr int TM = 4;
  constexpr int TN = 4;

  dim3 block_dim(BN / TN, BM / TM);
  dim3 grid_dim((N + BN - 1) / BN, (M + BM - 1) / BM);

  multiple_kernel<BM, BN, BK, TM, TN>
      <<<grid_dim, block_dim>>>(A, B, C, M, N, K);
}

3D Convolution

输入三维体 $\boldsymbol V_{D\times H\times W}$ 和卷积核 $\boldsymbol K_{K_D\times K_H\times K_W}$,计算无填充的三维卷积 $\boldsymbol O$:

$$ \boldsymbol O_{i,j,k}=\sum_{d=0}^{K_D-1}\sum_{r=0}^{K_H-1}\sum_{c=0}^{K_W-1}\boldsymbol V_{i+d,j+r,k+c}\boldsymbol K_{d,r,c} $$

例子:

Input:
V (2x2x2): [[[1.0, 2.0],
             [3.0, 4.0]],
            [[5.0, 6.0],
             [7.0, 8.0]]]

K (2x2x2): [[[1.0, 1.0],
             [1.0, 1.0]],
            [[1.0, 1.0],
             [1.0, 1.0]]]

Output:
O (1x1x1): [[[36.0]]]

其中:

  • $1\le D,H,W\le 256$
  • $1\le K_D,K_H,K_W\le 5$
  • $K_D\le D$、$K_H\le H$、$K_W\le W$

由于卷积核很小,简单的思路就是按照结果矩阵划分 Block,每个 Block 负责一个子矩阵。每个子矩阵需要知道所有的 $\boldsymbol V$ 矩阵中的元素个数比较少,直接分 Tile 把原矩阵对应的元素读出来,放到 Shared Memory 中即可。把卷积核这个矩阵可以放到常量区

代码:

#include <cuda_runtime.h>

#define OFFSET(d, r, c, rows, cols) \
  ((d) * ((rows) * (cols)) + (r) * (cols) + (c))

__constant__ float c_kernel[256];

template <const int BD, const int BR, const int BC, const int MAX_K>
__global__ void conv3d_kernel(const float* __restrict__ input, float* output,
                              int input_depth, int input_rows, int input_cols,
                              int kernel_depth, int kernel_rows,
                              int kernel_cols) {
  __shared__ float tile[BD + MAX_K - 1][BR + MAX_K - 1][BC + MAX_K - 1];

  for (int d = threadIdx.z; d < kernel_depth + BD - 1; d += BD) {
    for (int r = threadIdx.y; r < kernel_rows + BR - 1; r += BR) {
      for (int c = threadIdx.x; c < kernel_cols + BC - 1; c += BC) {
        const auto depth = d + blockIdx.z * BD;
        const auto row = r + blockIdx.y * BR;
        const auto col = c + blockIdx.x * BC;
        if (depth < input_depth && row < input_rows && col < input_cols) {
          tile[d][r][c] =
              input[OFFSET(depth, row, col, input_rows, input_cols)];
        } else {
          tile[d][r][c] = 0.0f;
        }
      }
    }
  }

  __syncthreads();

  const auto output_depth = input_depth - kernel_depth + 1;
  const auto output_rows = input_rows - kernel_rows + 1;
  const auto output_cols = input_cols - kernel_cols + 1;

  const auto depth = threadIdx.z + blockIdx.z * BD;
  const auto row = threadIdx.y + blockIdx.y * BR;
  const auto col = threadIdx.x + blockIdx.x * BC;

  if (depth < output_depth && row < output_rows && col < output_cols) {
    float sum = 0.0f;
    for (int kd = 0; kd < kernel_depth; ++kd) {
      for (int kr = 0; kr < kernel_rows; ++kr) {
        for (int kc = 0; kc < kernel_cols; ++kc) {
          sum += tile[threadIdx.z + kd][threadIdx.y + kr][threadIdx.x + kc]
                 * c_kernel[OFFSET(kd, kr, kc, kernel_rows, kernel_cols)];
        }
      }
    }
    output[OFFSET(depth, row, col, output_rows, output_cols)] = sum;
  }
}

// input, kernel, output are device pointers
extern "C" void solve(const float* input, const float* kernel, float* output,
                      int input_depth, int input_rows, int input_cols,
                      int kernel_depth, int kernel_rows, int kernel_cols) {
  const auto kernel_size = kernel_depth * kernel_rows * kernel_cols;
  cudaMemcpyToSymbol(c_kernel, kernel, kernel_size * sizeof(float), 0,
                     cudaMemcpyDeviceToDevice);

  constexpr auto BD = 8, BR = 8, BC = 8;

  const int output_depth = input_depth - kernel_depth + 1;
  const int output_rows = input_rows - kernel_rows + 1;
  const int output_cols = input_cols - kernel_cols + 1;

  const auto grid_dim =
      dim3((output_cols + BC - 1) / BC, (output_rows + BR - 1) / BR,
           (output_depth + BD - 1) / BD);
  const auto block_dim = dim3(BC, BR, BD);

  conv3d_kernel<BD, BR, BC, 5><<<grid_dim, block_dim>>>(
      input, output, input_depth, input_rows, input_cols, kernel_depth,
      kernel_rows, kernel_cols);
}

LeetGPU 的 Solution 中最快的代码使用了 REG_COL 来提高带宽,每个线程连续计算 4 列,并且尽可能地合并读取操作,代码:

#define OFFSET(d, r, c, rows, cols) \
  ((d) * ((rows) * (cols)) + (r) * (cols) + (c))

__constant__ float c_kernel[256];

template <const int BD, const int BR, const int BX, const int KD, const int KR,
          const int KC, const int REG_COL>
__global__ void conv3d_kernel(const float* __restrict__ input, float* output,
                              int input_depth, int input_rows, int input_cols) {
  // 一个对应 REG_COL 列,所以一个 block 总共 BX * REG_COL 列
  constexpr auto TILE_D = BD + KD - 1;
  constexpr auto TILE_R = BR + KR - 1;
  constexpr auto TILE_C = BX * REG_COL + KC - 1;

  // +1 为 降低 Bank Conflict
  __shared__ float tile[TILE_D][TILE_R][TILE_C + 1];

  const auto tid = threadIdx.x + threadIdx.y * BX + threadIdx.z * BX * BR;
  const auto num_threads = BX * BR * BD;
  const auto num_elements = TILE_D * TILE_R * TILE_C;

  // 线性协同加载可以提高访存合并效果
  for (int i = tid; i < num_elements; i += num_threads) {
    const auto d = i / (TILE_R * TILE_C);
    const auto r = (i % (TILE_R * TILE_C)) / TILE_C;
    const auto c = (i % (TILE_R * TILE_C)) % TILE_C;

    const auto depth = blockIdx.z * BD + d;
    const auto row = blockIdx.y * BR + r;
    const auto col = blockIdx.x * BX * REG_COL + c;

    if (depth < input_depth && row < input_rows && col < input_cols) {
      tile[d][r][c] = input[OFFSET(depth, row, col, input_rows, input_cols)];
    } else {
      tile[d][r][c] = 0.0f;
    }
  }

  __syncthreads();

  const auto output_depth = input_depth - KD + 1;
  const auto output_rows = input_rows - KR + 1;
  const auto output_cols = input_cols - KC + 1;

  float acc[REG_COL] = {};

#pragma unroll
  for (int kd = 0; kd < KD; ++kd) {
#pragma unroll
    for (int kr = 0; kr < KR; ++kr) {
      // 当前线程需要的当前行的所有元素
      float input_reg[KC + REG_COL - 1];
#pragma unroll
      for (int v = 0; v < KC + REG_COL - 1; ++v) {
        input_reg[v] =
            tile[threadIdx.z + kd][threadIdx.y + kr][threadIdx.x * REG_COL + v];
      }

#pragma unroll
      for (int kc = 0; kc < KC; ++kc) {
        const auto kval = c_kernel[OFFSET(kd, kr, kc, KR, KC)];
#pragma unroll
        for (int rc = 0; rc < REG_COL; ++rc) {
          acc[rc] += input_reg[rc + kc] * kval;
        }
      }
    }
  }

#pragma unroll
  for (int rc = 0; rc < REG_COL; ++rc) {
    const auto depth = threadIdx.z + blockIdx.z * BD;
    const auto row = threadIdx.y + blockIdx.y * BR;
    const auto col = (threadIdx.x + blockIdx.x * BX) * REG_COL + rc;

    if (depth < output_depth && row < output_rows && col < output_cols) {
      output[OFFSET(depth, row, col, output_rows, output_cols)] = acc[rc];
    }
  }
}

这其实和前面的矩阵乘法思想相同,可以叫做 Register Tiling,核心思想就是不要让一个线程只计算一个输出,而是让它计算一小组相关输出:

Block               负责一个输出 Tile
Thread              负责 Tile 内的一个微块
Register acc[...]   保存该微块的多个输出

然后利用这些输出之间共享的输入数据:

Global Memory → Shared Memory → Registers → 多次计算

每向更快一级的存储搬一次数据,就尽可能多使用几次

通常可以归纳成五步:

  1. 选择每线程负责的输出微块,例如 TM × TN 或连续 REG_COL 个输出
  2. 找出这些输出所需输入的并集
  3. 让整个 Block 合并地把输入从全局内存加载到 Shared Memory
  4. 每个线程把自己反复使用的 Shared Memory 数据取到寄存器
  5. 在寄存器累加器中完成多个输出,最后统一写回

Prefix Sum

输入 float 数组,输出一个每个数的前缀和

例子:

Input: [5.0, -2.0, 3.0, 1.0, -4.0]
Output: [5.0, 3.0, 6.0, 7.0, 3.0]

求前缀和这个操作被叫做 Scan,有两个普遍的算法

先说一下 Step Complexity(步数复杂度)和 Work Complexity(工作量复杂度)的概念:

  • 我们可以用“盖房子”来做个生动的比喻,假设要盖一座房子,总共需要砌 10000 块砖
  • Work Complexity:就是砌完这 10000 块砖的总工作量
  • Step Complexity:就是如果雇佣无限多的工人同时干活,受限于工序先后顺序(比如必须先筑基、再砌墙、最后盖屋顶),最快需要多少个时间步才能盖完

一个是 Hillis–Steele / Kogge–Stone 风格的 Scan 算法,它是 Inclusive(第 $i$ 个结果包含它自己)的

思路大概是每轮让元素与距离为 $2^k$ 的前驱合并,Step Complexity 是 $O(\log n)$,Work Complexity 是 $O(n\log n)$。图示如下:

img

另一种是 Belloch 算法,是一种 Exclusive Scan 算法。这个比较复杂,包括 Reduce(Up-Sweep)与 Down-Sweep 两个部分,Up-Sweep 计算子树和,Down-Sweep 计算 Exclusive 前缀和

Down-Sweep 阶段从根遍历到叶子节点,节点的值代表当前子树的 Exclusive 前缀和,即子树中最靠前的元素的 Exclusive 前缀和

开始时,设置根节点的值为 $0$(根节点的 Exclusive 前缀和为 $0$),向下遍历时:

  • 父节点的值表示子树的 Exclusive 前缀和,左右子节点的值仍然表示子树的和
  • 右子节点的值应该是左子树的和加上当前子树的 Exclusive 前缀和,即父节点的值加上原始左子节点的值
  • 左子节点的值应该是当前子树的 Exclusive 前缀和,即父节点的值

这样就得到了 Exclusive 前缀和,Step Complexity 是 $O(2\log n)$,Work Complexity 是 $O(2n)$,图示:

img

可以看到,Hillis–Steele 算法在 Step Complexity 占优,Blelloch 算法在 Work Complexity 占优

类似 Reduction 的步骤,先写一个 Warp Scan,用 Hillis–Steele 算法最好,然后再写 Block Scan,然后处理所有 Block 的前缀和,最后计算 Offset 加到所有的元素上

代码如下(假设 $N<250,000$):

#include <cooperative_groups.h>
#include <cuda_runtime.h>

__device__ __forceinline__ float warp_scan(float value) {
  // clang-format off
  // Hillis–Steele / Kogge–Stone 风格的 scan 算法,假设 warp 大小是 8
  //                         0    1      2        3          4            5              6                7                  8
  // round 1 (offset = 1):   0  0+1    1+2      2+3        3+4          4+5            5+6              6+7                7+8
  // round 2 (offset = 2):   0  0+1  0+1+2  0+1+2+3    1+2+3+4      2+3+4+5        3+4+5+6          4+5+6+7            5+6+7+8
  // round 3 (offset = 4):   0  0+1  0+1+2  0+1+2+3  0+1+2+3+4  0+1+2+3+4+5  0+1+2+3+4+5+6  0+1+2+3+4+5+6+7    1+2+3+4+5+6+7+8
  // round 4 (offset = 8):   0  0+1  0+1+2  0+1+2+3  0+1+2+3+4  0+1+2+3+4+5  0+1+2+3+4+5+6  0+1+2+3+4+5+6+7  0+1+2+3+4+5+6+7+8
  // clang-format on
  const auto lane = threadIdx.x & 31;
  float other_value = 0.0f;

#pragma unroll
  for (int offset = 1; offset < 32; offset <<= 1) {
    other_value = __shfl_up_sync(0xffffffff, value, offset);
    if (lane >= offset) {
      value += other_value;
    }
  }
  return value;
}

__device__ __forceinline__ float block_scan(float value) {
  const auto warp = threadIdx.x >> 5;
  const auto lane = threadIdx.x & 31;
  const auto num_warps = blockDim.x >> 5;

  __shared__ float warp_prefix_sums[32];

  value = warp_scan(value);
  if (lane == 31) {
    warp_prefix_sums[warp] = value;
  }

  __syncthreads();

  if (warp == 0) {
    float warp_sum = lane < num_warps ? warp_prefix_sums[lane] : 0.0f;
    warp_sum = warp_scan(warp_sum);
    if (lane < num_warps) {
      warp_prefix_sums[lane] = warp_sum;
    }
  }

  __syncthreads();

  if (warp == 0) {
    return value;
  } else {
    return value + warp_prefix_sums[warp - 1];
  }
}

__device__ float g_block_prefix_sums[128 + 1];

template <const int items_per_thread>
__global__ void prefix_sum_kernel(const float* __restrict__ input,
                                  float* __restrict__ output, int N) {
  cooperative_groups::grid_group grid = cooperative_groups::this_grid();
  const auto tid = threadIdx.x + blockIdx.x * blockDim.x;

  float thread_prefix_sums[items_per_thread];
  float thread_sum = 0.0f;

#pragma unroll
  for (int i = 0; i < items_per_thread; ++i) {
    const auto idx = tid * items_per_thread + i;
    if (idx < N) {
      thread_sum += input[idx];
    }
    thread_prefix_sums[i] = thread_sum;
  }

  const auto thread_prefix_sum = block_scan(thread_sum);
  const auto thread_exclusive_prefix_sum = thread_prefix_sum - thread_sum;

  if (threadIdx.x == blockDim.x - 1) {
    g_block_prefix_sums[blockIdx.x] = thread_prefix_sum;
  }

  grid.sync();

  if (blockIdx.x == 0) {
    float value =
        threadIdx.x < gridDim.x ? g_block_prefix_sums[threadIdx.x] : 0.0f;
    value = block_scan(value);
    if (threadIdx.x < gridDim.x) {
      g_block_prefix_sums[threadIdx.x] = value;
    }
  }

  grid.sync();

  const auto block_exclusive_prefix_sum =
      blockIdx.x == 0 ? 0.0f : g_block_prefix_sums[blockIdx.x - 1];

#pragma unroll
  for (int i = 0; i < items_per_thread; ++i) {
    const auto idx = tid * items_per_thread + i;
    if (idx < N) {
      output[idx] = thread_prefix_sums[i] + block_exclusive_prefix_sum
                    + thread_exclusive_prefix_sum;
    }
  }
}

// input, output are device pointers
extern "C" void solve(const float* input, float* output, int N) {
  constexpr auto items_per_thread = 8;

  // 256 / 32 = 8
  constexpr auto block_dim = 256;

  // 250,000 / (256 * 8) = 124
  // 124 < 160(最大 block 驻留数量)
  // 124 < 256
  const auto grid_dim =
      (N + (block_dim * items_per_thread) - 1) / (block_dim * items_per_thread);

  void* args[] = {(void*)&input, (void*)&output, (void*)&N};

  void* kernel = (void*)prefix_sum_kernel<items_per_thread>;
  cudaLaunchCooperativeKernel(kernel, grid_dim, block_dim, args, 0, 0);
}

General Matrix Multiplication(GEMM)

输入矩阵 $\boldsymbol A_{M\times K}$、$\boldsymbol B_{K\times N}$ 和 $\boldsymbol C_{M\times K}$ 和 $\alpha$、$\beta$,计算:

$$ \boldsymbol C=\alpha\boldsymbol A\boldsymbol B+\beta\boldsymbol C $$

输入的矩阵全部为 FP16,即 half 类型,$\alpha$ 和 $\beta$ 是 FP32 float 类型

Tensor Core 与 WMMA

Tensor Core 是 NVIDIA Volta 架构及其后续架构(如 Ampere、Hopper、Ada Lovelace 架构)中引入的一种特殊计算单元。它们专门用于深度学习任务中的张量计算,如矩阵乘法和卷积运算(卷积运算也可以转换为矩阵乘法),最核心的就是加速矩阵乘法

对于矩阵乘法来说,CUDA Core 可以看成是一堆 FMA(Fused Multiply-Add)单元,而 Tensor Core 在硬件层面上执行 $4\times 4\times 4$ 的矩阵乘加(MMA,Matrix Multiply and Accumulate),即 $\boldsymbol D=\boldsymbol A\boldsymbol B+\boldsymbol C$

相比于 CUDA Core,Tensor Core 吞吐量提高了很多,Volta 一个 SM 中有 64 个 FP32 CUDA Core 和 8 个 Tensor Core

  • 在一个周期内,Tensor Core 可以执行 $4\times 4\times 4=64$ 次 FMA,SM 吞吐量就是 $64\times 8=512$ 次 FMA
  • 而 CUDA Core 在一个周期内只能执行一次 FMA,所以整个 SM 总共 $64$ 次 FMA

大约性能提升了八倍

此外,Tensor Core 还减少了中间数据在寄存器文件、执行单元和线程之间来回搬运的成本。假设某个线程用 CUDA Core 算 $c_{ij}=\sum_k a_{ik}b_{kj}$,典型的流程是:Shared/Global Memory -> Register File -> FP32 FMA Unit -> Register File。先读入数据到寄存器,每做一次 FMA,都要从 Register File 中读出 $a$、$b$ 和 $acc$,计算完成后还要再写回 Register File

而矩阵乘法是一个数据复用度极高的操作,Tensor Core 中,很多数据一旦进入,就可以直接在专用 datapath 中复用。做到一次读取,广播给内部的多个乘法器,中间结果不需要像 CUDA Core 的 FMA 那样频繁地写回 Register File

所以英伟达宣传性能提升了 12 倍

NVIDIA 对外提供的 Tensor Core 最主要的接口是 WMMA(Warp Matrix Multiply and Accumulate):

template<typename Use, int m, int n, int k, typename T, typename Layout=void> class fragment;

// 读取矩阵数据到 fragment,然后 Warp 内部同步
void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm);
void load_matrix_sync(fragment<...> &a, const T* mptr, unsigned ldm, layout_t layout);

// 写入 fragment 到矩阵,然后 Warp 内部同步
void store_matrix_sync(T* mptr, const fragment<...> &a, unsigned ldm, layout_t layout);

// 初始化 fragment
void fill_fragment(fragment<...> &a, const T& v);

// MMA 运算 + Warp 内部同步
void mma_sync(fragment<...> &d, const fragment<...> &a, const fragment<...> &b, const fragment<...> &c, bool satf=false);

其中:

  • fragment:Tensor Core 数据存储类,支持 matrix_amatrix_baccumulator
  • load_matrix_sync:Tensor Core 数据加载 API,支持将矩阵数据从 Shared/Global Memory 加载到 fragment
  • store_matrix_sync:Tensor Core 结果存储 API,支持将计算结果从 fragment 存储到 Shared/Global Memory
  • fill_fragmentfragment 填充 API,支持常数值填充
  • mma_sync:Tensor Core 矩阵乘计算 API,支持 $\boldsymbol D = \boldsymbol A\boldsymbol B + \boldsymbol C$ 或者 $\boldsymbol C = \boldsymbol A\boldsymbol B + \boldsymbol C$

基于 WMMA 的 GEMM

引入 WMMA 之后,整体的矩阵计算层级就变成了:

image-20260904153710930

代码:

#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <mma.h>

#define OFFSET(row, col, cols) ((row) * (cols) + (col))
#define CEIL(x, y) (((x) + (y) - 1) / (y))

// A 分成 k_tile_m x k_tile_k 的块,B 分成 k_tile_k x k_tile_n 的块,C 分成
// k_tile_m x k_tile_n 的块
constexpr int k_tile_m = 16;
constexpr int k_tile_n = 16;
constexpr int k_tile_k = 16;

// 4 个 warp,分配到两个维度上:
//     0  1  n
//   ┌─────
// 0 │ 0  1
// 1 │ 2  3
// m
constexpr int k_threads_per_warp = 32;
constexpr int k_warps_per_block_dim_m = 2;
constexpr int k_warps_per_block_dim_n = 2;
constexpr int k_warps_per_block =
    k_warps_per_block_dim_m * k_warps_per_block_dim_n;
constexpr int k_rows_per_warp_group = k_warps_per_block_dim_m * k_tile_m;
constexpr int k_cols_per_warp_group = k_warps_per_block_dim_n * k_tile_n;

constexpr int k_threads_per_block = k_warps_per_block * k_threads_per_warp;

// 每个 warp 处理 4x4 个 tile,一个 block 处理 8x8 个 tile:
//      B0 B1 B2 B3 B4 B5 B6 B7
//    ┌────────────────────────
// A0 │ 0  1  0  1  0  1  0  1
// A1 │ 2  3  2  3  2  3  2  3
// A2 │ 0  1  0  1  0  1  0  1
// A3 │ 2  3  2  3  2  3  2  3
// A4 │ 0  1  0  1  0  1  0  1
// A5 │ 2  3  2  3  2  3  2  3
// A6 │ 0  1  0  1  0  1  0  1
// A7 │ 2  3  2  3  2  3  2  3
constexpr int k_tiles_per_warp_dim_m = 4;
constexpr int k_tiles_per_warp_dim_n = 4;
constexpr int k_tiles_per_warp =
    k_tiles_per_warp_dim_m * k_tiles_per_warp_dim_n;

constexpr int k_tiles_per_block_dim_m =
    k_warps_per_block_dim_m * k_tiles_per_warp_dim_m;
constexpr int k_tiles_per_block_dim_n =
    k_warps_per_block_dim_n * k_tiles_per_warp_dim_n;

constexpr int k_rows_per_block = k_tiles_per_block_dim_m * k_tile_m;
constexpr int k_cols_per_block = k_tiles_per_block_dim_n * k_tile_n;

// 向量化加载
constexpr int k_vec_size = 8;

#define FLOAT4(value) (reinterpret_cast<float4*>(&(value))[0])
#define CFLOAT4(value) (reinterpret_cast<const float4*>(&(value))[0])

// WMMA
namespace wmma = nvcuda::wmma;

using a_fragment = wmma::fragment<wmma::matrix_a, k_tile_m, k_tile_n, k_tile_k,
                                  half, wmma::row_major>;
using b_fragment = wmma::fragment<wmma::matrix_b, k_tile_m, k_tile_n, k_tile_k,
                                  half, wmma::row_major>;
using acc_fragment =
    wmma::fragment<wmma::accumulator, k_tile_m, k_tile_n, k_tile_k, float>;
using c_fragment =
    wmma::fragment<wmma::accumulator, k_tile_m, k_tile_n, k_tile_k, half>;

__global__ void gemm_kernel(const half* __restrict__ A,
                            const half* __restrict__ B, half* __restrict__ C,
                            int M, int N, int K, float alpha, float beta) {
  __shared__ half tile_a[2][k_rows_per_block][k_tile_k + 8];
  __shared__ half tile_b[2][k_tile_k][k_cols_per_block + 8];

  int stage = 0;

  const int block_row = blockIdx.y * k_rows_per_block;
  const int block_col = blockIdx.x * k_cols_per_block;

  const int warp = threadIdx.x / k_threads_per_warp;
  const int warp_id_dim_m = warp / k_warps_per_block_dim_n;
  const int warp_id_dim_n = warp % k_warps_per_block_dim_n;

  // warp 0 -> row offset 0,  col offset 0
  // warp 1 -> row offset 0,  col offset 16
  // warp 2 -> row offset 16, col offset 0
  // warp 3 -> row offset 16, col offset 16
  const int warp_tile_row = warp_id_dim_m * k_tile_m;
  const int warp_tile_col = warp_id_dim_n * k_tile_n;

  // 先把 K-stage 的第一部分读进来,后面会边读下一阶段的元素边计算

  // 针对矩阵 A 把 block 内的所有 thread 分配到两个维度上
  //       0   1   x
  //     ┌────────
  //  0  │ 0   1
  //  1  │ 2   3
  //  2  │ 4   5
  // ... │  ...
  // 64  │ 126 127
  // y
  constexpr int a_dim_x = k_tile_k / k_vec_size,
                a_dim_y = k_threads_per_block / a_dim_x;
  const int a_thread_id_x = threadIdx.x % a_dim_x;
  const int a_thread_id_y = threadIdx.x / a_dim_x;

  // 针对矩阵 B 把 block 内的所有 thread 分配到两个维度上
  //       0   1   ... 15   x
  //     ┌────────────────
  //  0  │ 0   1   ... 15
  //  1  │ 16  17  ... 31
  // ... │         ...
  //  8  │ 112 113 ... 127
  //  y
  constexpr int b_dim_x = k_cols_per_block / k_vec_size,
                b_dim_y = k_threads_per_block / b_dim_x;
  const int b_thread_id_x = threadIdx.x % b_dim_x;
  const int b_thread_id_y = threadIdx.x / b_dim_x;

#pragma unroll
  for (int i = 0; i < k_rows_per_block; i += a_dim_y) {
    const int local_row = i + a_thread_id_y;
    const int row = block_row + local_row;

    const int col = a_thread_id_x * k_vec_size;

    FLOAT4(tile_a[stage][local_row][col]) = CFLOAT4(A[OFFSET(row, col, K)]);
  }

#pragma unroll
  for (int i = 0; i < k_tile_k; i += b_dim_y) {
    const int row = i + b_thread_id_y;

    const int local_col = b_thread_id_x * k_vec_size;
    const int col = block_col + local_col;

    FLOAT4(tile_b[stage][row][local_col]) = CFLOAT4(B[OFFSET(row, col, N)]);
  }

  //      B0   B1   B2   B3
  //    ┌────┬────┬────┬────┐
  // A0 │acc0│acc1│acc2│acc3│
  //    ├────┼────┼────┼────┤
  // A1 │acc4│acc5│acc6│acc7│
  //    ├────┼────┼────┼────┤
  // A2 │acc8│acc9│ ...│ ...│
  //    ├────┼────┼────┼────┤
  // A3 │ ...│ ...│ ...│accF│
  //    └────┴────┴────┴────┘
  a_fragment a_frag[k_tiles_per_warp_dim_m];
  b_fragment b_frag[k_tiles_per_warp_dim_n];
  acc_fragment acc_frag[k_tiles_per_warp];
  c_fragment c_frag[k_tiles_per_warp];

#pragma unroll
  for (int i = 0; i < k_tiles_per_warp; ++i) {
    wmma::fill_fragment(acc_frag[i], 0.0f);
  }

  // 加载 C 矩阵的 tile 到 c_frag 中
#pragma unroll
  for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
    const auto row = block_row + (i * k_rows_per_warp_group) + warp_tile_row;
#pragma unroll
    for (int j = 0; j < k_tiles_per_warp_dim_n; ++j) {
      const auto col = block_col + (j * k_cols_per_warp_group) + warp_tile_col;
      wmma::load_matrix_sync(c_frag[OFFSET(i, j, k_tiles_per_warp_dim_n)],
                             &C[OFFSET(row, col, N)], N, wmma::mem_row_major);
    }
  }

  __syncthreads();

  // K-stage 循环,边计算边 prefetch 下一部分的 A 和 B

  // 用于存储下一部分 A 和 B 的 tile,这里只存储当前 thread 的对应的 tile 数据
  half stage_a[k_rows_per_block / a_dim_y * k_vec_size];
  half stage_b[k_tile_k / b_dim_y * k_vec_size];

  for (int k = 0; k < K; k += k_tile_k) {
    // prefetch 到 stage_a 和 stage_b 中
    if (k + k_tile_k < K) {
#pragma unroll
      for (int i = 0; i < k_rows_per_block; i += a_dim_y) {
        const int row = block_row + i + a_thread_id_y;
        const int col = k + k_tile_k + a_thread_id_x * k_vec_size;
        const int idx = i / a_dim_y * k_vec_size;
        FLOAT4(stage_a[idx]) = CFLOAT4(A[OFFSET(row, col, K)]);
      }

#pragma unroll
      for (int i = 0; i < k_tile_k; i += b_dim_y) {
        const int row = i + k + k_tile_k + b_thread_id_y;
        const int col = block_col + b_thread_id_x * k_vec_size;
        const int idx = i / b_dim_y * k_vec_size;
        FLOAT4(stage_b[idx]) = CFLOAT4(B[OFFSET(row, col, N)]);
      }
    }

    // 从 shared memory 加载 A 和 B 到 fragment
#pragma unroll
    for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
      const int row = (i * k_rows_per_warp_group) + warp_tile_row;
      wmma::load_matrix_sync(a_frag[i], &tile_a[stage][row][0], k_tile_k + 8);
    }
#pragma unroll
    for (int i = 0; i < k_tiles_per_warp_dim_n; ++i) {
      const int col = (i * k_cols_per_warp_group) + warp_tile_col;
      wmma::load_matrix_sync(b_frag[i], &tile_b[stage][0][col],
                             k_cols_per_block + 8);
    }

    // 执行 MMA 运算
#pragma unroll
    for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
#pragma unroll
      for (int j = 0; j < k_tiles_per_warp_dim_n; ++j) {
        const int acc_idx = OFFSET(i, j, k_tiles_per_warp_dim_n);
        wmma::mma_sync(acc_frag[acc_idx], a_frag[i], b_frag[j],
                       acc_frag[acc_idx]);
      }
    }

    // 将 prefetch 的东西写入 shared memory
    if (k + k_tile_k < K) {
#pragma unroll
      for (int i = 0; i < k_rows_per_block; i += a_dim_y) {
        const int idx = i / a_dim_y * k_vec_size;
        FLOAT4(
            tile_a[stage ^ 1][i + a_thread_id_y][a_thread_id_x * k_vec_size]) =
            FLOAT4(stage_a[idx]);
      }

#pragma unroll
      for (int i = 0; i < k_tile_k; i += b_dim_y) {
        const int idx = i / b_dim_y * k_vec_size;
        FLOAT4(
            tile_b[stage ^ 1][i + b_thread_id_y][b_thread_id_x * k_vec_size]) =
            FLOAT4(stage_b[idx]);
      }

      stage ^= 1;
      __syncthreads();
    }
  }

  // 将结果写回全局内存
#pragma unroll
  for (int i = 0; i < k_tiles_per_warp_dim_m; ++i) {
#pragma unroll
    for (int j = 0; j < k_tiles_per_warp_dim_n; ++j) {
      const int acc_idx = OFFSET(i, j, k_tiles_per_warp_dim_n);

      // fragment 是一个 warp-level 分布式对象,这里当前线程只处理自己负责的元素
      for (int t = 0; t < acc_frag[acc_idx].num_elements; ++t) {
        // 这里依赖了一个 UB,不一定 c_frag.x[t] 和 acc_frag.x[t] 一一对应
        const float ab = acc_frag[acc_idx].x[t];
        half& c = c_frag[acc_idx].x[t];
        c = __float2half(alpha * ab + beta * __half2float(c));
      }

      const int row =
          block_row + (i * k_warps_per_block_dim_m * k_tile_m) + warp_tile_row;
      const int col =
          block_col + (j * k_warps_per_block_dim_n * k_tile_n) + warp_tile_col;
      wmma::store_matrix_sync(&C[OFFSET(row, col, N)], c_frag[acc_idx], N,
                              wmma::mem_row_major);
    }
  }
}

// A, B, and C are device pointers
extern "C" void solve(const half* A, const half* B, half* C, int M, int N,
                      int K, float alpha, float beta) {
  const auto Mp = CEIL(M, k_rows_per_block) * k_rows_per_block;
  const auto Np = CEIL(N, k_cols_per_block) * k_cols_per_block;
  const auto Kp = CEIL(K, k_tile_k) * k_tile_k;

  constexpr auto block_dim = k_threads_per_block;
  const auto grid_dim =
      dim3(CEIL(Np, k_cols_per_block), CEIL(Mp, k_rows_per_block));

  if (Mp == M && Np == N && Kp == K) {
    // 如果矩阵大小匹配,直接使用原矩阵
    gemm_kernel<<<grid_dim, block_dim>>>(A, B, C, M, N, K, alpha, beta);
  } else {
    // 如果矩阵大小不匹配,那么申请新的内存,创建 padded 矩阵
    half *Ap = nullptr, *Bp = nullptr, *Cp = nullptr;
    cudaMalloc(&Ap, Mp * Kp * sizeof(half));
    cudaMalloc(&Bp, Kp * Np * sizeof(half));
    cudaMalloc(&Cp, Mp * Np * sizeof(half));

    cudaMemset(Ap, 0, Mp * Kp * sizeof(half));
    cudaMemset(Bp, 0, Kp * Np * sizeof(half));
    cudaMemset(Cp, 0, Mp * Np * sizeof(half));

    // cudaMemcpy2D(
    //     dst,        // 目标起始地址
    //     dst_pitch,  // 目标每一行的跨度(字节)
    //     src,        // 源起始地址
    //     src_pitch,  // 源每一行的跨度(字节)
    //     width,      // 每行实际复制多少字节
    //     height,     // 一共复制多少行
    //     kind        // 拷贝方向
    // );
    cudaMemcpy2D(Ap, Kp * sizeof(half), A, K * sizeof(half), K * sizeof(half),
                 M, cudaMemcpyDeviceToDevice);
    cudaMemcpy2D(Bp, Np * sizeof(half), B, N * sizeof(half), N * sizeof(half),
                 K, cudaMemcpyDeviceToDevice);
    cudaMemcpy2D(Cp, Np * sizeof(half), C, N * sizeof(half), N * sizeof(half),
                 M, cudaMemcpyDeviceToDevice);

    gemm_kernel<<<grid_dim, block_dim>>>(Ap, Bp, Cp, Mp, Np, Kp, alpha, beta);

    // 拷贝结果回原矩阵
    cudaMemcpy2D(C, N * sizeof(half), Cp, Np * sizeof(half), N * sizeof(half),
                 M, cudaMemcpyDeviceToDevice);
    cudaFree(Ap);
    cudaFree(Bp);
    cudaFree(Cp);
  }
}

这里用到了向量化加载,首先将矩阵 Pad 一下,然后再调用 Kernel,结束的时候把结果 Copy 回原始矩阵中,在加载的时候,每次加载 8 个 half 元素

此外,代码里面先读入 K-Stage 的第一阶段要用的数据,然后在 K-Stage 过程中边 Prefetch 下一部分的数据边计算

不 Prefetch 时,搬数据和计算基本是串行的:

K0: |------ global load ------|-- MMA --|
K1:                                     |------ global load ------|-- MMA --|
K2:                                                                         |------ global load ------|-- MMA --|

当前代码中,使用 Ping-Pong Buffer 将“从 Global Memory 中加载下一阶段的数据”和“MMA 计算”并行执行,流水线变成这样:

K0:      |========== MMA K0 ==========|
K1 load: |------ global memory K1 ------|
                                  SMEM buffer 1
K1:                                     |========== MMA K1 ==========|
K2 load:                                |------ global memory K2 ------|
                                                                  SMEM buffer 0

这样 Latency 变成了:

$$ T_{\text{GEMM}}:T_{\text{memory}}+T_{\text{compute}}\longrightarrow\max(T_{\text{memory}},T_{\text{compute}}) $$