FEATURED · 精选文章

CUDA优化实战:用Shared Memory Swizzling消除Bank Conflict

发布时间 / 2026/8/30 9:44:30
来源 / 创域科博编辑部
栏目 / 资讯中心
CUDA优化实战:用Shared Memory Swizzling消除Bank Conflict 在 CUDA 性能优化里Shared Memory Swizzling 是很多同学从“能写对”到“能写快”的分水岭。很多 kernel 跑起来结果没问题但用 Nsight Compute 一测shared memory 的 bank conflict 高达 20 多路一次访存被硬件拆成二十多次串行操作性能自然上不去。这里的核心问题往往不是算法复杂度而是数据在 shared memory 中的银行排列和线程访问方式不匹配。这篇文章会从 bank conflict 的底层原理讲起用矩阵转置这个经典案例把 padding、XOR swizzle 两种解法完整写出来并给出可直接编译运行的.cu代码、性能验证命令和排错思路。读完你可以做到三件事第一解释为什么“每行多开一个元素”的 padding 能解决转置冲突第二写出无 bank conflict 的 shared memory tile 版本第三用 Nsight Compute 自己验证优化效果。1. 这篇文章真正要解决的问题先说你最可能遇到的场景。假设你写了一个矩阵转置 kernelglobal memory 访问已经是合并访问了代码也能算出正确结果但性能就是比预期差。很多人第一反应是“是不是块大小不对”“是不是该用 float4”很少有人会想到问题出在写回 shared memory 的那一瞬间。更具体一点你用了 32×32 的 tile把数据先从 global memory 读进 shared memory再从 shared memory 转置写回。读入阶段没问题但写回阶段所有线程如果恰好访问同一个 bank 的不同地址硬件就会把这一条 shared memory load 指令拆成 32 个 wavefront 串行执行。这个代价在没有 profiler 的情况下很难直接看到但 ncu 会非常诚实地告诉你bank conflict 很高。所以要解决的痛点就是shared memory 访问因为 bank conflict 被串行化。这篇文章不是一个百科式的概念科普而是一条从“分析冲突”到“设计 swizzle 策略”再到“上机验证”的完整路径。适合已经开始写 CUDA kernel、希望提升性能的开发者也适合正在做矩阵运算、卷积优化、图像处理 tile 化的同学。2. 基础概念Bank、Bank Conflict 与 Wavefront2.1 Shared Memory 为什么快Shared memory 是 GPU 上每个 SMStreaming Multiprocessor内部的高速存储通常和 L1 cache 共享一块物理 SRAM。它比 global memory 快一个数量级因为它在芯片内部不需要经过 L2 和显存控制器。代价是容量有限一块 GPU 上每个 SM 通常只有几十 KB 到一两百 KB。Shared memory 的访问之所以快关键不是“单次访问延迟低”而是“一次可以并行服务很多个线程”。它把存储空间划分成若干bank同一周期内不同线程访问不同 bank 时可以并行完成。这就是 shared memory 带宽的来源。2.2 Bank 是什么以最常见的 NVIDIA GPU 为例shared memory 被组织为 32 个 bank每个 bank 的宽度是 4 字节。你可以把 shared memory 想象成 32 个独立的小仓库银行编号从 0 到 31。一个地址到底属于哪个 bank按下面的方式计算地址偏移 该地址相对 shared memory 起始地址的字节数 / 4对 float 来说就是数组下标bank index 地址偏移 % 32举个例子如果数组定义成__shared__ float s[32][32]那么元素s[row][col]的下标是row * 32 colbank index 是(row * 32 col) % 32。由于 32 的倍数对 32 取模等于 0你会发现s[0][0]、s[1][0]、s[2][0]这些同一列的元素全部落在 bank 0 上。2.3 Bank Conflict 和 Wavefront当一个 warp 内的线程访问 shared memory 时硬件希望每个线程访问不同的 bank这样一次指令就能完成。如果多个线程访问同一个 bank 的不同地址就发生了 bank conflict硬件必须把一次访问拆成多个 wavefront 串行完成。访问模式是否冲突说明32 个线程访问 32 个不同 bank无冲突一个 wavefront 完成32 个线程访问同一个地址无冲突硬件广播多数架构下不算冲突32 个线程访问同一 bank 的不同地址32 路冲突拆成 32 个 wavefront性能退化部分线程冲突N 路冲突拆成 N 个 wavefront 完成为了便于理解下面用一个简化的 4-bank 模型说明真实硬件是 32 bank。假设__shared__ float s[4][4]地址下标和 bank 的映射如下s[0][0] - bank 0 s[0][1] - bank 1 s[0][2] - bank 2 s[0][3] - bank 3 s[1][0] - bank 0 s[1][1] - bank 1 s[1][2] - bank 2 s[1][3] - bank 3 s[2][0] - bank 0 s[2][1] - bank 1 s[2][2] - bank 2 s[2][3] - bank 3 s[3][0] - bank 0 s[3][1] - bank 1 s[3][2] - bank 2 s[3][3] - bank 3如果同一列上有 4 个线程同时访问s[0][0]、s[1][0]、s[2][0]、s[3][0]它们全部打在 bank 0 上4 个地址各不相同于是发生 4 路冲突。Swizzle 的本质就是改变“逻辑下标到物理 bank 的映射”让原本会撞在同一个 bank 上的访问被分散到不同 bank。3. Swizzle 的三种常用策略Swizzle 不是一个固定 API而是一类“重新排列 shared memory 布局”的思路。实现方式很多最常见的三种如下。3.1 Padding加一行空隙这是最简单、最稳妥的策略。原本每行有TILE_DIM个元素现在偏要让它占TILE_DIM 1个位置__shared__ float tile[TILE_DIM][TILE_DIM 1];这样同一列的元素在物理上不再对齐。上一节 4-bank 的例子中s[row][col]的下标变成row * 5 col取模 4 之后s[0][0] - bank 0 s[0][1] - bank 1 s[0][2] - bank 2 s[0][3] - bank 3 s[1][0] - bank 1 s[1][1] - bank 2 s[1][2] - bank 3 s[1][3] - bank 0 s[2][0] - bank 2 s[2][1] - bank 3 s[2][2] - bank 0 s[2][3] - bank 1 s[3][0] - bank 3 s[3][1] - bank 0 s[3][2] - bank 1 s[3][3] - bank 2同一列上的元素被散到了不同 bank。这个技巧理解成本低几乎不会写错代价是多占一点点 shared memory。对 32×32 的 float tile 来说每行从 32 个 float 变成 33 个 float额外开销约 3%。3.2 XOR Swizzle当 tile 宽度是 2 的幂时可以用异或运算重新映射下标// 存储侧 tile[row][col ^ row] value; // 读取侧 value tile[row][col ^ row];为什么 XOR 能避免冲突以 32 为宽度时col ^ row会把 row 的 bit 混进 col 里让原本同一列的地址在 bank 编号上错开。这个策略比 padding 更“酷”但要求 tile 宽度必须是 2 的幂否则col ^ row可能越界而且存取两侧必须用同一套规则写错了非常难排查。3.3 更通用的 Mask / Permutation 变换如果你处理的场景不是简单的方阵转置而是卷积、直方图、三角矩阵等特殊访问模式padding 和 XOR 都可能不够用。这时的通用做法是设计一个__device__函数显式完成“逻辑下标 - 物理下标”的映射__device__ __forceinline__ int swizzle_index(int row, int col, int pitch) { // 例如把低 bit 和高 bit 交换或者与一个 mask 异或 return row * pitch ((col (row 3)) 31); }这种方式的优点是灵活缺点是需要针对具体访存模式分析 bank 分布代码可读性下降。我的建议是能 padding 就 paddingpadding 解决不了再用 XOR最后才考虑复杂 permutation。3.4 三种策略对比策略实现难度适用场景缺点Padding最低大多数按行存储、按列读取的场景额外 shared memory 开销XOR Swizzle中宽度为 2 的幂的方阵 tile宽度非 2 的幂时失效存取需一致自定义 Permutation高复杂、非规则访存模式分析成本高易出错4. 典型场景矩阵转置中的 Bank Conflict 分析矩阵转置是最适合讲清楚 swizzle 的案例因为它的访存模式非常对称读入时按行连续访问写回时按列访问。先用一个不使用 swizzle 的朴素 shared memory 版本作为基线#define TILE_DIM 32 __global__ void transpose_tiled(const float* in, float* out, int width) { __shared__ float tile[TILE_DIM][TILE_DIM]; int x blockIdx.x * TILE_DIM threadIdx.x; int y blockIdx.y * TILE_DIM threadIdx.y; if (x width y width) { tile[threadIdx.y][threadIdx.x] in[y * width x]; } __syncthreads(); int x2 blockIdx.y * TILE_DIM threadIdx.x; int y2 blockIdx.x * TILE_DIM threadIdx.y; if (x2 width y2 width) { out[y2 * width x2] tile[threadIdx.x][threadIdx.y]; } }分析这个 kernel 的两个阶段4.1 读入阶段无冲突读入时线程(threadIdx.x, threadIdx.y)访问tile[threadIdx.y][threadIdx.x]。在一个 warp 内threadIdx.y是相同的threadIdx.x从 0 到 31 变化所以访问的是同一行、不同列bank 编号正好是(row * 32 col) % 32 col。32 个线程访问 32 个不同 bank无冲突。4.2 写回阶段32 路冲突写回时线程(threadIdx.x, threadIdx.y)访问tile[threadIdx.x][threadIdx.y]。这里行和列交换了。在一个 warp 内threadIdx.y固定threadIdx.x变化于是访问的是不同行、同一列。以threadIdx.y 0为例地址下标为threadIdx.x * 32 0bank index 是(threadIdx.x * 32) % 32 0。也就是说这个名字上不同的 32 个地址全部落在 bank 0 上并且不是同一个地址于是产生 32 路 bank conflict。硬件会把一次 shared memory load 拆成 32 个 wavefront性能退化非常明显。4.3 Padding 如何解决把 shared memory 定义改成__shared__ float tile[TILE_DIM][TILE_DIM 1];此时地址下标变成threadIdx.x * (TILE_DIM 1) threadIdx.y。因为33 % 32 1所以 bank index 大约是(threadIdx.x threadIdx.y) % 32。warp 内threadIdx.y固定threadIdx.x从 0 到 31bank 编号各不相同。冲突被消除了。4.4 XOR Swizzle 如何解决XOR 版本会在存储和读取两侧都做异或__global__ void transpose_xor(const float* in, float* out, int width) { __shared__ float tile[TILE_DIM][TILE_DIM]; int x blockIdx.x * TILE_DIM threadIdx.x; int y blockIdx.y * TILE_DIM threadIdx.y; if (x width y width) { tile[threadIdx.y][threadIdx.x ^ threadIdx.y] in[y * width x]; } __syncthreads(); int x2 blockIdx.y * TILE_DIM threadIdx.x; int y2 blockIdx.x * TILE_DIM threadIdx.y; if (x2 width y2 width) { out[y2 * width x2] tile[threadIdx.x][threadIdx.y ^ threadIdx.x]; } }存储时线程(tx, ty)写入tile[ty][tx ^ ty]。warp 内ty固定tx变化所以列下标tx ^ ty各不相同bank 不冲突。读取时线程(tx, ty)读tile[tx][ty ^ tx]同理warp 内tx变化读到的物理下标ty ^ tx各不相同也不冲突。XOR 版本表面上没有增加 shared memory但它要求TILE_DIM是 2 的幂并且在转置时能保证存取规则自洽。如果 tile 宽度不是 2 的幂就不要硬套 XOR。5. 完整代码实现四个版本对比下面给出一份可以直接编译运行的完整代码。它包含四个 kerneltranspose_naive不经过 shared memory直接 global memory 转置。transpose_tiled共享内存 tile有 bank conflict 的版本。transpose_padded共享内存 tile padding。transpose_xor共享内存 tile XOR swizzle。// 文件路径transpose.cu // 编译nvcc -O3 -archcompute_86 -codesm_86 transpose.cu -o transpose // 运行./transpose 1024 #include cstdio #include cstdlib #include cmath #include cuda_runtime.h #define TILE_DIM 32 #define CHECK_CUDA(call) \ do { \ cudaError_t err_ (call); \ if (err_ ! cudaSuccess) { \ fprintf(stderr, CUDA error %s:%d: %s\n, __FILE__, __LINE__, \ cudaGetErrorString(err_)); \ exit(1); \ } \ } while (0) // 1. 朴素 global memory 转置 __global__ void transpose_naive(const float* in, float* out, int width) { int x blockIdx.x * blockDim.x threadIdx.x; int y blockIdx.y * blockDim.y threadIdx.y; if (x width y width) { out[x * width y] in[y * width x]; } } // 2. 共享内存 tile无 padding写回阶段有 bank conflict __global__ void transpose_tiled(const float* in, float* out, int width) { __shared__ float tile[TILE_DIM][TILE_DIM]; int x blockIdx.x * TILE_DIM threadIdx.x; int y blockIdx.y * TILE_DIM threadIdx.y; if (x width y width) {
RELATED — 相关阅读

相关资讯

LATEST — 最新资讯

最新发布

TODAY — 本日精选

新闻

WEEKLY — 本周精选

新闻

MONTHLY — 本月精选

新闻