
CUTLASS 2.x里的PitchLinearStripminedThreadMap说实话第一次翻到的时候我也愣了几秒。名字又长又拗口光把这个类名念顺就得练两遍。但搞懂之后会发现它是理解CUTLASS线程映射机制一个非常好的切入点gemm的epilogue输出、以及一大批tile迭代器背后都是它在默默安排“哪个线程应该碰哪块数据”。这篇文章我想把这个类彻底拆开讲清楚它解决什么问题、模板参数怎么理解、数据划分是怎么推出来的以及实际调参时会遇到哪些坑。1. 先搞清楚这名字到底是什么意思PitchLinearStripminedThreadMap拆成三段来看就清晰多了PitchLinear、Stripmined、ThreadMap。这三段分别描述了三件事内存布局、数据切分方式、线程与数据的对应关系。1.1 PitchLinear带Pitch的线性排列PitchLinear描述的是CUTLASS里一种非常基础的内存布局抽象。一个二维矩阵在全局内存里通常按行主序存放一行紧接一行。但如果每行末尾有填充或者矩阵只是某一大块缓冲区里的一个子区域那“行与行之间就不是紧挨着的”而是隔着一个固定的步长这个步长就叫Pitch。PitchLinear这个名字的意思就是在“带行距的线性内存”上做数据映射。这里的Linear是说每一行内部的数据是连续线性排列的而Pitch则处理行与行之间的跳跃。打个比方就像电影院的座位一排座位是连续的但排与排之间隔着一条过道从上一排到下一排不是走一步就能到得跨过过道。CUTLASS里用两个维度来描述这种布局Row是行方向Column是列方向。列方向是内存连续的方向行方向是带Pitch的方向。这一点后面很多推导都依赖它务必记牢。1.2 Stripmined把大块切成条带Stripmined是编译器优化里的一个老概念中文经常叫“条带化”。它的意思是不把一个大块一次性平均分给所有线程而是先把大块沿某个方向切成若干条带Strip然后每次迭代处理一条分多轮把整个块吃完。这和分页读书有点像。一本很厚的书你不会一次把整本摊在桌面上读而是几十页几十页地翻。翻完一叠再翻下一叠直到整本读完。Stripmined就是把一个大的Tile切成多个横向条带线程组每次处理一个条带。为什么要这么做一个直接原因是寄存器和shared memory都是稀缺资源。如果Tile是128x128的float一个线程要一次性处理整个Tile的1/128寄存器根本扛不住。切成条带后每个线程每轮只需要持有很小的数据量寄存器压力大幅下降也能把计算和访存更好地流水起来。1.3 ThreadMap线程和数据的对应关系ThreadMap解决的核心问题只有一个给定一块数据和一堆线程怎么规定每个线程负责哪些元素。听起来简单但加上向量化、合并访问、负载均衡这些约束后它实际上是一个挺精巧的排列组合问题。PitchLinearStripminedThreadMap就是PitchLinear布局下采用条带化切分的线程映射方案。它在CUTLASS 2.x里是应用最广的ThreadMap之一gemm epilogue里最常见的模板参数组合、PredicatedTileIterator的默认映射方式用的都是它。理解了这一个再看其他ThreadMap很多逻辑都是触类旁通。2. 模板参数逐个拆解PitchLinearStripminedThreadMap在CUTLASS 2.x源码里定义在cutlass/thread/linear_thread_map.h中模板参数包括Shape、Threads、Iterations、Alignment四个。每个参数都不是随便填的它们之间互相制约共同决定最终的线程映射结果。2.1 Shape整个Tile的形状Shape描述的是当前要处理的整个Tile的二维形状通常用cutlass::gemm::GemmShape来表示模板参数形式是ShapeRows, Columns, Depth。在PitchLinearStripminedThreadMap里重要的是前两个维度Rows和Columns。举例来说如果gemm的输出Tile是128x128那Shape就是Shape128, 128, 1其中Row128Column128。Column方向是内存连续方向也就是一行内从左到右的方向Row方向是跨行方向。这里的Shape指的是“线程映射任务的空间范围”它可以是global memory里的一个输出Tile也可以是shared memory里的一个数据块。不管是哪种映射逻辑是通用的线程组需要以某种方式把这块区域完整覆盖一遍。2.2 Threads与线程网格的分配逻辑Threads是参与映射的线程总数常见值是128、256通常是CTA里线程数或者某个协作组的大小。Threads这个参数本身不直接决定映射它需要通过一个推导逻辑组织成二维线程网格网格分成行向线程数沿Row方向和列向线程数沿Column方向。这里有一条核心原则列方向的线程数优先最大化。原因在于Column方向是连续内存方向列方向的线程越多单个线程在列方向负责的元素长度就越短而这一长度只要满足向量化要求就能让整个线程组在每行内的访问构成完美的合并访问。合并访问对全局内存吞吐的影响是数量级的所以CUTLASS在设计线程映射时总是先把列方向榨干剩下的线程再分配到行方向。2.3 Iterations一个Tile分几次吃完Iterations表示线程组把整个Tile吃完成需要迭代的次数也等于Tile沿Row方向被切成的条带数量。Iterations越大每个条带越矮单轮迭代的数据量越小寄存器压力和shared memory占用就越低但循环开销会增加。Iterations越小条带越高单轮数据量越大循环次数少但对寄存器的需求更高。实际选择Iterations时最理想的情况是Tile的Rows能被Iterations整除而且每条带的高度也能被行向线程数整除。这样所有线程在每一轮都是满负荷干活没有一个线程空转。如果除不尽CUTLASS会退化为带谓词判断的predicated访问功能上没问题但效率会打折。源码内部的推导逻辑是Strip的高度等于Shape::kRow除以Iterations。例如Shape128, 128配Iterations16时每个Strip的高度就是128/168行。2.4 Alignment向量化的硬约束Alignment表示内存访问的对齐宽度单位是“元素个数”而非字节。它的值直接决定了向量化访问的宽度。比如元素是float4字节Alignment4就对应16字节对齐的float4向量化访问Alignment8则对应32字节的向量化访问需要硬件支持。为什么Alignment是硬约束因为向量化访存指令比如float4、int4这类要求地址按向量宽度对齐如果地址没对齐就只能用标量访问性能差距非常明显。所以线程映射在设计时必须保证每个线程在列方向分配到的那段连续元素长度既能整除Tile的列数又是Alignment的整数倍。在PitchLinearStripminedThreadMap里Alignment还直接限制了列向线程数的上限因为每个线程至少要拿到Alignment个连续元素所以列向线程数不能超过Columns除以Alignment。3. ThreadLayout与StripShape是怎么推出来的参数本身好理解真正的关键在推导过程。PitchLinearStripminedThreadMap内部通过pitch_linear命名空间下的ThreadLayout和StripShape两个结构体把四个输入参数转换成线程网格形状和条带形状。3.1 线程网格分配源码逻辑CUTLASS 2.x源码里ThreadLayout的设计思路是在满足整除和对齐约束的前提下尽量找最大的列向线程数。源码逻辑可以简化为这样一个找候选数的过程// 简化后的逻辑用于理解 static int find_threads_in_row(int columns, int threads, int alignment) { // 候选列向线程数不能超过 columns / alignment int max_candidates min(threads, columns / alignment); for (int candidate max_candidates; candidate 1; --candidate) { if (threads % candidate 0 columns % candidate 0 (columns / candidate) % alignment 0) { return candidate; } } return 1; }这三个条件分别保证线程总数能被列向线程数整除Tile列数能被列向线程数整除每个线程在列方向分到的元素数满足对齐要求。三个条件全部满足后列向线程数取最大值剩下的线程自然就是行向线程数。这里源码里的命名有点容易看岔kThreadsInRow这个变量名在PitchLinear语境里其实指的是“线程网格里每一行有几个线程”而这一行恰恰对应的是数据的列方向映射也就是连续内存方向。读源码的时候不要按常规直觉去理解盯着PitchLinear布局的定义看就不会绕晕。3.2 一个具体的推导实例用最常见的输出Tile配置来走一遍Shape是128x128Threads128Alignment8Iterations16。先求列向线程数。columns / alignment 128 / 8 16所以候选最大是16。检验三个条件128 % 16 0成立128 % 16 0成立(128 / 16) % 8 0成立也就是每个线程在列方向正好分配8个连续float。于是列向线程数16行向线程数128/168。StripShape呢Tile行数128除以Iterations 16每条带高度8行宽度就是整个128列。所以一次迭代里整个线程组覆盖一个8行x128列的条带。每个线程在这个条带内负责的位置是行方向由行向线程数决定8行对应8个线程每个线程正好负责一行列方向16个线程均分128列每个线程负责8列。恰好每个线程一次迭代就是处理一行里的8个连续float一个float4向量化访问干净利落。再看一个稍微不同的例子加深理解。Shape是64x256Threads128Alignment8Iterations8。columns/alignment32候选最大32。三个条件128%320256%320(256/32)%80全部满足。列向线程数32行向线程数4。StripShape为(64/88行256列)每条带高8行每线程列方向分256/328个元素行方向8行/4线程每线程2行。这一轮每个线程实际上要处理2行每行8个float。3.3 一次迭代里线程组到底覆盖了哪些数据把上面的推导再形象化一下。以128x128 Tile、128线程、Alignment8、Iterations16为例线程组是这样工作的第一轮迭代线程组整体落在Tile最上面的8行条带里。第0行到第7行每行128列。列方向上16个线程分工第0个线程负责第0到7列第1个线程负责第8到15列依此类推。行方向上8个线程各管一行于是每个线程这一轮把自己的8个float用一条向量化指令读出来或者写出去。第一轮结束整个线程组集体下移8行进入第二轮覆盖第8到15行。重复16轮之后128行全部覆盖完毕。注意这个移动是所有线程一起移动的不存在有的线程往前走、有的线程原地等的情况。这就是条带化最直观的工作方式整个线程组作为一个整体像一行工兵排着队往前平推。每一轮迭代里线程组在行方向覆盖的高度恰好等于行向线程数这是最理想的情况。如果行方向线程数和Strip高度不完全匹配比如Strip高度是12、行向线程数是8就会出现部分线程在这一轮多跑几行、下一轮少跑几行的局面负载就会略有不均性能会有一定折损。4. 它在CUTLASS里的典型应用场景理论讲完来看实际场景。PitchLinearStripminedThreadMap在CUTLASS 2.x里最典型的藏身之处是epilogue阶段的数据回写路径以及各种以Tile为单位做数据搬移的迭代器。4.1 epilogue输出回写PredicatedTileIterator用过CUTLASS 2.x的人对PredicatedTileIterator应该不陌生。它负责把gemm计算出来的累加器结果从寄存器写回全局内存。这个迭代器内部就往死了依赖PitchLinearStripminedThreadMap来安排线程和数据的对应关系。为什么输出回写特别适合用这个ThreadMap因为gemm输出Tile天然就是一个PitchLinear布局的二维区域行主序排列、列方向连续、行与行之间有确定的pitch。更重要的是输出回写必须追求向量化store一个线程拿着连续多个float一次性写出去吞吐和只写单个float完全不是一个量级。PredicatedTileIterator这个名字里的Predicated表示它对边界情况有谓词保护。当输出Tile落在全局内存边缘比如矩阵宽度不是向量化宽度的整数倍迭代器会用谓词判断哪些store是合法的哪些需要跳过。这是PitchLinearStripminedThreadMap本身不处理的事二者一个负责映射逻辑一个负责边界安全配合得很明确。4.2 mainloop里的数据加载思路对比mainloop阶段加载A、B矩阵到shared memory时当然也可以用这个ThreadMap但很多情况下会用别的专门优化过的方案。加载global memory到shared memory时除了向量化还要考虑shared memory bank conflict以及数据重排swizzle等额外问题。对比一下就能看出PitchLinearStripminedThreadMap的设计取舍它把“线程到数据的映射”这一件事做得非常纯粹只关心怎么覆盖Tile、怎么保证向量化、怎么让所有线程同步推进。至于shared memory bank冲突、swizzle pattern这类存储相关的优化它不管那是上层迭代器或者swizzle机制要考虑的。这种职责分离是CUTLASS设计上很值得借鉴的地方每个组件解决一个问题组合起来却能应对各种复杂情况。4.3 为什么说这个ThreadMap很“通用”我在源码里翻它出现的地方印象里gemm的输出、splitk的partial sum、以及不少convolution的epilogue里都能看到它。它通用性好核心在于它的抽象粒度选得准只依赖Shape、Threads、Iterations、Alignment四个参数不关心数据具体是float还是half也不关心是读还是写。这种通用性让同一个ThreadMap能适配完全不同的计算场景。你想把一个算子从gemm搬到卷积或者从fp16改成fp32只要Tile形状和线程数、对齐方式合理ThreadMap这一层基本不用动。真正需要改的是上层迭代器如何解释数据格式而不是底层线程映射怎么排布。5. 实战中容易踩的坑和排查思路源码读明白了真到自己改算子、调性能的时候还是会踩到一些和ThreadMap相关的坑。下面这几个是我自己实际折腾CUTLASS时遇到过的问题整理成速查式的心得。5.1 Alignment不满足时的修正方法遇到过最典型的性能问题就是Tile列数不是Alignment的整数倍导致向量化store退化成标量store性能大幅下降。比如矩阵宽度是132而你想用float4访问Alignment4。132对4来说是整除的132/433各线程列方向能分到整数个float4没有问题。但如果宽度是130130%4!0那必然有线程的访问不是4字节对齐的这时候要么把Tile宽度调整成4的倍数靠PredicatedTileIterator做边界保护要么把Alignment降到1放弃向量化。实际调优时我的习惯是先检查输出矩阵的leading dimension是不是16字节对齐的。很多情况下矩阵本身宽度没问题但leading dimension没有对齐CUTLASS照样没法用宽向量化访问。这种情况调整内存分配的对齐方式往往比改ThreadMap参数性价比高得多。5.2 线程数与列数不匹配时的表现Threads和Tile列数不匹配时会出现一个容易忽略的问题列向线程数没达到理论上限导致单个线程在列方向分到的元素数超过Alignment向量化宽度变大但没被利用或者访问pattern变得不规整。举个例子Shape是128x128Threads96Alignment8。columns/alignment16候选最多16但96%16!0所以不能选16。试1596%15!0试1496%14!0试1296%120成立128%12!0不满足。最后能成立的是8128%80(128/8)%8096%80。列向线程数就落到8行向线程数12。每个线程在列方向分16个float相当于两个float4但线程网格的形状变成12行x8列一次迭代的覆盖形状和128线程的版本差异很大。排查这类问题时建议自己在调试器里打印ThreadLayout和StripShape确认线程网格形状是否符合预期。不要凭直觉以为线程数翻倍就一定列方向线程数也翻倍整除性约束会给出很多反直觉的结果。5.3 性能调优Iterations和Alignment怎么权衡Iterations和Alignment之间存在一个联动关系。Iterations增加会让Strip高度变矮如果Strip高度不是行向线程数的整数倍就会引入负载不均。Alignment如果设得过大列向线程数上限变小可能导致行向线程数变大单个线程在行方向的跳跃步数变多行方向访存的代价上升。我调一个epilogue性能时常用的套路是先固定Alignment4对应float4保证向量化然后用Tile行数除以预期的行向线程数来确定Iterations让Strip高度刚好等于行向线程数的整数倍最后用硬件性能分析工具看访存事务数量确认没有发生非合并访问。还有一个经验是Iterations并不是越小越好。有个案例里我把Iterations从16改成8Strip高度从8变成16寄存器占用看着没涨多少但因为单轮数据变大流水线排布反而更紧张性能反而下降了几个百分点。这个参数对性能的影响不是单调的必须针对具体算子和硬件实测验证。我在实际使用PitchLinearStripminedThreadMap的过程中最深的体会是它其实代表了一种CUTLASS式的思维方式——先定义清楚内存布局PitchLinear再决定切分策略Stripmined最后在严格的向量化和整除约束下推导线程映射。每一步都有明确的物理含义没有模糊空间。读这种源码时多花点时间在推导过程上后面排查性能问题会少走很多弯路。如果你手头正在改CUTLASS算子或者想自己写一个迭代器建议先把128x128这个经典配置的推导亲手算一遍建立起直觉再去看那些花哨的变体会发现套路基本都是一样的。