资讯动态

CUTLASS中PitchLinearStripminedThreadMap解读

发布时间:2026/10/2 4:52:41 来源:尧图企业网站定制
写CUDA kernel的人应该都有过这种经历naive版的GEMM在小规模数据上跑得还行一旦线程块被切成不规则的形状或者你想针对特定架构的访存特性做向量化优化线程ID和坐标之间的换算就能把人绕晕。NVIDIA开源的CUTLASS把这一层彻底抽象化了PitchLinearStripminedThreadMap就是其中负责“线程块内线程到全局内存tile切片”映射的关键类之一。这篇源码解读会从模板参数、继承关系、坐标换算、谓词机制、以及与兄弟类的差异几个角度把这个类彻底讲透最后附上我在实际修改CUTLASS时踩过的几个坑。适合对CUTLASS有一定了解、想深入layout目录下thread_map.h源码的读者。1. ThreadMap在CUTLASS整个GEMM/卷积流水线里的位置1.1 全局内存Tiling与Shared Memory Stage之间被忽略的一层在CUTLASS的GEMM kernel里数据流大体是这样global memory → shared memory → register → tensor core / CUDA core → shared memory → global memory。中间有两个关键步骤一是把一整块数据从global读到shared的加载阶段二是把结果从shared写回global的写回阶段。这两个阶段都要面对同一个问题线程块内有那么多线程每个线程到底负责读哪几个元素这个问题看似简单约束条件却不少。要保证向量化访存——比如一次LDG.128读16字节要保证同一个warp的线程访问的地址尽量合并要避免shared memory的bank conflict还得让数据落到shared memory里的位置和后续register stage的读取方式匹配。ThreadMap就是在这些约束下做坐标分配的核心抽象。PitchLinearStripminedThreadMap处理的是其中最常见的一类规则形状的2D tile、PitchLinear布局下的线程映射。CUTLASS里整个tile的调度路径是TileScheduler决定这个线程块负责输出矩阵的哪一块然后把这一块表示为一个TileShape接下去进入mainloop由ThreadMap决定线程块内部256个线程怎么瓜分这个TileShape。很多人读CUTLASS源码时直接跳去看GEMM的mainloop和accumulator但真正决定访存效率的往往就是ThreadMap里那几行整数运算。1.2 从TensorLayout到ThreadMap为什么抽象边界划在这里CUTLASS里layout相关的类分两层。一层是TensorLayout描述数据在内存中的排布方式比如RowMajor就是“行的末尾紧跟下一行开头pitch等于列数乘以元素大小”另一层是ThreadMap描述线程到TensorLayout上元素坐标的映射。ThreadMap不关心数据在内存里具体是不是连续它只关心给一个thread_id和一个iteration序号返回对应的(row, column)。把边界划在这里有个实际好处kernel主体代码写的是同一个GmemIterator接口换不同ThreadMap就能适配不同的tile切分策略和架构特性不用改kernel逻辑。PitchLinearStripminedThreadMap是这一族接口在“规则2D tile”场景下用得最多的实现。对新手来说先搞清楚这个类的输入输出再回头去看GEMM的mainloop很多疑惑会自然解开——比如为什么有时候一个线程连续访问的元素在列上有时候在行上完全由ThreadMap的模板参数决定。为什么要叫PitchLinear这个名字含义是“带pitch的线性布局”也就是二维数据按行主序线性排列相邻行之间可能有额外的pitch偏移而不是恰好等于列宽。这在GEMM里很常见因为内存对齐要求会让一行末尾空出几个字节。ThreadMap必须在坐标换算时把这个pitch考虑进去否则地址就算错了。2. 模板参数与继承链PitchLinearStripminedThreadMap的骨架2.1 三个模板参数背后的三个问题打开cutlass/layout/thread_map.h找到PitchLinearStripminedThreadMap以2.8.0版本为例类声明长这样template typename Shape_, typename Threads_, int Iterations_ struct PitchLinearStripminedThreadMap : public PitchLinearThreadMapShape_, Threads_, 1 { ... };三个模板参数分别回答了三个问题切多大一块数据、用多少人来搬、每个人搬几趟。Shape_tile形状用cutlass::Shaperow, column传入例如Shape64, 64表示一个64行64列的tile切片。Threads_线程布局用cutlass::Threads1, 128这种写法表示128个线程按PitchLinear的默认习惯排在列方向。Iterations_整型常量表示每个线程需要迭代的次数源码里对应kIterations。这三个参数不是独立的它们被一个等式绑定tile总元素数 ÷ 线程数 ÷ 向量宽度 Iterations。比如一个64×64的tile128个线程向量宽度4那么每个线程需要访问32个元素也就是8次迭代Iterations8。如果你设的Iterations不满足这个等式CUTLASS在编译期就会static_assert报错或者运行期出现访问越界前者还好后者很难查。2.2 为什么继承一个固定ElementsPerAccess1的基类注意模板声明里的第三行继承的是PitchLinearThreadMapShape_, Threads_, 1最后一个“1”是基类模板参数ElementsPerAccess_。这个设计初看有点绕都被stripmined了为什么基类只允许每趟访问一个元素我个人的理解是PitchLinearThreadMap这个基类承担的是最底层的“线程坐标分布”职责它把线程均匀摊在tile的各维度上算出每个线程初始的(row, column)以及stride。而“向量宽度”“迭代次数”这些属于更上层的访存策略放在子类里处理。这样基类可以保持简单任何需要PitchLinear坐标映射的地方都能复用而不是被向量化参数绑死。PitchLinearThreadMap本身的完整模板参数是四个template typename Shape_, typename Threads_, int ElementsPerAccess_, int ThreadsPerRow_ 0 struct PitchLinearThreadMap { ... };ElementsPerAccess_表示一次性连续访问的元素个数ThreadsPerRow_可以覆盖默认的线程分布方式。StripMined版本把ElementsPerAccess_固定为1就是为了让基类只处理“每个线程访问一批标量元素”的基础分布而把真正决定向量宽度的逻辑交给上层去组合。2.3 kIterations参与strip-mined循环的方式Strip-mining是一个经典的循环变换术语把一个大循环切成多段等长的“strip”。在这个类里语义完全对应大循环是“这个线程负责访问的所有元素”每个strip是一次迭代访问的一块连续元素宽度通常是kElementsPerAccess对齐到向量宽度。在GEMM的mainloop里你会看到这样的循环结构for (int iter 0; iter kIterations; iter) { // 通过GmemIterator计算地址访问(row iter * row_stride, column)位置的数据 }迭代之间在行方向或列方向移动一个固定stridestride由tile的尺寸和迭代次数决定。这个循环在编译期就能被完全展开因为kIterations是模板常量循环边界零开销地址计算也能被大量常量折叠。这就是把Iterations放进模板参数的深层原因——不是图方便是要让编译器看到完整的访存序列生成更紧凑的指令序列。2.4 Threads布局的两种写法Threads_参数本身也是一个Shape常见写法是Threads1, 128和Threads128, 1。这两个含义完全不同。Threads1, 128表示128个线程在列方向排开适合列方向元素连续、按列合并访存的场景Threads128, 1表示128个线程在行方向排开适合按行访存的场景。PitchLinearStripminedThreadMap内部读取这个Threads_的方式是kThreads Threads::kCount也就是只取总线程数。至于你传入的是1, 128还是128, 1其实并不改变总线程数的一半计算但会影响基类PitchLinearThreadMap在计算线程分布时的语义。实际使用中我见到的绝大多数GEMM配置都是Threads1, 线程数因为CUDA的合并访存天然偏向连续线程访问连续地址列方向优先正好匹配这个硬件特性。3. 坐标换算的核心逻辑线程ID怎么变成(row, column)3.1 列优先分配先从列方向分线程这个类的映射策略是“优先列方向”。具体来说对于一个ShapeM, N的tile坐标换算分为几步先把列方向按向量宽度分成组column_groups N / kElementsPerAccessthread_id对column_groups取模得到线程落在哪个列组列坐标 列组索引 × kElementsPerAccessthread_id除以column_groups得到“行方向的线程序号”行方向的线程序号配合kIterations决定每次迭代访问哪一行写成伪代码int column_group thread_id % column_groups; int column column_group * kElementsPerAccess; int row_anchor thread_id / column_groups; // 第iter次迭代的行坐标 int row row_anchor iter * row_stride;其中row_stride通常等于M / kIterations。前提是能整除不能整除的情况靠谓词保护后面单独讲。这个设计背后的动机是访存合并。连续的thread_id对应连续的column_group也就是同一行上连续的一段地址。一个warp里的32个线程访问的地址空间是连续聚拢的LDG的合并效率很高。如果反过来让连续线程沿行方向分布同一时刻的访存会分散在好几行里合并性就差很多。对于global memory这种cache line粒度128字节的存储器来说合并与否可能就是5%和95%带宽利用率的差别。3.2 一个64x64 tile的手算验证拿一个具体例子验证tile是64×64128个线程向量宽度4Iterations8。总元素是4096128线程 × 8迭代 × 4向量宽度 4096等式成立。column_groups 64 / 4 16。来看thread_id 50这个线程怎么映射column_group 50 % 16 2column 2 × 4 8也就是第8到第11列。row_anchor 50 / 16 3。row_stride 64 / 8 8。所以第iter次迭代访问的行是3 iter × 8即第3、11、19、27、35、43、51、59行列固定在8到11。整个访问序列是分布在8行×4列的一个子块8次迭代刚好覆盖64行中的8个行带每行4个连续元素。再看相邻线程的行为。thread_id 49和thread_id 51分别落在column_group 1和3上列区间是4到7和12到15都在同一行带里。整个warpthread_id 32到63在第一次迭代时访问的是第3行的第0到63列合并成连续的64×4字节256字节正好是多个cache line的连续区域。这就是这个ThreadMap能达到高访存带宽的直接原因。3.3 行方向为什么用除法而不是取模列优先分配里列用mod、行用div。这个选择和PitchLinear布局的定义是一致的。PitchLinear看重column方向的连续性所以column必须分给连续thread_id行方向只承担“多余的线程编号”和“迭代推进”两个职责。理解了这个后面看它和PitchLinearThreadMap的差异时就顺了。还有一点值得注意M和N在Shape里是编译期常量所以上面的mod和div运算在编译期会被编译器优化成乘法和移位而不是真正的整数除法指令。这也是为什么ThreadMap的映射逻辑可以写成看起来很朴素的代码实际性能开销却几乎为零。如果M、N是运行期变量这套映射的性能模型就得重新评估了。3.4 和ColumnMajor布局的映射对比如果你熟悉ColumnMajor存储可能会问列主序的情况下PitchLinearStripminedThreadMap还适用吗答案是适用但访存模式会变差。因为PitchLinear的语义默认是“列方向连续”而ColumnMajor下连续的是行方向两者直接套用会导致线程访问的地址跨行跨列合并性大减。CUTLASS里对ColumnMajor tile有另外的ThreadMap变体本质上就是把这个类的行列角色对调。所以实际使用时先确认布局是RowMajor还是ColumnMajor再决定要不要用这个类别拿过来就套。4. 从坐标到地址GmemIterator与谓词机制怎么配合4.1 GmemIterator的职责边界PitchLinearStripminedThreadMap本身并不直接做内存访问它定义的是映射规则。真正拿坐标去算地址、发load/store指令的是GmemIterator。这个迭代器以嵌套类型的形式定义在ThreadMap内部在kernel里承担实际的数据搬运。看一段GmemIterator构造的大致逻辑GmemIterator( typename Layout::Stride stride, typename Layout::TensorCoord const extent, int thread_id, VectorType const* pointer) { // 根据thread_id和映射规则计算初始偏移 int column_group thread_id % column_groups; ... this-pointer_ pointer row * stride.row column * stride.column; }构造函数接收layout的stride、tile的extent、thread_id和基地址指针。构造之后迭代器内部就保存了当前线程对应的首地址和一系列步长。每次调用add_tile_offset或者执行operator它按照映射规则把指针推进到下一个目标。比较有意思的是CUTLASS把“坐标如何变成地址”和“下一个元素的坐标是什么”这两件事拆开了。GmemIterator只需要知道pitch和当前坐标ThreadMap负责告诉它下一个坐标怎么走。这个解耦让同一个迭代器能用在不同的ThreadMap上也正是自定义ThreadMap有意义的前提——你只要遵循迭代器的接口契约替换映射规则不影响其他代码。4.2 谓词什么时候该保护访问规则形状的tile在理想情况下不需要任何边界判断。但实际GEMM里当tile的M或N不能被线程数和迭代次数整除时总会有几个线程在某些迭代上越界。PitchLinearStripminedThreadMap内部提供Predicate嵌套类调用generate(thread_id, iteration)返回bool。比如64×64的例子如果M改成63row_stride按8算最后一个线程组的某些迭代就会访问不存在的第64行。这时候必须在load之前做predicated access否则就是非法地址访问。CUTLASS的谓词不是每次迭代都现场计算的。在mainloop的早期阶段谓词位通过位运算批量算好存成一个整型掩码循环体里直接根据掩码做分支或者predicated load。这个优化对性能很关键因为地址计算和谓词判断的开销被摊到了整个循环体之外循环内部只剩下纯粹的访存指令。看一个简化版的谓词生成struct Predicate { static bool generate(int thread_id, int iteration) { int column_group thread_id % column_groups; int row thread_id / column_groups iteration * row_stride; return row M column_group * kElementsPerAccess N; } };4.3 常见的误解谓词只保护tile边缘吗很多第一次读源码的人以为谓词只保护最右下角的那一个线程。实际上由于strip-mined映射把行方向按迭代切分当M不能被kIterations整除时中间行带的尾部迭代也可能越界。这取决于row_stride的取整方向。CUTLASS默认是向下取整尽量不越界所以尾部迭代就是谓词保护的重点区域。举个例子M 63kIterations 8row_stride 63 / 8 7整数除法向下取整。那么每个线程最后一次迭代访问的行是row_anchor 7 × 7 row_anchor 49。对于row_anchor 8的线程最后一次迭代访问第57行还在63以内没越界但如果row_stride取8row_anchor 0的线程最后一次迭代访问第56行row_anchor 8的线程访问第64行就越界了。两种取整方式下越界的位置完全不同谓词的边界条件也跟着变。修改kIterations时必须同步检查谓词的边界条件是否还成立。4.4 访存类型和向量宽度的对齐GmemIterator实际load时用的是VectorType指针通常是float4、half4这类16字节类型。CUTLASS里通过kElementsPerAccess来体现向量宽度而VectorType AlignedTypeT, kElementsPerAccess。这里有个硬性要求起始地址必须按16字节对齐否则LDG.128会直接报错或产生性能回退。地址对齐由tile在global memory中的偏移和pitch共同决定ThreadMap本身不负责对齐检查但CUTLASS的布局类在计算tile起始偏移时会尽量保证对齐。如果你自定义了tile切分方式这个对齐就是你的责任了。5. 和PitchLinearThreadMap/ImplicitGemmThreadMap的差异为什么要这样设计5.1 与PitchLinearThreadMap的本质区别不看源码的话很多人以为StripMined版只是普通版加了个循环。但从继承关系看PitchLinearThreadMap的ElementsPerAccess参数在strip-mined版里被固定成1了真正的差异在更高层PitchLinearThreadMap定义“线程在tile上的基本分布”只处理单次访问循环次数由调用者自己管理。PitchLinearStripminedThreadMap在分布基础上显式引入kIterations把“每个线程访问多少个元素”这个维度纳入类型系统让循环结构成为类型的一部分。好处是编译器可以在编译期完全展开这个循环访问序列完全确定循环边界零开销。坏处是类型变得更具体组合性变差——这也是为什么CUTLASS里还有一堆更“奇形怪状”的ThreadMap。打个比方PitchLinearThreadMap是给你一张地图和一辆车告诉你从哪出发、路怎么走但开几趟你随意PitchLinearStripminedThreadMap则把行程规划成固定的班车时刻表每趟车几点发、停哪几站全部写死。固定班车的调度效率高但灵活性差。5.2 ImplicitGemmThreadMap为什么另起炉灶在卷积kernel里CUTLASS用的是implicit GEMMtile的一个维度是output pixel另一个维度是filter × channel的组合。这个shape本身就不是规则的2D矩阵。如果硬用PitchLinearStripminedThreadMap要么把filter维强行拉平要么就得在坐标转换里做一堆除法取模访存模式会很碎。所以CUTLASS为卷积专门设计了ImplicitGemmThreadMap多了一层filter维度的处理逻辑。它和PitchLinearStripminedThreadMap解决的虽然是同一个问题——线程到坐标的映射但因为坐标系的语义不同内部的整数运算和谓词策略完全不同。比如ImplicitGemmThreadMap要考虑filter的滑动窗口偏移同一个output位置的多个filter共享一部分输入数据这些在PitchLinear版本里根本不存在。还有一个细节ImplicitGemmThreadMap的线程分布往往不是纯列优先而是把filter维和pixel维做了混合排布为了同时保证LDG合并和shared memory写入不冲突。这也说明了ThreadMap的选择不是拍脑袋决定的每个类的设计都对应一个具体的访存约束集。5.3 其他变体Swizzle、ColumnMajor和自定义ThreadMapthread_map.h里除了这两个还有SwizzleThreadMap等变体。SwizzleThreadMap主要处理shared memory的bank conflict通过交换线程到地址映射的某些位把访存冲突摊开。它和PitchLinearStripminedThreadMap是不同层面的东西一个是“数据在shared memory里怎么摆放”一个是“线程访问tile的哪个位置”但组合使用时效果会互相影响。如果你觉得自己碰到的访存模式是块状分布比如每个线程负责一个4×4微块而不是一条列上的竖条那么直接改PitchLinearStripminedThreadMap的模板参数可能不够需要考虑自定义ThreadMap。PitchLinearStripminedThreadMap的价值恰恰在于它定义了一个清晰的“ThreadMap契约”给你thread_id和iteration还你坐标。遵循这个契约你完全可以写出自己的实现替换进去。CUTLASS的kernel代码依赖的是GmemIterator接口而不是某个具体的ThreadMap类型所以替换是安全的。6. 二次封装和修改时我踩过的三个坑6.1 改了kIterations忘了同步谓词有次为了调大访存粒度我把kIterations从8改成4同时把向量宽度从4改成8。线程总数没变tile尺寸没变编译也通过了但跑出来的结果全是花屏。查了两天发现谓词逻辑里还在按kIterations8的边界算行号导致某些行越界访问读到了脏数据。这个坑的教训是ThreadMap三个模板参数是绑定的改任何一个都要重新验证“总元素数 线程数 × 迭代次数 × 向量宽度”这个等式以及相应的谓词范围。尤其是从其他项目复制配置时tile尺寸和线程数往往不成比例一定要逐个参数对一遍。6.2 线程数没有对齐warp sizeThreads参数传了一个120线程的布局以为CUTLASS会自动补齐。实际不会。CUTLASS完全信任你传入的线程数120线程意味着第4个warp只有24个活跃线程不仅访存合并效率下降而且某些架构上会导致部分bank的load宽度异常shared memory写入时还会出现bank conflict。老老实实用32的倍数更严格说是warp_size的倍数是我在这个项目上得到的最直接教训。这可能听起来很基础但当你从别人项目里抄配置时很容易忽略这个细节。CUTLASS有static_assert检查Threads::kCount是否合理但并不会强制它是32的倍数——这个责任在调用方。6.3 自定义映射时没处理pitch的坑实际GEMM里一行元素在global memory中的字节数不一定是N × sizeof(T)可能是对齐后的pitch。比如有些场景要求每行起始地址16字节对齐那pitch就会比N × sizeof(T)大。TensorLayout层面的pitch是另一个话题了但如果你手写ThreadMap或者改映射逻辑记住坐标换算算出来的是逻辑坐标真正的地址偏移 row × pitch column × sizeof(T)pitch必须从TensorLayout里取不能直接用列数乘。PitchLinearStripminedThreadMap基类内部处理了这一点但一旦你自己覆写映射逻辑这个细节就变成了你的责任。我见过有人照着行主序的公式手写地址计算结果在pitch不为默认值的矩阵上全部错位排查很久才发现是pitch的问题。6.4 静态断言错误信息看不懂怎么办CUTLASS的static_assert报错信息有时候很长几十行模板错误堆叠在一起。遇到这类编译错误我的做法是先从第一个error开始看通常是static_assert失败后面的全是模板展开噪音。然后回到“总元素数 线程数 × 迭代次数 × 向量宽度”这个等式逐一核对模板参数。如果还找不到把tile尺寸改小——比如64×64改成16×16——往往能更快暴露问题因为小tile的约束检查和错误信息更直接。6.5 调试ThreadMap映射的两个实用方法想验证自己对某个ThreadMap映射的理解是否正确最快的办法不是读注释而是在kernel里把thread_id、iteration、计算出的row/column用printf打出来。更推荐的做法是写一个host端的单元测试用纯C代码照抄映射公式算一遍坐标跟设备端printf的结果做比对。CUTLASS的ThreadMap映射逻辑是纯函数式的非常适合做这种独立验证。我每次改这类layout代码都先做这一步省下的调试时间远超写测试的时间。另外把ThreadMap的模板参数当作测试用例的输入在host端一次性测完所有边界情况比反复编译部署kernel要快得多。6.6 从性能角度做一次微调建议如果你已经跑通了一个CUTLASS GEMM kernel想通过调整PitchLinearStripminedThreadMap的参数来提升带宽利用率先看两个指标一是实际的访存字节数除以理论访问字节数二是warp级别的访存合并程度。前者可以用CUPTI或者Nsight Compute看后者可以数一下映射后同一个warp内地址的连续性。大部分情况下把向量宽度从4B提升到16B的效果最明显前提是数据和pitch对齐。对齐没问题但带宽还是上不去再考虑调整Iterations和线程数的比例让每次迭代访问的行带更宽或更窄。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑