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
相关产品推荐
相关产品推荐

