You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

CUDA矩阵乘法内核中共享内存存储Bank冲突减少的原因探究

CUDA GEMM内核的共享内存Bank冲突分析

我正在使用三个用于矩阵乘法的CUDA内核:

  • gemm3:共享内存GEMM的基线实现
  • gemm4:X维度线程块数量减少的实现
  • gemm5:X、Y维度线程块数量均减少的实现

经过性能分析(profiling)后,我发现这些内核的共享内存存储bank冲突数量逐渐减少,在gemm5中甚至完全消失。尽管gemm3中存储时将threadIdx.x作为最终下标本应无冲突,但profiling结果显示gemm3的冲突数量最多。同时我难以理解gemm4和gemm5中冲突减少的原理,希望深入学习该主题。

内核代码

__global__ void gemm3(const float* A, const float* B, float* C, int M, int N, int K){
    int xid = threadIdx.x  + blockIdx.x * blockDim.x;
    int yid = threadIdx.y  + blockIdx.y * blockDim.y;
    __shared__ float smem_A[BLOCK_SIZE][BLOCK_SIZE];
    __shared__ float smem_B[BLOCK_SIZE][BLOCK_SIZE];
    float sum = 0.0f;
    for (int kid = 0;kid < K / BLOCK_SIZE;kid++){
        smem_A[threadIdx.y][threadIdx.x] = A[yid * K + kid * BLOCK_SIZE + threadIdx.x];
        smem_B[threadIdx.y][threadIdx.x] = B[(kid * BLOCK_SIZE + threadIdx.y) * N + xid];
        __syncthreads();
        for (int sub_k = 0; sub_k <BLOCK_SIZE;sub_k++){
            sum += smem_A[threadIdx.y][sub_k] * smem_B[sub_k][threadIdx.x] ;
        } 
        __syncthreads();
    }    
    C[yid * N + xid] = sum;
}
__global__ void gemm4(const float* A, const float* B, float* C, int M, int N, int K){
    const int tx = threadIdx.x;
    const int ty = threadIdx.y;
    const int bx = blockIdx.x;
    const int by = blockIdx.y;
    int xid = tx + bx * blockDim.x;
    int yid = ty + by * blockDim.y;
    const int block_offset = BLOCK_SIZE * 4;
    __shared__ float smem_A[BLOCK_SIZE][BLOCK_SIZE];
    __shared__ float smem_B[BLOCK_SIZE][BLOCK_SIZE * 4];
    float sum[4] = {0,0,0,0};
    for (int kid = 0;kid < K / BLOCK_SIZE;kid++){
        smem_A[threadIdx.y][threadIdx.x] = A[yid * K + kid * BLOCK_SIZE + threadIdx.x];
        smem_B[threadIdx.y][threadIdx.x * 4]     = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4];
        smem_B[threadIdx.y][threadIdx.x * 4 + 1] = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4 + 1];
        smem_B[threadIdx.y][threadIdx.x * 4 + 2] = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4 + 2];
        smem_B[threadIdx.y][threadIdx.x * 4 + 3] = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4 + 3];
        __syncthreads();
        for (int sub_k = 0; sub_k <BLOCK_SIZE;sub_k++){
            sum[0] = fma(smem_A[threadIdx.y][sub_k],  smem_B[sub_k][threadIdx.x*4], sum[0]);
            sum[1] = fma(smem_A[threadIdx.y][sub_k],  smem_B[sub_k][threadIdx.x*4+1], sum[1]);
            sum[2] = fma(smem_A[threadIdx.y][sub_k],  smem_B[sub_k][threadIdx.x*4+2], sum[2]);
            sum[3] = fma(smem_A[threadIdx.y][sub_k],  smem_B[sub_k][threadIdx.x*4+3], sum[3]);
        } 
        __syncthreads();
    }    
    C[yid * N + bx * block_offset+ threadIdx.x * 4] = sum[0];
    C[yid * N + bx * block_offset+ threadIdx.x * 4 + 1] = sum[1];
    C[yid * N + bx * block_offset+ threadIdx.x * 4 + 2] = sum[2];
    C[yid * N + bx * block_offset+ threadIdx.x * 4 + 3] = sum[3];
}
__global__ void gemm5(const float* A, const float* B, float* C, int M, int N, int K){
    const int tx = threadIdx.x;
    const int ty = threadIdx.y;
    const int bx = blockIdx.x;
    const int by = blockIdx.y;
    int xid = tx + bx * blockDim.x;
    int yid = ty + by * blockDim.y;
    const int block_offset = BLOCK_SIZE * 4;
    __shared__ float smem_A[BLOCK_SIZE * 4][BLOCK_SIZE];
    __shared__ float smem_B[BLOCK_SIZE][BLOCK_SIZE * 4];
    float sum[4][4]= {0.f};
    for (int kid = 0;kid < K / BLOCK_SIZE;kid++){
        smem_A[threadIdx.y * 4][threadIdx.x] = A[(by * block_offset + threadIdx.y * 4)  * K + kid * BLOCK_SIZE + threadIdx.x];
        smem_A[threadIdx.y * 4 + 1][threadIdx.x] = A[(by * block_offset + threadIdx.y * 4 + 1)  * K + kid * BLOCK_SIZE + threadIdx.x];
        smem_A[threadIdx.y * 4 + 2][threadIdx.x] = A[(by * block_offset + threadIdx.y * 4 + 2)  * K + kid * BLOCK_SIZE + threadIdx.x];
        smem_A[threadIdx.y * 4 + 3][threadIdx.x] = A[(by * block_offset + threadIdx.y * 4 + 3)  * K + kid * BLOCK_SIZE + threadIdx.x];

        smem_B[threadIdx.y][threadIdx.x * 4]     = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4];
        smem_B[threadIdx.y][threadIdx.x * 4 + 1] = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4 + 1];
        smem_B[threadIdx.y][threadIdx.x * 4 + 2] = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4 + 2];
        smem_B[threadIdx.y][threadIdx.x * 4 + 3] = B[(kid * BLOCK_SIZE + threadIdx.y) * N + bx * block_offset+ threadIdx.x * 4 + 3];
        __syncthreads();
        for (int sub_k = 0; sub_k <BLOCK_SIZE;sub_k++){
            for (int i = 0; i < 4; i++){
                for (int j = 0; j < 4; j++){
                    sum[i][j] = fma(smem_A[threadIdx.y * 4 + i][sub_k],  smem_B[sub_k][threadIdx.x * 4 + j], sum[i][j]);
                }
            }
        } 
        __syncthreads();
    }
    for (int i = 0; i < 4; i++){
        for (int j = 0; j < 4; j++){
            C[(by * block_offset + threadIdx.y * 4 + i) * N + bx * block_offset + threadIdx.x * 4 + j] = sum[i][j];
        }
    }
}

性能分析结果

gemm3(...), Context 1, Stream 19
----------------------------------------------------------------------
l1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_ld.sum      0
l1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_st.sum  160811
----------------------------------------------------------------------

gemm4(...),  Context 1, Stream 20
----------------------------------------------------------------------
l1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_ld.sum      0
l1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_st.sum  57492
----------------------------------------------------------------------

gemm5(...), Context 1, Stream 21
----------------------------------------------------------------------
l1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_ld.sum      0
l1tex__data_bank_conflicts_pipe_lsu_mem_shared_op_st.sum      0
----------------------------------------------------------------------

内容的提问来源于stack exchange,提问作者Worldbuffer

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.12 17:15:59