From f45230fb2c33483181dee3d44fe9eab6f3f754e6 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Wed, 26 Aug 2026 15:43:09 +0800 Subject: [PATCH] perf: pair b fragments into ldmatrix x4 loads MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - fold the two adjacent nt B fragments of each pair into one ldmatrix.x4: lanes 0-7/8-15 address rows n0..n7 chunks c/c+1, lanes 16-23/24-31 the same chunks of rows n8..n15, so {r0,r1} feed the even nt mma and {r2,r3} the odd nt — 4 x4 B loads per k-tile instead of 8 x2 (12 LDSM total, matching the decompiled cuBLAS and CUTLASS loop shapes) - the +8-row half never reaches the XOR-swizzle source bits for kK <= 64 (row[2:1]), so the pairing rides the existing per-lane address closure with one extra term (rh16 * 8 * kK); kK=128 swizzles on row[2:0] and keeps the x2 path - decompilation trail: nsys shows cuBLAS never split-Ks on the gap shapes (grid.z=1, no atomics; it fills waves with 64x128/64x64 tiles instead), and a CUTLASS 3.8 reference at our exact 128x128 s3 geometry reaches 196.5T at 2048^3 vs our 177.9T with NOP=0 and 12 LDSM — proving the loop shape is reachable from CUDA C++ (see perf/fp8_next_ideas.md F/C) Benchmark: L20 (sm_89), kernel-level sweep: 2048^3 176.7->178.1T, 4096^3 196.3->197.2T, 8192^3 ->208.7T, 4096x512x4096 137.3->139.1T, 896x1152x4096 121->123.4T. CUDA-graph e2e: 512^3 54.2->55.3T, 1536^3 142.3->143.4T, 2048^3 177.9->179.4T, 8192^3 198.2->199.0T. Four-layout C++ suite, 596 pytests pass. --- csrc/kernels/fp8/gemm.cuh | 61 ++++++++++++++++++++++++++++++++------- 1 file changed, 51 insertions(+), 10 deletions(-) diff --git a/csrc/kernels/fp8/gemm.cuh b/csrc/kernels/fp8/gemm.cuh index e6cdfab..a2680df 100644 --- a/csrc/kernels/fp8/gemm.cuh +++ b/csrc/kernels/fp8/gemm.cuh @@ -474,6 +474,20 @@ __global__ void __launch_bounds__(Traits::kCtaThreads, (a_row0 + rh8 * 8 + r7) * kK + ((rh16 ^ lswz) << 4)); const unsigned b_lane_off = static_cast((b_row0 + r7) * kK + ((rh8 ^ lswz) << 4)); + // x4-paired B loads (cuBLAS/CUTLASS loop shape): one ldmatrix.x4 feeds + // the two adjacent nt fragments — 2 x4 instead of 4 x2 per k_seg (12 + // LDSM per k-tile instead of 16). Lane contract: lanes 0-7 address + // rows n0..n7 chunk c, lanes 8-15 rows n0..n7 chunk c+1, lanes 16-23 + // rows n8..n15 chunk c, lanes 24-31 rows n8..n15 chunk c+1; regs + // {r0,r1} are the even nt's k-halves, {r2,r3} the odd nt's. The +8-row + // step never reaches the swizzle source bits for kK <= 64 (kChunks<=4: + // bits row[2:1]), so lanes 16-31 reuse the same lswz and each pair + // address is the even-nt base + p*(16*kK). kK=128 swizzles on row[2:0] + // where +8 flips bits — that config keeps the x2 loads. + constexpr bool kPairB = kK / 16 <= 4; + static_assert(!kPairB || kNt % 2 == 0, "B pairing needs even kNt"); + constexpr unsigned kPairStep = 16 * kK; // bytes per nt-pair row step + const unsigned b4_lane_off = b_lane_off + rh16 * kPairStep / 2; // Prime the pipeline. Each committed group occupies one circular shared // memory stage; the loop also handles K dimensions smaller than kStages. @@ -517,7 +531,8 @@ __global__ void __launch_bounds__(Traits::kCtaThreads, const T8* a_tile = a_base + (size_t)(tile_index % kARing) * kAStageBytes; const T8* b_tile = b_base + (size_t)(tile_index % kBRing) * kBStageBytes; const unsigned a_addr = __cvta_generic_to_shared(a_tile) + a_lane_off; - const unsigned b_addr = __cvta_generic_to_shared(b_tile) + b_lane_off; + const unsigned b_addr = __cvta_generic_to_shared(b_tile) + + (kPairB ? b4_lane_off : b_lane_off); // Per-k_seg base pair (cuBLAS's scheme): seg s lives at the seg-0 // base XOR (s<<5) — one LOP3 per extra seg per k-tile, never per // fragment. Every LDSM below addresses [base + immediate]. @@ -531,19 +546,41 @@ __global__ void __launch_bounds__(Traits::kCtaThreads, // kNt ldmatrix.x2 (B) + kMt ldmatrix.x4 (A) feed kMt*kNt*2 mma.sync // per k_seg — 0.5 load instructions per MMA, versus 4.5 scalar LDS // per MMA in the 128x64-tile version (the kernel was LSU-issue-bound - // there). B fragments double-buffer across k_segs. + // there). B fragments double-buffer across k_segs. kPairB folds the + // two adjacent nt fragments of one pair into a single x4 (see + // b4_lane_off above): kNt/2 x4 loads, regs {r0,r1}/{r2,r3} feeding + // the even/odd nt MMAs respectively. unsigned b_frag[2][kNt][2]; + unsigned b_frag4[2][kNt / 2][4]; #pragma unroll - for (int nt = 0; nt < kNt; ++nt) - astrai::ldmatrix_x2_lane(b_frag[0][nt], b_seg[0] + nt * kNtStep); + for (int p = 0; p < kNt / 2; ++p) + if constexpr (kPairB) + astrai::ldmatrix_x4_lane(b_frag4[0][p], + b_seg[0] + p * kPairStep); + else { + astrai::ldmatrix_x2_lane(b_frag[0][p * 2], + b_seg[0] + p * 2 * kNtStep); + astrai::ldmatrix_x2_lane(b_frag[0][p * 2 + 1], + b_seg[0] + (p * 2 + 1) * kNtStep); + } #pragma unroll for (int k_seg = 0; k_seg < kSegs; ++k_seg) { const int bcur = k_seg & 1, bnext = bcur ^ 1; if (k_seg + 1 < kSegs) { #pragma unroll - for (int nt = 0; nt < kNt; ++nt) - astrai::ldmatrix_x2_lane(b_frag[bnext][nt], - b_seg[k_seg + 1] + nt * kNtStep); + for (int p = 0; p < kNt / 2; ++p) + if constexpr (kPairB) + astrai::ldmatrix_x4_lane( + b_frag4[bnext][p], + b_seg[k_seg + 1] + p * kPairStep); + else { + astrai::ldmatrix_x2_lane( + b_frag[bnext][p * 2], + b_seg[k_seg + 1] + p * 2 * kNtStep); + astrai::ldmatrix_x2_lane( + b_frag[bnext][p * 2 + 1], + b_seg[k_seg + 1] + (p * 2 + 1) * kNtStep); + } } // Software-pipelined A fragments: the ldmatrix.x4 for row mt+1 // is issued before the MMAs consuming row mt, so the LDS fixed @@ -561,9 +598,13 @@ __global__ void __launch_bounds__(Traits::kCtaThreads, astrai::ldmatrix_x4_lane(a_frag[mt + 1], a_seg[k_seg] + (mt + 1) * kMtStep); #pragma unroll - for (int nt = 0; nt < kNt; ++nt) - astrai::mma_sync(acc[nt][mt], a_frag[mt], - b_frag[bcur][nt], acc[nt][mt]); + for (int nt = 0; nt < kNt; ++nt) { + const unsigned* bops = + kPairB ? (b_frag4[bcur][nt >> 1] + (nt & 1) * 2) + : b_frag[bcur][nt]; + astrai::mma_sync(acc[nt][mt], a_frag[mt], bops, + acc[nt][mt]); + } } } // Barrier 4 (lean-ring only): every thread finished reading this