Skip to content
KernelIndex
Search⌘K

submission 505548

mufeez-amjad · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1706 lines, June 9 Researcher Reciprocity License v1.0.

v4d.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-505548?include=source"
interfacepython
Compatibility
measured onNVIDIA B200
declared hardwareNVIDIA B200
architecturessm_100
dtypesfp8_e4m3, nvfp4

Benchmark evidence

1 measurement across 1 GPU, fastest first.

Operation / workload
Hardware
Latency
Rank
Observed
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
21.0µs
#101 of 310
2026-02-21

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:68b4f80cbfdecadc8afeecf3aff87266efec17588c85a9d66872b62a6501818b
license declaredunknown
license concludedunknown
authorsmufeez-amjad
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

fused-epiloguetemplate <int BLOCK_M, int BLOCK_N, bool LOW_M_EPILOGUE>
mbarrier__device__ inline void mbarrier_init(int mbar_addr, int count) {
num-warps = 8constexpr int TMA_NUM_WARPS = 8;
persistent-kerneltemplate <bool PERSISTENT, int BLOCK_M, int BLOCK_N, int NS, int CLUSTER_SIZE = 1>
shared-memory__device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z,
tcgen05asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
tile-k = 256constexpr int TMA_BLOCK_K = 256;
tile-n = 128constexpr int TMA_BLOCK_N = 128;
tmaCUtensorMap A_tmap;
vector-width = half2half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset);

Kernel source

v4d.py1706 lines
#!POPCORN leaderboard nvfp4_group_gemm
#!POPCORN gpu NVIDIA

from __future__ import annotations

from functools import lru_cache
from typing import cast

import torch
from torch.utils.cpp_extension import load_inline

from task import input_t, output_t

"""
g: 8; k: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; m: [80, 176, 128, 72, 64, 248, 96, 160]; n: [4096, 4096, 4096, 4096, 4096, 4096, 4096, 4096]; seed: 1111
 ⏱ 47.5 ± 0.00 µs
 ⚡ 47.4 µs 🐌 47.5 µs

g: 8; k: [2048, 2048, 2048, 2048, 2048, 2048, 2048, 2048]; m: [40, 76, 168, 72, 164, 148, 196, 160]; n: [7168, 7168, 7168, 7168, 7168, 7168, 7168, 7168]; seed: 1111
 ⏱ 45.0 ± 0.04 µs
 ⚡ 44.6 µs 🐌 45.3 µs

g: 2; k: [4096, 4096]; m: [192, 320]; n: [3072, 3072]; seed: 1111
 ⏱ 14.3 ± 0.01 µs
 ⚡ 13.9 µs 🐌 14.6 µs

g: 2; k: [1536, 1536]; m: [128, 384]; n: [4096, 4096]; seed: 1111
 ⏱ 10.5 ± 0.01 µs
 ⚡ 10.4 µs 🐌 10.7 µs
"""

CUDA_SRC = """
#include <algorithm>
#include <cstring>
#include <limits>
#include <vector>
#include <unordered_map>
#include <cstdint>
#include <cuda.h>
#include <cudaTypedefs.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAException.h>

static inline int ceil_div(int a, int b) { return (a + b - 1) / b; }

#define CUDA_CHECK(expr)                                                       \\
  do {                                                                         \\
    cudaError_t _err = (expr);                                                 \\
    TORCH_CHECK(_err == cudaSuccess, "CUDA error: ", cudaGetErrorString(_err)); \\
  } while (0)

static inline uint64_t hash_combine_u64(uint64_t h, uint64_t x) {
  h ^= x;
  h *= 1099511628211ULL;
  return h;
}

constexpr int WARP_SIZE = 32;
constexpr int MMA_K = 64;

constexpr uint64_t EVICT_FIRST = 0x12F0000000000000;
constexpr uint64_t EVICT_LAST = 0x14F0000000000000;

struct WorkItem {
  int problem_idx;
  int tile_m;
  int tile_n;
};

// Per-problem metadata.
// Align on 128 byte boundary, useful since this is read by many CTAs.
struct __align__(128) ProblemInfo {
  CUtensorMap A_tmap;
  CUtensorMap B_tmap;
  CUtensorMap B_tmap_256;
  const char* SFA_ptr;
  const char* SFB_ptr;
  half* C_ptr;
  int M, N, K;
  int64_t Cs0, Cs1;
};

// tcgen05 descriptors encode shared-memory addresses in 16-byte units.
// Mask to the HW-supported address width and drop the 16B alignment bits.
__device__ inline constexpr uint64_t desc_encode(uint64_t x) {
  return (x & 0x3'FFFFULL) >> 4ULL;
}

// elect.sync: use this to have a single lane issue TMA/tcgen05 instructions while the
// whole warp stays converged.
__device__ uint32_t elect_sync() {
  uint32_t pred = 0;
  asm volatile(
    "{\\n\\t"
    ".reg .pred %%px;\\n\\t"
    "elect.sync _|%%px, %1;\\n\\t"
    "@%%px mov.s32 %0, 1;\\n\\t"
    "}"
    : "+r"(pred)
    : "r"(0xFFFFFFFF)
  );
  return pred;
}

__device__ inline uint32_t get_cluster_ctarank() {
  uint32_t rank;
  asm volatile("mov.u32 %0, %%cluster_ctarank;" : "=r"(rank));
  return rank;
}

__device__ inline void cluster_sync() {
  asm volatile("barrier.cluster.arrive;" ::: "memory");
  asm volatile("barrier.cluster.wait;" ::: "memory");
}

// Shared-memory mbarrier helpers.
// Used for:
// - TMA completion barrier: consumer waits for bytes to arrive in shared memory.
// - Stage reuse barrier: producer waits until MMA is done with a stage before
//   overwriting it.
__device__ inline void mbarrier_init(int mbar_addr, int count) {
  asm volatile("mbarrier.init.shared::cta.b64 [%0], %1;" :: "r"(mbar_addr), "r"(count));
}

// Program the expected byte count for a TMA stage and arrive.
// This must happen before issuing any cp.async.bulk.* that completes to the
// barrier.
__device__ inline void mbarrier_arrive_expect_tx(int mbar_addr, int size) {
  asm volatile("mbarrier.arrive.expect_tx.release.cta.shared::cta.b64 _, [%0], %1;"
              :: "r"(mbar_addr), "r"(size) : "memory");
}

constexpr uint32_t MBAR_WAIT_HINT_TMA = 0x989680;
constexpr uint32_t MBAR_WAIT_HINT_REUSE = 64;

__device__ __forceinline__ void mbarrier_wait_hint(int mbar_addr, int phase, uint32_t suspend_time_hint) {
  asm volatile(
    "{\\n\\t"
    ".reg .pred P1;\\n\\t"
    "LAB_WAIT:\\n\\t"
    "mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\\n\\t"
    "@!P1 bra.uni LAB_WAIT;\\n\\t"
    "}"
    :: "r"(mbar_addr), "r"(phase), "r"(suspend_time_hint)
  );
}

__device__ __forceinline__ void mbarrier_wait_tma(int mbar_addr, int phase) {
  mbarrier_wait_hint(mbar_addr, phase, MBAR_WAIT_HINT_TMA);
}

__device__ __forceinline__ void mbarrier_wait_reuse(int mbar_addr, int phase) {
  mbarrier_wait_hint(mbar_addr, phase, MBAR_WAIT_HINT_REUSE);
}

// TMA: 3D tensor-map load from global -> shared memory.
// The (x,y,z) coordinates correspond to the CUtensorMap encoding in
// init_AB_tmap_u4.
template <int CTA_GROUP = 1>
__device__ inline void tma_3d_gmem2smem(int dst, const void *tmap_ptr, int x, int y, int z,
                                         int mbar_addr, uint64_t cache_policy) {
  asm volatile(
    "cp.async.bulk.tensor.3d.shared::cta.global.mbarrier::complete_tx::bytes.cta_group::%7.L2::cache_hint "
    "[%0], [%1, {%2, %3, %4}], [%5], %6;"
    :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "l"(cache_policy), "n"(CTA_GROUP)
    : "memory"
  );
}

__device__ inline int mapa_cta_to_cluster(int cta_addr, int dest_cta) {
  int cluster_addr;
  asm volatile("mapa.shared::cluster.u32 %0, %1, %2;" : "=r"(cluster_addr) : "r"(cta_addr), "r"(dest_cta));
  return cluster_addr;
}

__device__ inline void mbarrier_arrive_cluster(int mbar_cluster_addr) {
  asm volatile("mbarrier.arrive.shared::cluster.b64 _, [%0];"
    :: "r"(mbar_cluster_addr) : "memory");
}

// Cluster multicast variant of 3D TMA.
// dst and mbar_addr are in shared::cluster address space.
// The multicast mask selects which CTAs in the cluster receive the data.
__device__ inline void tma_3d_gmem2smem_multicast(int dst, const void *tmap_ptr, int x, int y, int z,
                                                   int mbar_addr, uint16_t multicast_mask) {
  asm volatile(
    "cp.async.bulk.tensor.3d.shared::cluster.global.mbarrier::complete_tx::bytes.multicast::cluster.cta_group::1 "
    "[%0], [%1, {%2, %3, %4}], [%5], %6;"
    :: "r"(dst), "l"(tmap_ptr), "r"(x), "r"(y), "r"(z), "r"(mbar_addr), "h"(multicast_mask)
    : "memory"
  );
}

// Linear bulk copy global -> shared.
// Used for scale-factor tensors (SFA/SFB).
__device__ inline void tma_gmem2smem(int dst, const void *src, int size, int mbar_addr, uint64_t cache_policy) {
  asm volatile(
    "cp.async.bulk.shared::cta.global.mbarrier::complete_tx::bytes.L2::cache_hint "
    "[%0], [%1], %2, [%3], %4;"
    :: "r"(dst), "l"(src), "r"(size), "r"(mbar_addr), "l"(cache_policy)
    : "memory"
  );
}

__device__ inline void tcgen05_cp_nvfp4(int taddr, uint64_t s_desc) {
  asm volatile("tcgen05.cp.cta_group::1.32x128b.warpx4 [%0], %1;" :: "r"(taddr), "l"(s_desc));
}

__device__ inline void tcgen05_commit(int mbar_addr) {
  asm volatile(
    "tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%0];\\n"
    :: "r"(mbar_addr) : "memory"
  );
}

__device__ inline void tcgen05_mma_nvfp4(
  uint64_t a_desc, uint64_t b_desc, uint32_t i_desc,
  int scale_A_tmem, int scale_B_tmem, int enable_input_d
) {
  const int d_tmem = 0;
  asm volatile(
    "{\\n\\t"
    ".reg .pred p;\\n\\t"
    "setp.ne.b32 p, %6, 0;\\n\\t"
    "tcgen05.mma.cta_group::1.kind::mxf4nvf4.block_scale.block16 [%0], %1, %2, %3, [%4], [%5], p;\\n\\t"
    "}"
    :: "r"(d_tmem), "l"(a_desc), "l"(b_desc), "r"(i_desc),
       "r"(scale_A_tmem), "r"(scale_B_tmem), "r"(enable_input_d)
  );
}

inline constexpr char SHAPE_16x256b[] = ".16x256b";
inline constexpr char NUM_x8[] = ".x8";
inline constexpr char NUM_x16[] = ".x16";

template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_32regs(float *tmp, int row, int col) {
  asm volatile("tcgen05.ld.sync.aligned%33%34.b32 "
              "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
              "  %8,  %9, %10, %11, %12, %13, %14, %15, "
              " %16, %17, %18, %19, %20, %21, %22, %23, "
              " %24, %25, %26, %27, %28, %29, %30, %31}, [%32];"
              : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
                "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
                "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
                "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31])
              : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}

template <const char *SHAPE, const char *NUM>
__device__ inline
void tcgen05_ld_64regs(float *tmp, int row, int col) {
  asm volatile("tcgen05.ld.sync.aligned%65%66.b32 "
              "{ %0,  %1,  %2,  %3,  %4,  %5,  %6,  %7, "
              "  %8,  %9, %10, %11, %12, %13, %14, %15, "
              " %16, %17, %18, %19, %20, %21, %22, %23, "
              " %24, %25, %26, %27, %28, %29, %30, %31, "
              " %32, %33, %34, %35, %36, %37, %38, %39, "
              " %40, %41, %42, %43, %44, %45, %46, %47, "
              " %48, %49, %50, %51, %52, %53, %54, %55, "
              " %56, %57, %58, %59, %60, %61, %62, %63}, [%64];"
              : "=f"(tmp[ 0]), "=f"(tmp[ 1]), "=f"(tmp[ 2]), "=f"(tmp[ 3]), "=f"(tmp[ 4]), "=f"(tmp[ 5]), "=f"(tmp[ 6]), "=f"(tmp[ 7]),
                "=f"(tmp[ 8]), "=f"(tmp[ 9]), "=f"(tmp[10]), "=f"(tmp[11]), "=f"(tmp[12]), "=f"(tmp[13]), "=f"(tmp[14]), "=f"(tmp[15]),
                "=f"(tmp[16]), "=f"(tmp[17]), "=f"(tmp[18]), "=f"(tmp[19]), "=f"(tmp[20]), "=f"(tmp[21]), "=f"(tmp[22]), "=f"(tmp[23]),
                "=f"(tmp[24]), "=f"(tmp[25]), "=f"(tmp[26]), "=f"(tmp[27]), "=f"(tmp[28]), "=f"(tmp[29]), "=f"(tmp[30]), "=f"(tmp[31]),
                "=f"(tmp[32]), "=f"(tmp[33]), "=f"(tmp[34]), "=f"(tmp[35]), "=f"(tmp[36]), "=f"(tmp[37]), "=f"(tmp[38]), "=f"(tmp[39]),
                "=f"(tmp[40]), "=f"(tmp[41]), "=f"(tmp[42]), "=f"(tmp[43]), "=f"(tmp[44]), "=f"(tmp[45]), "=f"(tmp[46]), "=f"(tmp[47]),
                "=f"(tmp[48]), "=f"(tmp[49]), "=f"(tmp[50]), "=f"(tmp[51]), "=f"(tmp[52]), "=f"(tmp[53]), "=f"(tmp[54]), "=f"(tmp[55]),
                "=f"(tmp[56]), "=f"(tmp[57]), "=f"(tmp[58]), "=f"(tmp[59]), "=f"(tmp[60]), "=f"(tmp[61]), "=f"(tmp[62]), "=f"(tmp[63])
              : "r"((row << 16) | col), "C"(SHAPE), "C"(NUM));
}

__device__ inline void tcgen05_ld_16x256b_x8(float *tmp, int row, int col) {
  tcgen05_ld_32regs<SHAPE_16x256b, NUM_x8>(tmp, row, col);
}

__device__ inline void tcgen05_ld_16x256b_x16(float *tmp, int row, int col) {
  tcgen05_ld_64regs<SHAPE_16x256b, NUM_x16>(tmp, row, col);
}

__device__ __forceinline__ void tcgen05_dealloc_cols_cta1(uint32_t tmem, int count) {
  asm volatile(
    "tcgen05.dealloc.cta_group::1.sync.aligned.b32 %0, %1;\\n"
    :: "r"(tmem), "r"(count)
    : "memory"
  );
}

constexpr int TMA_BLOCK_N = 128;
constexpr int TMA_BLOCK_K = 256;
constexpr int TMA_NUM_WARPS = 8;
constexpr int MMA_M = 128;

constexpr int LOW_M_THRESHOLD = 96;

constexpr int NS_DEEP_HI = 6;
constexpr int NS_WIDE_HI = 4;
constexpr int NS_64_HI = 8;
constexpr int NS_64_LO = 4;
constexpr int CLUSTER_SIZE_256 = 4;

enum AlgoKind : uint8_t {
  ALGO_128 = 0,
  ALGO_256 = 1,
  ALGO_64 = 2,
};

constexpr int mbar_bytes(int ns, int arrivals_per_stage) {
  return ((arrivals_per_stage * ns * 8 + 63) & ~63);
}

constexpr int stage_bytes(int block_n) {
  const int sfb_width = (block_n < 128) ? 128 : block_n;
  return MMA_M * (TMA_BLOCK_K / 2) + block_n * (TMA_BLOCK_K / 2)
       + MMA_M * (TMA_BLOCK_K / 16) + sfb_width * (TMA_BLOCK_K / 16);
}

constexpr int smem_bytes(int block_n, int ns, int arrivals_per_stage, int mbar_sets = 1) {
  return stage_bytes(block_n) * ns + mbar_bytes(ns, arrivals_per_stage) * mbar_sets;
}

constexpr int block_n_for_algo(uint8_t algo) {
  return (algo == ALGO_64) ? 64 : ((algo == ALGO_256) ? 256 : 128);
}

constexpr int TMA_WARP = 4;
// Use a second (otherwise idle) warp to issue B/SFB TMA in parallel.
// This cuts producer-side latency and reduces MMA-side barrier stalls.
constexpr int TMA_WARP_B = 6;
constexpr int MMA_WARP = 5;

constexpr int tma_expected_tx_bytes(bool do_A, int a_bytes, int b_bytes, int sfa_bytes, int sfb_bytes) {
  return do_A ? (a_bytes + sfa_bytes) : (b_bytes + sfb_bytes);
}

// Epilogue: read fp32 accumulators from TMEM (tcgen05.ld) and store fp16 C.
// Only threads with tid < BLOCK_M participate; this maps 4 warps (0..3) to the
// 128 output rows, with each warp handling a 32-row stripe.
template <int BLOCK_M, int BLOCK_N, bool LOW_M_EPILOGUE>
__device__ __forceinline__ void epilogue_store(
    const ProblemInfo& prob,
    int m_offset,
    int n_offset,
    int tid,
    int warp_id,
    int lane_id
) {
  if (tid >= BLOCK_M) return;

  const int M = prob.M;
  const int N = prob.N;
  half* C_ptr = prob.C_ptr;
  const int64_t Cs0 = prob.Cs0;
  const int64_t Cs1 = prob.Cs1;

  const bool full_n = (n_offset + BLOCK_N <= N);
  const bool full_m = (m_offset + BLOCK_M <= M);
  const bool full_tile = full_n && full_m;
  const bool contiguous = (Cs1 == 1);

  const int warp_row_base = m_offset + warp_id * 32;
  if (LOW_M_EPILOGUE && warp_row_base >= M) return;

  int m_iters = 2;
  if (LOW_M_EPILOGUE) {
    const int remaining = M - warp_row_base;
    m_iters = (remaining <= 16) ? 1 : 2;
  }

  // Lane mapping: each lane owns two columns (half2) and one of 8 rows.
  const int lane_row = lane_id >> 2;
  const int lane_col = (lane_id & 3) * 2;

  if constexpr (BLOCK_N == 64) {
    constexpr int WIDTH = 64;
    for (int m = 0; m < m_iters; ++m) {
      float tmp[32];
      tcgen05_ld_16x256b_x8(tmp, warp_id * 32 + m * 16, 0);
      asm volatile("tcgen05.wait::ld.sync.aligned;");

      const int row0 = warp_row_base + m * 16 + lane_row;
      const int row1 = row0 + 8;

      if (contiguous) {
        if (full_tile) {
          half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset);
          half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset);
          #pragma unroll
          for (int i = 0; i < WIDTH / 8; i++) {
            const int idx = i * 4;
            const int col = i * 8 + lane_col;
            const int h2_idx = col >> 1;
            row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
            row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
          }
          continue;
        }

        const bool row0_in = row0 < M;
        const bool row1_in = row1 < M;
        if (full_n) {
          half2* row0_ptr = row0_in ? reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset) : nullptr;
          half2* row1_ptr = row1_in ? reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset) : nullptr;
          #pragma unroll
          for (int i = 0; i < WIDTH / 8; i++) {
            const int idx = i * 4;
            const int col = i * 8 + lane_col;
            const int h2_idx = col >> 1;
            if (row0_in) {
              row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
            }
            if (row1_in) {
              row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
            }
          }
        } else {
          #pragma unroll
          for (int i = 0; i < WIDTH / 8; i++) {
            const int idx = i * 4;
            const int col = n_offset + i * 8 + lane_col;
            if (col < N) {
              const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
              const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
              if (row0_in) {
                if (col + 1 < N) {
                  reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = h2_row0;
                } else {
                  C_ptr[row0 * Cs0 + col] = __low2half(h2_row0);
                }
              }
              if (row1_in) {
                if (col + 1 < N) {
                  reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = h2_row1;
                } else {
                  C_ptr[row1 * Cs0 + col] = __low2half(h2_row1);
                }
              }
            }
          }
        }
      } else {
        const bool row0_in = row0 < M;
        const bool row1_in = row1 < M;
        #pragma unroll
        for (int i = 0; i < WIDTH / 8; i++) {
          const int idx = i * 4;
          const int col = n_offset + i * 8 + lane_col;
          if (col < N) {
            const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
            const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
            const half h00 = __low2half(h2_row0);
            const half h01 = __high2half(h2_row0);
            const half h10 = __low2half(h2_row1);
            const half h11 = __high2half(h2_row1);
            if (row0_in) {
              C_ptr[row0 * Cs0 + col * Cs1] = h00;
              if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
            }
            if (row1_in) {
              C_ptr[row1 * Cs0 + col * Cs1] = h10;
              if (col + 1 < N) C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
            }
          }
        }
      }
    }
  } else {
      // Original 128-wide logic
      // We load/store in 128-column halves so tcgen05.ld has a fixed shape.
      constexpr int HALF_N = 128;
      const int halves = BLOCK_N / HALF_N;

      for (int m = 0; m < m_iters; ++m) {
        for (int half_idx = 0; half_idx < halves; ++half_idx) {
          float tmp[HALF_N / 2];
          const int col_base = half_idx * HALF_N;
          // TMEM coordinates are relative to the CTA's output tile.
          tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, col_base);
          asm volatile("tcgen05.wait::ld.sync.aligned;");

        const int row0 = warp_row_base + m * 16 + lane_row;
        const int row1 = row0 + 8;

        if (contiguous) {
          if (full_tile) {
            half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base);
            half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset + col_base);
            #pragma unroll
            for (int i = 0; i < HALF_N / 8; i++) {
              const int idx = i * 4;
              const int col = i * 8 + lane_col;
              const int h2_idx = col >> 1;
              row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
              row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
            }
            continue;
          }

          const bool row0_in = row0 < M;
          const bool row1_in = row1 < M;
          if (full_n) {
            half2* row0_ptr = row0_in ? reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base) : nullptr;
            half2* row1_ptr = row1_in ? reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset + col_base) : nullptr;
            #pragma unroll
            for (int i = 0; i < HALF_N / 8; i++) {
              const int idx = i * 4;
              const int col = i * 8 + lane_col;
              const int h2_idx = col >> 1;
              if (row0_in) {
                row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
              }
              if (row1_in) {
                row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
              }
            }
          } else {
            #pragma unroll
            for (int i = 0; i < HALF_N / 8; i++) {
              const int idx = i * 4;
              const int col = n_offset + col_base + i * 8 + lane_col;
              if (col < N) {
                const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
                const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
                if (row0_in) {
                  if (col + 1 < N) {
                    reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + col)[0] = h2_row0;
                  } else {
                    C_ptr[row0 * Cs0 + col] = __low2half(h2_row0);
                  }
                }
                if (row1_in) {
                  if (col + 1 < N) {
                    reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + col)[0] = h2_row1;
                  } else {
                    C_ptr[row1 * Cs0 + col] = __low2half(h2_row1);
                  }
                }
              }
            }
          }
        } else {
          const bool row0_in = row0 < M;
          const bool row1_in = row1 < M;
          #pragma unroll
          for (int i = 0; i < HALF_N / 8; i++) {
            const int idx = i * 4;
            const int col = n_offset + col_base + i * 8 + lane_col;
            if (col < N) {
              const half2 h2_row0 = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
              const half2 h2_row1 = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
              const half h00 = __low2half(h2_row0);
              const half h01 = __high2half(h2_row0);
              const half h10 = __low2half(h2_row1);
              const half h11 = __high2half(h2_row1);
              if (row0_in) {
                C_ptr[row0 * Cs0 + col * Cs1] = h00;
                if (col + 1 < N) C_ptr[row0 * Cs0 + (col + 1) * Cs1] = h01;
              }
              if (row1_in) {
                C_ptr[row1 * Cs0 + col * Cs1] = h10;
                if (col + 1 < N) C_ptr[row1 * Cs0 + (col + 1) * Cs1] = h11;
              }
            }
          }
        }
        }
      }
  }
}

template <int BLOCK_M, int BLOCK_N>
__device__ __forceinline__ void epilogue_store_fulltile_contiguous(
    const ProblemInfo& prob,
    int m_offset,
    int n_offset,
    int tid,
    int warp_id,
    int lane_id
) {
  if (tid >= BLOCK_M) return;

  half* C_ptr = prob.C_ptr;
  const int64_t Cs0 = prob.Cs0;

  const int warp_row_base = m_offset + warp_id * 32;

  // Lane mapping: each lane owns two columns (half2) and one of 8 rows.
  const int lane_row = lane_id >> 2;
  const int lane_col = (lane_id & 3) * 2;

  if constexpr (BLOCK_N == 64) {
    constexpr int WIDTH = 64;
    #pragma unroll
    for (int m = 0; m < 2; ++m) {
      const int row0 = warp_row_base + m * 16 + lane_row;
      const int row1 = row0 + 8;
      float tmp[32];
      tcgen05_ld_16x256b_x8(tmp, warp_id * 32 + m * 16, 0);
      asm volatile("tcgen05.wait::ld.sync.aligned;");

      half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset);
      half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset);
      #pragma unroll
      for (int i = 0; i < WIDTH / 8; i++) {
        const int idx = i * 4;
        const int col = i * 8 + lane_col;
        const int h2_idx = col >> 1;
        row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
        row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
      }
    }
  } else {
    constexpr int HALF_N = 128;
    const int halves = BLOCK_N / HALF_N;

    #pragma unroll
    for (int m = 0; m < 2; ++m) {
      const int row0 = warp_row_base + m * 16 + lane_row;
      const int row1 = row0 + 8;

      #pragma unroll
      for (int half_idx = 0; half_idx < halves; ++half_idx) {
        float tmp[HALF_N / 2];
        const int col_base = half_idx * HALF_N;
        tcgen05_ld_16x256b_x16(tmp, warp_id * 32 + m * 16, col_base);
        asm volatile("tcgen05.wait::ld.sync.aligned;\\n");

        half2* row0_ptr = reinterpret_cast<half2*>(C_ptr + row0 * Cs0 + n_offset + col_base);
        half2* row1_ptr = reinterpret_cast<half2*>(C_ptr + row1 * Cs0 + n_offset + col_base);
        #pragma unroll
        for (int i = 0; i < HALF_N / 8; i++) {
          const int idx = i * 4;
          const int col = i * 8 + lane_col;
          const int h2_idx = col >> 1;
          row0_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 0], tmp[idx + 1]));
          row1_ptr[h2_idx] = __float22half2_rn(make_float2(tmp[idx + 2], tmp[idx + 3]));
        }
      }
    }
  }
}

template <bool PERSISTENT, int BLOCK_M, int BLOCK_N, int NS, int CLUSTER_SIZE = 1>
__global__ __launch_bounds__(TMA_NUM_WARPS * WARP_SIZE)
void grouped_gemm_kernel_v4(
  const ProblemInfo* __restrict__ global_probs,
  const WorkItem* __restrict__ work_items,
  int num_items
) {
  constexpr int TMA_A_SMEM_BYTES = BLOCK_M * (TMA_BLOCK_K / 2);
  constexpr int TMA_B_SMEM_BYTES = BLOCK_N * (TMA_BLOCK_K / 2);
  constexpr int TMA_SFA_SMEM_BYTES = MMA_M * (TMA_BLOCK_K / 16);
  // Always allocate at least 128-wide equivalent for SFB to match alignment
  constexpr int SFB_WIDTH = (BLOCK_N < 128) ? 128 : BLOCK_N;
  constexpr int TMA_SFB_SMEM_BYTES = SFB_WIDTH * (TMA_BLOCK_K / 16);
  constexpr int STAGE_SIZE = TMA_A_SMEM_BYTES + TMA_B_SMEM_BYTES + TMA_SFA_SMEM_BYTES + TMA_SFB_SMEM_BYTES;
  const int tid = threadIdx.x;
  const int lane_id = tid % WARP_SIZE;
  const int warp_id = tid / WARP_SIZE;

  uint32_t cta_rank = 0;
  if constexpr (CLUSTER_SIZE > 1) {
    cta_rank = get_cluster_ctarank();
  }

  // Shared memory is used as a multi-stage ring buffer.
  // Per stage: [A tile][B tile][SFA][SFB]. After all stages we place mbarriers.
  extern __shared__ __align__(1024) char smem_ptr[];
  const int smem_base = static_cast<int>(__cvta_generic_to_shared(smem_ptr));

  constexpr int B_off = TMA_A_SMEM_BYTES;
  constexpr int SFA_off = B_off + TMA_B_SMEM_BYTES;
  constexpr int SFB_off = SFA_off + TMA_SFA_SMEM_BYTES;

  constexpr int MBAR_ARRIVALS = (CLUSTER_SIZE > 1) ? 3 : 2;
  constexpr int MBAR_SETS = (PERSISTENT && CLUSTER_SIZE > 1) ? 2 : 1;
  constexpr int MBAR_SET_BYTES = mbar_bytes(NS, MBAR_ARRIVALS);
  const int mbar_base = smem_base + STAGE_SIZE * NS;

  // TMEM allocation is in columns. We need 2 columns per output column because
  // accumulators are fp32.
  constexpr int TMEM_COLS = BLOCK_N * 2;
  constexpr int SFA_tmem = BLOCK_N;
  constexpr int SFB_tmem = SFA_tmem + 4 * (TMA_BLOCK_K / MMA_K);
  
  constexpr uint32_t idesc = (1U << 7U) | (1U << 10U)
                           | ((uint32_t)BLOCK_N >> 3U << 17U)
                           | ((uint32_t)MMA_M >> 7U << 27U);

  if (warp_id == 0) {
    asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem_base), "r"(TMEM_COLS));
  }
  __syncthreads();

  __shared__ int shared_work_idx;

  int work_idx = blockIdx.x;
  int work_epoch = 0;

  if (work_idx < num_items) {
    if (tid == 0) {
      for (int set = 0; set < MBAR_SETS; ++set) {
        const int set_base = mbar_base + set * MBAR_SET_BYTES;
        for (int i = 0; i < NS; ++i) {
          // mbarrier[stage]: TMA completion barrier.
          // Two producer warps arrive (A/SFA and B/SFB).
          mbarrier_init(set_base + i * 8, 2);

          // mbarrier[NS+stage]: stage reuse barrier.
          // The MMA warp commits once per stage.
          mbarrier_init(set_base + (NS + i) * 8, 1);

          if constexpr (CLUSTER_SIZE > 1) {
            if (cta_rank == 0) {
              // Only CTA rank 0 initializes the cluster-wide barrier used to
              // guard multicast stage reuse.
              mbarrier_init(set_base + (2*NS + i) * 8, CLUSTER_SIZE);
            }
          }
        }
      }
      asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
    }
    __syncthreads();
  }

  // Ensure all CTAs have initialized mbarriers before any multicast TMA.
  if constexpr (CLUSTER_SIZE > 1) {
    cluster_sync();
  }

  while (work_idx < num_items) {
    const WorkItem& work = work_items[work_idx];
    const ProblemInfo& prob = global_probs[work.problem_idx];
    const int mbar_work_base = mbar_base + ((work_epoch % MBAR_SETS) * MBAR_SET_BYTES);

    if (warp_id == 1 && elect_sync()) {
      int prefetch_idx;
      if constexpr (PERSISTENT) {
        prefetch_idx = work_idx + gridDim.x;
      } else {
        prefetch_idx = work_idx + 1;
      }
      if (prefetch_idx < num_items) {
        const ProblemInfo* next_prob = &global_probs[work_items[prefetch_idx].problem_idx];
        asm volatile("prefetch.tensormap [%0];" :: "l"(&next_prob->A_tmap) : "memory");
        if constexpr (BLOCK_N == 256) {
          asm volatile("prefetch.tensormap [%0];" :: "l"(&next_prob->B_tmap_256) : "memory");
        } else {
          asm volatile("prefetch.tensormap [%0];" :: "l"(&next_prob->B_tmap) : "memory");
        }
      }
    }

    const int m_offset = work.tile_m * BLOCK_M;
    const int n_offset = work.tile_n * BLOCK_N;
    const int K = prob.K;
    const int num_k_iters = K / TMA_BLOCK_K;

    if ((warp_id == TMA_WARP || warp_id == TMA_WARP_B) && elect_sync()) {
      constexpr uint64_t cache_A = EVICT_LAST;
      constexpr uint64_t cache_B = EVICT_FIRST;

        const bool do_A = (warp_id == TMA_WARP);
        const bool do_B = (warp_id == TMA_WARP_B);

        auto issue_tma = [&](int k_iter, int stage) {
          const int mbar_addr = mbar_work_base + stage * 8;
          const int stage_base = smem_base + stage * STAGE_SIZE;
          const int off_k = k_iter * TMA_BLOCK_K;

          // Program expect_tx before issuing any TMA that completes to this
          // barrier. A completion arriving before expect_tx is set can leave the
          // consumer stuck in mbarrier_wait.
          const int expect_bytes = tma_expected_tx_bytes(
            do_A,
            TMA_A_SMEM_BYTES,
            TMA_B_SMEM_BYTES,
            TMA_SFA_SMEM_BYTES,
            TMA_SFB_SMEM_BYTES
          );
          mbarrier_arrive_expect_tx(mbar_addr, expect_bytes);
          int issued_bytes = 0;

          if (do_A) {
            if constexpr (CLUSTER_SIZE > 1) {
              // Cluster path: CTA rank 0 multicasts A to the whole cluster.
              // dst and mbarrier are passed as shared::cluster addresses.
              if (cta_rank == 0) {
                uint16_t mc = (1 << CLUSTER_SIZE) - 1;
                int cluster_dst = mapa_cta_to_cluster(stage_base, 0);
                int cluster_mbar = mapa_cta_to_cluster(mbar_addr, 0);
                tma_3d_gmem2smem_multicast(cluster_dst, &prob.A_tmap, 0, m_offset, off_k / 256, cluster_mbar, mc);
              }
            } else {
              tma_3d_gmem2smem<1>(stage_base, &prob.A_tmap, 0, m_offset, off_k / 256, mbar_addr, cache_A);
            }
            issued_bytes += TMA_A_SMEM_BYTES;

            // SFA scale blocks are indexed by (m_tile, k_blk) and stored as
            // 512B blocks (matching tcgen05_cp_nvfp4 granularity).
            const int rest_k = K / 16 / 4;
            const int k_blk = off_k / (16 * 4);
            const char* SFA_src = prob.SFA_ptr + ((m_offset / 128) * rest_k + k_blk) * 512;
            tma_gmem2smem(stage_base + SFA_off, SFA_src, TMA_SFA_SMEM_BYTES, mbar_addr, cache_A);
            issued_bytes += TMA_SFA_SMEM_BYTES;
          } else if (do_B) {
            if constexpr (BLOCK_N == 256) {
              tma_3d_gmem2smem<1>(stage_base + B_off, &prob.B_tmap_256, 0, n_offset, off_k / 256, mbar_addr, cache_B);
            } else {
              tma_3d_gmem2smem<1>(stage_base + B_off, &prob.B_tmap, 0, n_offset, off_k / 256, mbar_addr, cache_B);
            }
            issued_bytes += TMA_B_SMEM_BYTES;

            // SFB scale blocks are indexed by (n_tile, k_blk).
            const int rest_k = K / 16 / 4;
            const int k_blk = off_k / (16 * 4);
            if constexpr (BLOCK_N == 256) {
              constexpr int SFB_HALF_BYTES = 128 * (TMA_BLOCK_K / 16);
              const char* SFB_src0 = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;
              const char* SFB_src1 = prob.SFB_ptr + (((n_offset / 128) + 1) * rest_k + k_blk) * 512;
              tma_gmem2smem(stage_base + SFB_off, SFB_src0, SFB_HALF_BYTES, mbar_addr, cache_B);
              tma_gmem2smem(stage_base + SFB_off + SFB_HALF_BYTES, SFB_src1, SFB_HALF_BYTES, mbar_addr, cache_B);
              issued_bytes += 2 * SFB_HALF_BYTES;
            } else {
              // For BLOCK_N=128 or 64, we load a single 128-wide SFB block.
              // Logic relies on SFB_ptr being 128-aligned/blocked.
              const char* SFB_src = prob.SFB_ptr + ((n_offset / 128) * rest_k + k_blk) * 512;
              tma_gmem2smem(stage_base + SFB_off, SFB_src, TMA_SFB_SMEM_BYTES, mbar_addr, cache_B);
              issued_bytes += TMA_SFB_SMEM_BYTES;
            }
          }

          if (issued_bytes != expect_bytes) {
            asm volatile("trap;");
          }
        };

      for (int k_iter = 0; k_iter < NS && k_iter < num_k_iters; k_iter++) {
        issue_tma(k_iter, k_iter);
      }

      int stage = 0;
      int mma_phase = 0;
      for (int k_iter = NS; k_iter < num_k_iters; k_iter++) {
        if constexpr (CLUSTER_SIZE > 1) {
            if (do_A && cta_rank == 0) {
              // A is shared across the cluster via multicast. Before reusing a
              // ring-buffer stage for the next multicast, rank 0 must wait for
              // all CTAs to finish consuming the current stage.
               mbarrier_wait_reuse(mbar_work_base + (2*NS + stage) * 8, mma_phase);
             } else {
               mbarrier_wait_reuse(mbar_work_base + (NS + stage) * 8, mma_phase);
             }
         } else {
           mbarrier_wait_reuse(mbar_work_base + (NS + stage) * 8, mma_phase);
         }
        issue_tma(k_iter, stage);
        stage++;
        if (stage == NS) {
          stage = 0;
          mma_phase ^= 1;
        }
      }
    }

    else if (warp_id == MMA_WARP && elect_sync()) {
      auto make_desc_AB = [](int addr) -> uint64_t {
        const int SBO = 8 * 128;
        // Descriptor encoding is coupled to the shared-memory swizzle and the
        // tcgen05 operand layout. SBO matches 128B swizzle.
        return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL) | (2ULL << 61ULL);
      };
      auto make_desc_SF = [](int addr) -> uint64_t {
        // Scale-factor loads use a different stride (16B) but the same address
        // encoding (16B units).
        const int SBO = 8 * 16;
        return desc_encode(addr) | (desc_encode(SBO) << 32ULL) | (1ULL << 46ULL);
      };

      int stage = 0;
      int tma_phase = 0;
      for (int k_iter = 0; k_iter < num_k_iters; k_iter++) {
        mbarrier_wait_tma(mbar_work_base + stage * 8, tma_phase);

        const int stage_base = smem_base + stage * STAGE_SIZE;

        const uint64_t SF_desc = make_desc_SF(0);
        const uint64_t SFA_desc = SF_desc + ((uint64_t)(stage_base + SFA_off) >> 4ULL);
        const uint64_t SFB_desc = SF_desc + ((uint64_t)(stage_base + SFB_off) >> 4ULL);

        // Copy scale factors from shared memory into TMEM.
        #pragma unroll
        for (int k = 0; k < TMA_BLOCK_K / MMA_K; k++) {
          tcgen05_cp_nvfp4(SFA_tmem + k * 4, SFA_desc + (uint64_t)k * (512ULL >> 4ULL));
          if constexpr (BLOCK_N == 256) {
            constexpr uint64_t SFB_HALF_DESC = (uint64_t)(128 * (TMA_BLOCK_K / 16)) >> 4ULL;
            tcgen05_cp_nvfp4(SFB_tmem + k * 8, SFB_desc + (uint64_t)k * (512ULL >> 4ULL));
            tcgen05_cp_nvfp4(SFB_tmem + k * 8 + 4, SFB_desc + SFB_HALF_DESC + (uint64_t)k * (512ULL >> 4ULL));
          } else {
            tcgen05_cp_nvfp4(SFB_tmem + k * 4, SFB_desc + (uint64_t)k * (512ULL >> 4ULL));
          }
        }

        // MMA loop over the 256-wide K tile in 64-wide chunks.
        #pragma unroll
        for (int k = 0; k < TMA_BLOCK_K / MMA_K; k++) {
          uint64_t a_desc = make_desc_AB(stage_base + k * 32);
          uint64_t b_desc = make_desc_AB(stage_base + B_off + k * 32);

          // scale_A_tmem = SFA_tmem + k * 4;
          const int scale_A_tmem = SFA_tmem + k * 4;
          int scale_B_tmem;
          if constexpr (BLOCK_N == 256) {
            scale_B_tmem = SFB_tmem + k * 8;
          } else if constexpr (BLOCK_N == 128) {
            scale_B_tmem = SFB_tmem + k * 4;
          } else {
            // BLOCK_N=64: use (tile_n % 2) * 2 to select the 64-wide slice of the 128-wide SFB
            scale_B_tmem = SFB_tmem + k * 4 + (work.tile_n % 2) * 2;
          }

          // First MMA uses D=0, subsequent MMAs accumulate.
          const int enable_input_d = (k_iter == 0 && k == 0) ? 0 : 1;
          tcgen05_mma_nvfp4(a_desc, b_desc, idesc, scale_A_tmem, scale_B_tmem, enable_input_d);
        }

        tcgen05_commit(mbar_work_base + (NS + stage) * 8);
        if constexpr (CLUSTER_SIZE > 1) {
          int cm = mapa_cta_to_cluster(mbar_work_base + (2*NS + stage) * 8, 0);
          mbarrier_arrive_cluster(cm);
        }

        stage++;
        if (stage == NS) {
          stage = 0;
          tma_phase ^= 1;
        }
      }

      const int last_stage = (num_k_iters - 1) % NS;
      const int last_phase = ((num_k_iters - 1) / NS) % 2;
      mbarrier_wait_reuse(mbar_work_base + (NS + last_stage) * 8, last_phase);
    }

    __syncthreads();
    asm volatile("tcgen05.fence::after_thread_sync;" ::: "memory");

    const bool short_k = (K <= 2048);
    const bool full_tile = (m_offset + BLOCK_M <= prob.M) && (n_offset + BLOCK_N <= prob.N);
    if (short_k && full_tile && prob.Cs1 == 1) {
      epilogue_store_fulltile_contiguous<BLOCK_M, BLOCK_N>(prob, m_offset, n_offset, tid, warp_id, lane_id);
    } else {
      const bool adaptive_m_epilogue = (prob.M <= LOW_M_THRESHOLD) || short_k;
      if (adaptive_m_epilogue) {
        epilogue_store<BLOCK_M, BLOCK_N, true>(prob, m_offset, n_offset, tid, warp_id, lane_id);
      } else {
        epilogue_store<BLOCK_M, BLOCK_N, false>(prob, m_offset, n_offset, tid, warp_id, lane_id);
      }
    }

    // In cluster mode we must synchronize across CTAs before re-initializing
    // mbarriers; __syncthreads is CTA-local and does not order cluster-wide
    // mbarrier arrivals.
    if constexpr (PERSISTENT) {
      if (warp_id == TMA_WARP && elect_sync()) {
        // Static grid-stride work distribution avoids global atomics and
        // smooths the tail when num_items slightly exceeds one wave.
        shared_work_idx = work_idx + gridDim.x;

        if constexpr (CLUSTER_SIZE == 1) {
          if (shared_work_idx < num_items) {
            for (int i = 0; i < NS; ++i) {
              mbarrier_init(mbar_base + i * 8, 2);
              mbarrier_init(mbar_base + (NS + i) * 8, 1);
              if constexpr (CLUSTER_SIZE > 1) {
                if (cta_rank == 0) {
                  mbarrier_init(mbar_base + (2*NS + i) * 8, CLUSTER_SIZE);
                }
              }
            }
            asm volatile("fence.mbarrier_init.release.cluster;" ::: "memory");
          }
        }
      }
    }

    if constexpr (PERSISTENT) {
      __syncthreads();
    }

    if constexpr (PERSISTENT) {
      work_idx = shared_work_idx;
      if constexpr (CLUSTER_SIZE > 1) {
        work_epoch++;
      }
    } else {
      break;
    }
  }

  if (warp_id == 0) {
    tcgen05_dealloc_cols_cta1(0, TMEM_COLS);
  }
}

static inline uint64_t max_tmap_rows_u4(const at::Tensor& t, uint64_t global_width) {
  TORCH_CHECK(global_width >= 256 && (global_width % 256) == 0, "K must be multiple of 256");

  const uint64_t logical_height = (uint64_t)t.size(0);
  if (t.dim() < 2) {
    return logical_height;
  }

  const int64_t elem_size = (int64_t)t.element_size();
  const int64_t stride0 = t.stride(0);
  const int64_t row_bytes = (int64_t)(global_width / 2);
  if (elem_size <= 0 || stride0 <= 0 || (row_bytes % elem_size) != 0) {
    return logical_height;
  }

  const int64_t row_elems = row_bytes / elem_size;
  if (stride0 < row_elems) {
    return logical_height;
  }

  const int64_t storage_nbytes = (int64_t)t.storage().nbytes();
  const int64_t storage_offset_bytes = (int64_t)t.storage_offset() * elem_size;
  if (storage_offset_bytes > storage_nbytes) {
    return logical_height;
  }

  const int64_t available_bytes = storage_nbytes - storage_offset_bytes;
  if (available_bytes < row_bytes) {
    return logical_height;
  }

  const int64_t stride0_bytes = stride0 * elem_size;
  const uint64_t max_rows = (uint64_t)(1 + (available_bytes - row_bytes) / stride0_bytes);
  return std::max(logical_height, max_rows);
}

void init_AB_tmap_u4(
  CUtensorMap *tmap,
  const void *ptr,
  uint64_t global_height, uint64_t global_width,
  uint32_t shared_height, uint32_t shared_width,
  uint64_t max_safe_height
) {
  TORCH_CHECK(ptr != nullptr, "ptr is null");
  TORCH_CHECK(((uintptr_t)ptr % 16) == 0, "ptr must be 16-byte aligned");
  TORCH_CHECK(global_width >= 256 && (global_width % 256) == 0, "K must be multiple of 256");
  TORCH_CHECK(shared_width == 256, "shared_width must be 256");

  const uint64_t aligned_height = (global_height + 127ULL) & ~127ULL;
  if (max_safe_height >= aligned_height) {
    global_height = aligned_height;
  }

  constexpr uint32_t rank = 3;
  uint64_t globalDim[rank]       = {256, global_height, global_width / 256};
  uint64_t globalStrides[rank-1] = {global_width / 2, 128};
  uint32_t boxDim[rank]          = {256, shared_height, 1};
  uint32_t elementStrides[rank]  = {1, 1, 1};

  // Swizzle must match the shared-memory layout expected by tcgen05.
  constexpr CUtensorMapSwizzle swizzle = CU_TENSOR_MAP_SWIZZLE_128B;

  // Cache cuTensorMap templates by shape.
  struct ShapeKey { uint64_t gh, gw; uint32_t sh, sw; };
  struct ShapeHash {
    size_t operator()(const ShapeKey& k) const noexcept {
      uint64_t h = k.gh;
      h ^= (k.gw + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2));
      h ^= ((uint64_t)k.sh << 32) ^ (uint64_t)k.sw;
      return (size_t)h;
    }
  };
  struct ShapeEq {
    bool operator()(const ShapeKey& a, const ShapeKey& b) const noexcept {
      return a.gh == b.gh && a.gw == b.gw && a.sh == b.sh && a.sw == b.sw;
    }
  };
  struct PtrKey { uint64_t gh, gw; uint32_t sh, sw; const void* ptr; };
  struct PtrHash {
    size_t operator()(const PtrKey& k) const noexcept {
      uint64_t h = k.gh;
      h ^= (k.gw + 0x9e3779b97f4a7c15ULL + (h << 6) + (h >> 2));
      h ^= ((uint64_t)k.sh << 32) ^ (uint64_t)k.sw;
      h ^= ((uint64_t)k.ptr >> 4);
      return (size_t)h;
    }
  };
  struct PtrEq {
    bool operator()(const PtrKey& a, const PtrKey& b) const noexcept {
      return a.gh == b.gh && a.gw == b.gw && a.sh == b.sh && a.sw == b.sw && a.ptr == b.ptr;
    }
  };

  static thread_local std::unordered_map<ShapeKey, CUtensorMap, ShapeHash, ShapeEq> tmpl_cache;
  static thread_local std::unordered_map<PtrKey, CUtensorMap, PtrHash, PtrEq> ptr_cache;

  PtrKey pkey{global_height, global_width, shared_height, shared_width, ptr};
  auto pit = ptr_cache.find(pkey);
  if (pit != ptr_cache.end()) { *tmap = pit->second; return; }

  ShapeKey skey{global_height, global_width, shared_height, shared_width};
  auto sit = tmpl_cache.find(skey);
  if (sit == tmpl_cache.end()) {
    CUtensorMap tmp;
    auto err = cuTensorMapEncodeTiled(
      &tmp, CU_TENSOR_MAP_DATA_TYPE_16U4_ALIGN8B,
      rank, (void*)ptr, globalDim, globalStrides, boxDim, elementStrides,
      CU_TENSOR_MAP_INTERLEAVE_NONE, swizzle,
      CU_TENSOR_MAP_L2_PROMOTION_NONE, CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE
    );
    TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapEncodeTiled failed");
    sit = tmpl_cache.emplace(skey, tmp).first;
  }

  CUtensorMap tmp = sit->second;
  auto err = cuTensorMapReplaceAddress(&tmp, (void*)ptr);
  TORCH_CHECK(err == CUDA_SUCCESS, "cuTensorMapReplaceAddress failed");
  ptr_cache.emplace(pkey, tmp);
  *tmap = tmp;
}

std::vector<at::Tensor> group_gemm(
    std::vector<at::Tensor> A_list,
    std::vector<at::Tensor> B_list,
    std::vector<at::Tensor> C_list,
    std::vector<at::Tensor> sfa_list,
    std::vector<at::Tensor> sfb_list,
    at::Tensor sizes_cpu
) {
    int64_t G = A_list.size();
    auto dev = A_list[0].device();
    c10::cuda::CUDAGuard device_guard(dev);
    auto sizes_accessor = sizes_cpu.accessor<int64_t, 2>();

    constexpr int sm_count = 148;  // B200 / SM100

    static bool attrs_set = false;
    if (!attrs_set) {
      constexpr int SMEM_128_HI_K256 = smem_bytes(128, NS_DEEP_HI, 2);
      constexpr int SMEM_256_HI_K256 = smem_bytes(256, NS_WIDE_HI, 2);
      constexpr int SMEM_256_CLUSTER_HI_K256 = smem_bytes(256, NS_WIDE_HI, 3, 2);
      constexpr int SMEM_64_HI_K256 = smem_bytes(64, NS_64_HI, 2);
      constexpr int SMEM_64_LO_K256 = smem_bytes(64, NS_64_LO, 2);

      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_HI_K256));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_128_HI_K256));

      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_HI_K256));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_HI_K256));

      constexpr int CL = CLUSTER_SIZE_256;
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_HI_K256));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_256_CLUSTER_HI_K256));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL>), cudaFuncAttributeNonPortableClusterSizeAllowed, 1));

      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 64, NS_64_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_64_HI_K256));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 64, NS_64_HI>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_64_HI_K256));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<true, 128, 64, NS_64_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_64_LO_K256));
      CUDA_CHECK(cudaFuncSetAttribute((grouped_gemm_kernel_v4<false, 128, 64, NS_64_LO>), cudaFuncAttributeMaxDynamicSharedMemorySize, SMEM_64_LO_K256));

      attrs_set = true;
    }

    static thread_local bool occ_set = false;
    static thread_local int occ_64_hi = 0, occ_64_lo = 0;
    static thread_local int occ_128_hi = 0;
    static thread_local int occ_256_hi = 0;
    static thread_local int occ_256_cluster_hi = 0;
    if (!occ_set) {
      constexpr int SMEM_128_HI_K256 = smem_bytes(128, NS_DEEP_HI, 2);
      constexpr int SMEM_256_HI_K256 = smem_bytes(256, NS_WIDE_HI, 2);
      constexpr int SMEM_256_CLUSTER_HI_K256 = smem_bytes(256, NS_WIDE_HI, 3, 2);
      constexpr int SMEM_64_HI_K256 = smem_bytes(64, NS_64_HI, 2);
      constexpr int SMEM_64_LO_K256 = smem_bytes(64, NS_64_LO, 2);

      constexpr int THREADS = TMA_NUM_WARPS * WARP_SIZE;
      CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
          &occ_64_hi, grouped_gemm_kernel_v4<false, 128, 64, NS_64_HI>, THREADS, SMEM_64_HI_K256));
      CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
          &occ_64_lo, grouped_gemm_kernel_v4<false, 128, 64, NS_64_LO>, THREADS, SMEM_64_LO_K256));
      CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
          &occ_128_hi, grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI>, THREADS, SMEM_128_HI_K256));
      CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
          &occ_256_hi, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI>, THREADS, SMEM_256_HI_K256));
      CUDA_CHECK(cudaOccupancyMaxActiveBlocksPerMultiprocessor(
          &occ_256_cluster_hi, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CLUSTER_SIZE_256>, THREADS, SMEM_256_CLUSTER_HI_K256));

      TORCH_CHECK(occ_128_hi > 0 && occ_256_hi > 0 && occ_256_cluster_hi > 0 && occ_64_hi > 0 && occ_64_lo > 0,
                  "occupancy query returned zero blocks/SM");
      occ_set = true;
    }

    auto options = at::TensorOptions().dtype(at::kByte).device(dev);
    auto host_pinned_options = at::TensorOptions().dtype(at::kByte).device(at::kCPU).pinned_memory(true);

    const int64_t probs_bytes = (int64_t)(G * sizeof(ProblemInfo));

    static thread_local at::Tensor h_probs_cache;
    static thread_local std::vector<ProblemInfo> h_probs_pageable;
    static thread_local at::Tensor h_work_64_cache;
    static thread_local at::Tensor h_work_128_cache;
    static thread_local at::Tensor h_work_256_cache;
    static thread_local at::Tensor h_work_256_cluster_cache;
    if ((int64_t)h_probs_pageable.size() < G) {
      h_probs_pageable.resize((size_t)G);
    }

    ProblemInfo* problem_infos = h_probs_pageable.data();

    std::vector<uint8_t> active(G, 0);
    std::vector<uint8_t> algo_kind(G, ALGO_128);
    std::vector<uint8_t> use_cluster_256(G, 0);
    std::vector<int> num_tiles_m(G, 0);
    std::vector<int> num_tiles_n(G, 0);
    std::vector<int64_t> Ms(G, 0), Ns(G, 0), Ks(G, 0);

    for (int64_t i = 0; i < G; i++) {
        const int64_t M = sizes_accessor[i][0], N = sizes_accessor[i][1], K = sizes_accessor[i][2];
        Ms[(size_t)i] = M; Ns[(size_t)i] = N; Ks[(size_t)i] = K;
        if (A_list[i].stride(1) != 1 || B_list[i].stride(1) != 1) {
          continue;
        }
        active[(size_t)i] = 1;
        const bool n_aligned_256 = ((N & 255) == 0);
        const bool wide_n = (N >= 6144);

        // Dispatch heuristic.
        const bool is_64 = (wide_n && K >= 4096) || (wide_n && !n_aligned_256 && K >= 2048);
        const bool is_256 = !is_64 && (N >= 6144) && (K >= 2048) && n_aligned_256;

        const int tiles_n_256 = ceil_div((int)N, 256);
        const bool enough_tiles_for_cluster = (tiles_n_256 >= 8);
        const bool is_256_cluster = is_256 && (K >= 4096) && enough_tiles_for_cluster;

        uint8_t algo = ALGO_128;
        if (is_64) algo = ALGO_64;
        else if (is_256) algo = ALGO_256;
        algo_kind[(size_t)i] = algo;
        use_cluster_256[(size_t)i] = is_256_cluster;

        const int block_n = block_n_for_algo(algo);
        num_tiles_m[(size_t)i] = ceil_div((int)M, MMA_M);
        num_tiles_n[(size_t)i] = ceil_div((int)N, block_n);
    }

    uint64_t work_hash = 1469598103934665603ULL;
    work_hash = hash_combine_u64(work_hash, (uint64_t)TMA_BLOCK_K);
    for (int64_t i = 0; i < G; i++) {
        const bool is_active = (active[(size_t)i] != 0);
        if (!is_active) {
          work_hash = hash_combine_u64(work_hash, 0);
          continue;
        }
        const int tiles_m = num_tiles_m[(size_t)i];
        const int tiles_n = num_tiles_n[(size_t)i];
        const uint8_t algo = algo_kind[(size_t)i];
        const bool is_256_cluster = (use_cluster_256[(size_t)i] != 0);
        
        work_hash = hash_combine_u64(work_hash, (uint64_t)tiles_m);
        work_hash = hash_combine_u64(work_hash, (uint64_t)tiles_n);
        work_hash = hash_combine_u64(work_hash, (uint64_t)algo);
        work_hash = hash_combine_u64(work_hash, (uint64_t)is_256_cluster);
    }

    int64_t num_items_64_i64 = 0;
    int64_t num_items_128_i64 = 0;
    int64_t num_items_256_i64 = 0;
    int64_t num_items_256_cluster_i64 = 0;
    for (int64_t i = 0; i < G; i++) {
      if (!active[(size_t)i]) continue;
      const int64_t tiles_m = num_tiles_m[(size_t)i];
      const int64_t tiles_n = num_tiles_n[(size_t)i];
      const uint8_t algo = algo_kind[(size_t)i];

      if (algo == ALGO_64) {
        num_items_64_i64 += tiles_m * tiles_n;
      } else if (algo == ALGO_128) {
        num_items_128_i64 += tiles_m * tiles_n;
      } else {
        if (use_cluster_256[(size_t)i]) {
          const int64_t tiles_n_padded = ((tiles_n + CLUSTER_SIZE_256 - 1) / CLUSTER_SIZE_256) * CLUSTER_SIZE_256;
          num_items_256_cluster_i64 += tiles_m * tiles_n_padded;
        } else {
          num_items_256_i64 += tiles_m * tiles_n;
        }
      }
    }

    TORCH_CHECK(num_items_64_i64 <= std::numeric_limits<int>::max(), "too many ALGO_64 work items");
    TORCH_CHECK(num_items_128_i64 <= std::numeric_limits<int>::max(), "too many ALGO_128 work items");
    TORCH_CHECK(num_items_256_i64 <= std::numeric_limits<int>::max(), "too many ALGO_256 work items");
    TORCH_CHECK(num_items_256_cluster_i64 <= std::numeric_limits<int>::max(), "too many ALGO_256 cluster work items");

    const int num_items_64 = (int)num_items_64_i64;
    const int num_items_128 = (int)num_items_128_i64;
    const int num_items_256 = (int)num_items_256_i64;
    const int num_items_256_cluster = (int)num_items_256_cluster_i64;

    uint64_t probs_hash = 1469598103934665603ULL;
    probs_hash = hash_combine_u64(probs_hash, (uint64_t)TMA_BLOCK_K);
    for (int64_t i = 0; i < G; i++) {
        if (!active[(size_t)i]) continue;
        const int64_t M = Ms[(size_t)i], N = Ns[(size_t)i], K = Ks[(size_t)i];

        ProblemInfo& p = problem_infos[i];
        p.M = M; p.N = N; p.K = K;
        p.Cs0 = C_list[i].stride(0); p.Cs1 = C_list[i].stride(1);
        p.C_ptr = (half*)C_list[i].data_ptr();
        p.SFA_ptr = (const char*)sfa_list[i].data_ptr();
        p.SFB_ptr = (const char*)sfb_list[i].data_ptr();

        const uint8_t algo = algo_kind[(size_t)i];
        
        const uint64_t A_max_safe_height = max_tmap_rows_u4(A_list[i], (uint64_t)K);
        const uint64_t B_max_safe_height = max_tmap_rows_u4(B_list[i], (uint64_t)K);

        init_AB_tmap_u4(&p.A_tmap, A_list[i].data_ptr(), A_list[i].size(0), K, 128, 256, A_max_safe_height);
        const int block_n = block_n_for_algo(algo);
        const int tmap_b_height = (block_n == 64) ? 64 : 128;
        init_AB_tmap_u4(&p.B_tmap, B_list[i].data_ptr(), B_list[i].size(0), K, tmap_b_height, 256, B_max_safe_height);
        
        if (algo == ALGO_256) {
          init_AB_tmap_u4(&p.B_tmap_256, B_list[i].data_ptr(), B_list[i].size(0), K, 256, 256, B_max_safe_height);
        } else {
          p.B_tmap_256 = p.B_tmap;
        }

        probs_hash = hash_combine_u64(probs_hash, (uint64_t)i);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)M);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)N);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)K);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)A_list[i].data_ptr());
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)B_list[i].data_ptr());
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.C_ptr);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.SFA_ptr);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)(uintptr_t)p.SFB_ptr);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs0);
        probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs1);
    }

    if (num_items_64 == 0 && num_items_128 == 0 && num_items_256 == 0 && num_items_256_cluster == 0) return C_list;

    const int wave_hi_64 = sm_count * occ_64_hi;
    const int wave_lo_64 = sm_count * occ_64_lo;
    const bool use_lo_64 = (wave_lo_64 > wave_hi_64) && (num_items_64 > 2 * wave_hi_64);
    const int wave_cap_64 = use_lo_64 ? wave_lo_64 : wave_hi_64;
    const int wave_cap_128 = sm_count * occ_128_hi;
    const int wave_cap_256 = sm_count * occ_256_hi;
    const int wave_cap_256_cluster = (sm_count * occ_256_cluster_hi / CLUSTER_SIZE_256) * CLUSTER_SIZE_256;

    constexpr int LPT_MIN_WAVES = 3;
    const bool enable_lpt_64 = (num_items_64 > LPT_MIN_WAVES * wave_cap_64);
    const bool enable_lpt_128 = (num_items_128 > LPT_MIN_WAVES * wave_cap_128);
    const bool enable_lpt_256 = (num_items_256 > LPT_MIN_WAVES * wave_cap_256);
    const bool enable_lpt_256_cluster = (num_items_256_cluster > LPT_MIN_WAVES * wave_cap_256_cluster);

    static thread_local at::Tensor d_probs_cache;
    static thread_local at::Tensor d_work_cache_64;
    static thread_local at::Tensor d_work_cache_128;
    static thread_local at::Tensor d_work_cache_256;
    static thread_local at::Tensor d_work_cache_256_cluster;
    static thread_local uint64_t last_probs_hash = 0;
    static thread_local uint64_t last_work_hash = 0;
    static thread_local bool last_hash_valid = false;

    const int64_t work_bytes_64 = (int64_t)(num_items_64 * sizeof(WorkItem));
    const int64_t work_bytes_128 = (int64_t)(num_items_128 * sizeof(WorkItem));
    const int64_t work_bytes_256 = (int64_t)(num_items_256 * sizeof(WorkItem));
    const int64_t work_bytes_256_cluster = (int64_t)(num_items_256_cluster * sizeof(WorkItem));

    bool probs_realloc = false;
    bool work_realloc = false;
    if (!d_probs_cache.defined() || d_probs_cache.device() != dev || d_probs_cache.scalar_type() != at::kByte || d_probs_cache.numel() < probs_bytes) {
      d_probs_cache = at::empty({probs_bytes}, options);
      probs_realloc = true;
    }
    if (work_bytes_64 > 0 && (!d_work_cache_64.defined() || d_work_cache_64.device() != dev || d_work_cache_64.scalar_type() != at::kByte || d_work_cache_64.numel() < work_bytes_64)) {
      d_work_cache_64 = at::empty({work_bytes_64}, options);
      work_realloc = true;
    }
    if (work_bytes_128 > 0 && (!d_work_cache_128.defined() || d_work_cache_128.device() != dev || d_work_cache_128.scalar_type() != at::kByte || d_work_cache_128.numel() < work_bytes_128)) {
      d_work_cache_128 = at::empty({work_bytes_128}, options);
      work_realloc = true;
    }
    if (work_bytes_256 > 0 && (!d_work_cache_256.defined() || d_work_cache_256.device() != dev || d_work_cache_256.scalar_type() != at::kByte || d_work_cache_256.numel() < work_bytes_256)) {
      d_work_cache_256 = at::empty({work_bytes_256}, options);
      work_realloc = true;
    }
    if (work_bytes_256_cluster > 0 && (!d_work_cache_256_cluster.defined() || d_work_cache_256_cluster.device() != dev || d_work_cache_256_cluster.scalar_type() != at::kByte || d_work_cache_256_cluster.numel() < work_bytes_256_cluster)) {
      d_work_cache_256_cluster = at::empty({work_bytes_256_cluster}, options);
      work_realloc = true;
    }

    if (probs_realloc || !last_hash_valid || last_probs_hash != probs_hash) {
      const bool use_pinned_probs_copy = (probs_bytes >= (64 * 1024));
      const void* probs_src = problem_infos;
      if (use_pinned_probs_copy) {
        if (!h_probs_cache.defined() || h_probs_cache.device().type() != at::kCPU || !h_probs_cache.is_pinned() || h_probs_cache.scalar_type() != at::kByte || h_probs_cache.numel() < probs_bytes) {
          h_probs_cache = at::empty({probs_bytes}, host_pinned_options);
        }
        std::memcpy(h_probs_cache.data_ptr(), problem_infos, (size_t)probs_bytes);
        probs_src = h_probs_cache.data_ptr();
      }
      CUDA_CHECK(cudaMemcpyAsync(d_probs_cache.data_ptr(), probs_src, probs_bytes, cudaMemcpyHostToDevice));
      last_probs_hash = probs_hash;
    }
    if (work_realloc || !last_hash_valid || last_work_hash != work_hash) {
      if (work_bytes_64 > 0 && (!h_work_64_cache.defined() || h_work_64_cache.device().type() != at::kCPU || !h_work_64_cache.is_pinned() || h_work_64_cache.scalar_type() != at::kByte || h_work_64_cache.numel() < work_bytes_64)) {
        h_work_64_cache = at::empty({work_bytes_64}, host_pinned_options);
      }
      if (work_bytes_128 > 0 && (!h_work_128_cache.defined() || h_work_128_cache.device().type() != at::kCPU || !h_work_128_cache.is_pinned() || h_work_128_cache.scalar_type() != at::kByte || h_work_128_cache.numel() < work_bytes_128)) {
        h_work_128_cache = at::empty({work_bytes_128}, host_pinned_options);
      }
      if (work_bytes_256 > 0 && (!h_work_256_cache.defined() || h_work_256_cache.device().type() != at::kCPU || !h_work_256_cache.is_pinned() || h_work_256_cache.scalar_type() != at::kByte || h_work_256_cache.numel() < work_bytes_256)) {
        h_work_256_cache = at::empty({work_bytes_256}, host_pinned_options);
      }
      if (work_bytes_256_cluster > 0 && (!h_work_256_cluster_cache.defined() || h_work_256_cluster_cache.device().type() != at::kCPU || !h_work_256_cluster_cache.is_pinned() || h_work_256_cluster_cache.scalar_type() != at::kByte || h_work_256_cluster_cache.numel() < work_bytes_256_cluster)) {
        h_work_256_cluster_cache = at::empty({work_bytes_256_cluster}, host_pinned_options);
      }

      WorkItem* h_work_64 = (work_bytes_64 > 0) ? reinterpret_cast<WorkItem*>(h_work_64_cache.data_ptr()) : nullptr;
      WorkItem* h_work_128 = (work_bytes_128 > 0) ? reinterpret_cast<WorkItem*>(h_work_128_cache.data_ptr()) : nullptr;
      WorkItem* h_work_256 = (work_bytes_256 > 0) ? reinterpret_cast<WorkItem*>(h_work_256_cache.data_ptr()) : nullptr;
      WorkItem* h_work_256_cluster = (work_bytes_256_cluster > 0) ? reinterpret_cast<WorkItem*>(h_work_256_cluster_cache.data_ptr()) : nullptr;

      int out_64 = 0;
      int out_128 = 0;
      int out_256 = 0;
      int out_256_cluster = 0;

      int max_tn_64 = 0;
      for (int64_t i = 0; i < G; ++i) {
        if (active[(size_t)i] && algo_kind[(size_t)i] == ALGO_64) {
          max_tn_64 = std::max(max_tn_64, num_tiles_n[(size_t)i]);
        }
      }
      for (int tn = 0; tn < max_tn_64; ++tn) {
        for (int64_t i = 0; i < G; ++i) {
          if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_64 || tn >= num_tiles_n[(size_t)i]) continue;
          for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
            h_work_64[out_64++] = {(int)i, tm, tn};
          }
        }
      }

      int max_tn_128 = 0;
      for (int64_t i = 0; i < G; ++i) {
        if (active[(size_t)i] && algo_kind[(size_t)i] == ALGO_128) {
          max_tn_128 = std::max(max_tn_128, num_tiles_n[(size_t)i]);
        }
      }
      for (int tn = 0; tn < max_tn_128; ++tn) {
        for (int64_t i = 0; i < G; ++i) {
          if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_128 || tn >= num_tiles_n[(size_t)i]) continue;
          for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
            h_work_128[out_128++] = {(int)i, tm, tn};
          }
        }
      }

      struct ClusterRow {
        int problem_idx;
        int tile_m;
        int tiles_n;
        int pad;
      };
      std::vector<ClusterRow> cluster_rows;
      for (int64_t i = 0; i < G; ++i) {
        if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_256) continue;
        if (!use_cluster_256[(size_t)i]) {
          for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
            for (int tn = 0; tn < num_tiles_n[(size_t)i]; ++tn) {
              h_work_256[out_256++] = {(int)i, tm, tn};
            }
          }
          continue;
        }
        const int tiles_n = num_tiles_n[(size_t)i];
        const int remainder = tiles_n % CLUSTER_SIZE_256;
        const int pad = (remainder == 0) ? 0 : (CLUSTER_SIZE_256 - remainder);
        for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
          cluster_rows.push_back({(int)i, tm, tiles_n, pad});
        }
      }

      auto sort_by_volume = [&](WorkItem* items, int count) {
        if (count <= 1) return;
        std::sort(items, items + count, [&](const WorkItem& a, const WorkItem& b) {
          const int m_a = std::min(128, (int)Ms[(size_t)a.problem_idx] - a.tile_m * 128);
          const int m_b = std::min(128, (int)Ms[(size_t)b.problem_idx] - b.tile_m * 128);
          if (m_a == m_b) return Ks[(size_t)a.problem_idx] > Ks[(size_t)b.problem_idx];
          return m_a > m_b;
        });
      };

      if (enable_lpt_64) sort_by_volume(h_work_64, out_64);
      if (enable_lpt_128) sort_by_volume(h_work_128, out_128);
      if (enable_lpt_256) sort_by_volume(h_work_256, out_256);

      if (enable_lpt_256_cluster && !cluster_rows.empty()) {
        std::sort(cluster_rows.begin(), cluster_rows.end(), [&](const ClusterRow& a, const ClusterRow& b) {
          const int m_a = std::min(128, (int)Ms[(size_t)a.problem_idx] - a.tile_m * 128);
          const int m_b = std::min(128, (int)Ms[(size_t)b.problem_idx] - b.tile_m * 128);
          if (m_a == m_b) return Ks[(size_t)a.problem_idx] > Ks[(size_t)b.problem_idx];
          return m_a > m_b;
        });
      }

      for (const ClusterRow& row : cluster_rows) {
        for (int tn = 0; tn < row.tiles_n; ++tn) {
          h_work_256_cluster[out_256_cluster++] = {row.problem_idx, row.tile_m, tn};
        }
        for (int p = 0; p < row.pad; ++p) {
          h_work_256_cluster[out_256_cluster++] = {row.problem_idx, row.tile_m, 0};
        }
      }

      TORCH_CHECK(out_64 == num_items_64, "ALGO_64 work-item count mismatch");
      TORCH_CHECK(out_128 == num_items_128, "ALGO_128 work-item count mismatch");
      TORCH_CHECK(out_256 == num_items_256, "ALGO_256 work-item count mismatch");
      TORCH_CHECK(out_256_cluster == num_items_256_cluster, "ALGO_256 cluster work-item count mismatch");

      if (work_bytes_64 > 0) {
        CUDA_CHECK(cudaMemcpyAsync(d_work_cache_64.data_ptr(), h_work_64, work_bytes_64, cudaMemcpyHostToDevice));
      }
      if (work_bytes_128 > 0) {
        CUDA_CHECK(cudaMemcpyAsync(d_work_cache_128.data_ptr(), h_work_128, work_bytes_128, cudaMemcpyHostToDevice));
      }
      if (work_bytes_256 > 0) {
        CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256.data_ptr(), h_work_256, work_bytes_256, cudaMemcpyHostToDevice));
      }
      if (work_bytes_256_cluster > 0) {
        CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256_cluster.data_ptr(), h_work_256_cluster, work_bytes_256_cluster, cudaMemcpyHostToDevice));
      }

      last_work_hash = work_hash;
    }
    last_hash_valid = true;

    if (num_items_64 > 0) {
      const int wave_hi = sm_count * occ_64_hi;
      const int wave_lo = sm_count * occ_64_lo;
      const bool use_lo = (wave_lo > wave_hi) && (num_items_64 > 2 * wave_hi);

      const int ns = use_lo ? NS_64_LO : NS_64_HI;
      const int occ = use_lo ? occ_64_lo : occ_64_hi;
      const int wave_cap = sm_count * occ;

      const int SMEM_SIZE_64 = smem_bytes(64, ns, 2);

      const bool persistent_64 = (num_items_64 > wave_cap);
      const int launch_ctas_64 = persistent_64 ? wave_cap : num_items_64;

      if (persistent_64) {
        if (use_lo) {
          grouped_gemm_kernel_v4<true, 128, 64, NS_64_LO><<<launch_ctas_64, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_64>>>(
              (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_64.data_ptr(), num_items_64);
        } else {
          grouped_gemm_kernel_v4<true, 128, 64, NS_64_HI><<<launch_ctas_64, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_64>>>(
              (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_64.data_ptr(), num_items_64);
        }
      } else {
         if (use_lo) {
          grouped_gemm_kernel_v4<false, 128, 64, NS_64_LO><<<launch_ctas_64, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_64>>>(
              (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_64.data_ptr(), num_items_64);
        } else {
          grouped_gemm_kernel_v4<false, 128, 64, NS_64_HI><<<launch_ctas_64, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_64>>>(
              (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_64.data_ptr(), num_items_64);
        }
      }
    }

    if (num_items_128 > 0) {
      const int ns = NS_DEEP_HI;
      const int occ = occ_128_hi;
      const int wave_cap = sm_count * occ;

      const int SMEM_SIZE_128 = smem_bytes(128, ns, 2);

      const bool persistent_128 = (num_items_128 > wave_cap);
      const int launch_ctas_128 = persistent_128 ? wave_cap : num_items_128;

      if (persistent_128) {
        grouped_gemm_kernel_v4<true, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(
            (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);
      } else {
        grouped_gemm_kernel_v4<false, 128, 128, NS_DEEP_HI><<<launch_ctas_128, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_128>>>(
            (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_128.data_ptr(), num_items_128);
      }
    }

    if (num_items_256 > 0) {
      const int ns = NS_WIDE_HI;
      const int occ = occ_256_hi;
      const int wave_cap = sm_count * occ;

      const int SMEM_SIZE_256 = smem_bytes(256, ns, 2);

      const bool persistent_256 = (num_items_256 > wave_cap);
      const int launch_ctas_256 = persistent_256 ? wave_cap : num_items_256;

      if (persistent_256) {
        grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(
            (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);
      } else {
        grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI><<<launch_ctas_256, TMA_NUM_WARPS * WARP_SIZE, SMEM_SIZE_256>>>(
            (ProblemInfo*)d_probs_cache.data_ptr(), (WorkItem*)d_work_cache_256.data_ptr(), num_items_256);
      }
    }

    if (num_items_256_cluster > 0) {
      constexpr int CL256 = CLUSTER_SIZE_256;

      const int ns = NS_WIDE_HI;
      const int occ = occ_256_cluster_hi;
      const int wave_cap = (sm_count * occ / CL256) * CL256;

      const int SMEM_SIZE_256 = smem_bytes(256, ns, 3, 2);

      const bool persistent_256 = (num_items_256_cluster > wave_cap);
      const int launch_ctas_256 = persistent_256 ? wave_cap : num_items_256_cluster;

      const ProblemInfo* d_probs_ptr = (const ProblemInfo*)d_probs_cache.data_ptr();
      const WorkItem* d_work_ptr = (const WorkItem*)d_work_cache_256_cluster.data_ptr();

      cudaLaunchConfig_t config = {};
      config.gridDim = dim3(launch_ctas_256);
      config.blockDim = dim3(TMA_NUM_WARPS * WARP_SIZE);
      config.dynamicSmemBytes = SMEM_SIZE_256;

      cudaLaunchAttribute launch_attrs[1];
      launch_attrs[0].id = cudaLaunchAttributeClusterDimension;
      launch_attrs[0].val.clusterDim = {CL256, 1, 1};
      config.attrs = launch_attrs;
      config.numAttrs = 1;

      if (persistent_256) {
        CUDA_CHECK(cudaLaunchKernelEx(&config, grouped_gemm_kernel_v4<true, 128, 256, NS_WIDE_HI, CL256>,
            d_probs_ptr, d_work_ptr, num_items_256_cluster));
      } else {
        CUDA_CHECK(cudaLaunchKernelEx(&config, grouped_gemm_kernel_v4<false, 128, 256, NS_WIDE_HI, CL256>,
            d_probs_ptr, d_work_ptr, num_items_256_cluster));
      }
    }
    CUDA_CHECK(cudaGetLastError());
    return C_list;
}

TORCH_LIBRARY(my_module, m) {
    m.def("group_gemm(Tensor[] a, Tensor[] b, Tensor[] c, Tensor[] sfa, Tensor[] sfb, Tensor sizes) -> Tensor[]");
    m.impl("group_gemm", &group_gemm);
}
"""

load_inline(
    "group_gemm",
    cpp_sources="",
    cuda_sources=CUDA_SRC,
    verbose=True,
    is_python_module=False,
    no_implicit_headers=True,
    extra_cuda_cflags=[
        "-O3",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--use_fast_math",
        "--expt-relaxed-constexpr",
        "--relocatable-device-code=false",
        "-lineinfo",
        "-Xptxas=-v",
    ],
    extra_ldflags=["-lcuda"],
)
group_gemm = torch.ops.my_module.group_gemm


@lru_cache(maxsize=128)
def _sizes_cpu_cached(key: tuple[tuple[int, int, int], ...]) -> torch.Tensor:
    return torch.tensor(key, dtype=torch.int64, device="cpu")


def custom_kernel(data: input_t) -> output_t:
    abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data
    A_list = [t[0] for t in abc_tensors]
    B_list = [t[1] for t in abc_tensors]
    C_list = [t[2] for t in abc_tensors]
    sfa_list = [t[0] for t in sfasfb_reordered_tensors]
    sfb_list = [t[1] for t in sfasfb_reordered_tensors]

    key = tuple(tuple(int(v) for v in x) for x in problem_sizes)
    sizes_cpu = _sizes_cpu_cached(key)
    out = group_gemm(A_list, B_list, C_list, sfa_list, sfb_list, sizes_cpu)
    return cast(output_t, out)
scrolls · 1706 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 500897.

⋯ 30 unchanged lines
CUDA_SRC = """
#include <algorithm>
+ #include <cstring>
+ #include <limits>
#include <vector>
#include <unordered_map>
#include <cstdint>
⋯ 96 unchanged lines
:: "r"(mbar_addr), "r"(size) : "memory");
}
- __device__ void mbarrier_wait(int mbar_addr, int phase) {
- uint32_t ticks = 0x989680;
+ constexpr uint32_t MBAR_WAIT_HINT_TMA = 0x989680;
+ constexpr uint32_t MBAR_WAIT_HINT_REUSE = 64;
+
+ __device__ __forceinline__ void mbarrier_wait_hint(int mbar_addr, int phase, uint32_t suspend_time_hint) {
asm volatile(
"{\\n\\t"
".reg .pred P1;\\n\\t"
⋯ 1 unchanged lines
"mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 P1, [%0], %1, %2;\\n\\t"
"@!P1 bra.uni LAB_WAIT;\\n\\t"
"}"
- :: "r"(mbar_addr), "r"(phase), "r"(ticks)
+ :: "r"(mbar_addr), "r"(phase), "r"(suspend_time_hint)
);
}
+ __device__ __forceinline__ void mbarrier_wait_tma(int mbar_addr, int phase) {
+ mbarrier_wait_hint(mbar_addr, phase, MBAR_WAIT_HINT_TMA);
+ }
+
+ __device__ __forceinline__ void mbarrier_wait_reuse(int mbar_addr, int phase) {
+ mbarrier_wait_hint(mbar_addr, phase, MBAR_WAIT_HINT_REUSE);
+ }
+
// TMA: 3D tensor-map load from global -> shared memory.
// The (x,y,z) coordinates correspond to the CUtensorMap encoding in
// init_AB_tmap_u4.
⋯ 70 unchanged lines
);
}
- struct SHAPE {
- static constexpr char _16x256b[] = ".16x256b";
- };
- struct NUM {
- static constexpr char x8[] = ".x8";
- static constexpr char x16[] = ".x16";
- };
+ inline constexpr char SHAPE_16x256b[] = ".16x256b";
+ inline constexpr char NUM_x8[] = ".x8";
+ inline constexpr char NUM_x16[] = ".x16";
template <const char *SHAPE, const char *NUM>
__device__ inline
⋯ 34 unchanged lines
}
__device__ inline void tcgen05_ld_16x256b_x8(float *tmp, int row, int col) {
- tcgen05_ld_32regs<SHAPE::_16x256b, NUM::x8>(tmp, row, col);
+ tcgen05_ld_32regs<SHAPE_16x256b, NUM_x8>(tmp, row, col);
}
__device__ inline void tcgen05_ld_16x256b_x16(float *tmp, int row, int col) {
- tcgen05_ld_64regs<SHAPE::_16x256b, NUM::x16>(tmp, row, col);
+ tcgen05_ld_64regs<SHAPE_16x256b, NUM_x16>(tmp, row, col);
}
__device__ __forceinline__ void tcgen05_dealloc_cols_cta1(uint32_t tmem, int count) {
⋯ 47 unchanged lines
constexpr int TMA_WARP_B = 6;
constexpr int MMA_WARP = 5;
- constexpr bool ENABLE_ROLLING_TMAP_PREFETCH = true;
- constexpr bool ENABLE_MBARRIER_DOUBLE_BUFFER = true;
- constexpr int MBARRIER_CLUSTER_SETS = ENABLE_MBARRIER_DOUBLE_BUFFER ? 2 : 1;
-
constexpr int tma_expected_tx_bytes(bool do_A, int a_bytes, int b_bytes, int sfa_bytes, int sfb_bytes) {
return do_A ? (a_bytes + sfa_bytes) : (b_bytes + sfb_bytes);
}
⋯ 338 unchanged lines
constexpr int SFB_off = SFA_off + TMA_SFA_SMEM_BYTES;
constexpr int MBAR_ARRIVALS = (CLUSTER_SIZE > 1) ? 3 : 2;
- constexpr int MBAR_SETS = (PERSISTENT && CLUSTER_SIZE > 1 && ENABLE_MBARRIER_DOUBLE_BUFFER) ? MBARRIER_CLUSTER_SETS : 1;
+ constexpr int MBAR_SETS = (PERSISTENT && CLUSTER_SIZE > 1) ? 2 : 1;
constexpr int MBAR_SET_BYTES = mbar_bytes(NS, MBAR_ARRIVALS);
const int mbar_base = smem_base + STAGE_SIZE * NS;
⋯ 9 unchanged lines
if (warp_id == 0) {
asm volatile("tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [%0], %1;" :: "r"(smem_base), "r"(TMEM_COLS));
- } else if (!ENABLE_ROLLING_TMAP_PREFETCH && warp_id == 1 && elect_sync()) {
- // Best-effort tensormap prefetch for the first few work items.
-
- for (int i = 0; i < num_items && i < 8; ++i) {
- const ProblemInfo* prob = &global_probs[work_items[i].problem_idx];
- asm volatile("prefetch.tensormap [%0];" :: "l"(&prob->A_tmap) : "memory");
- if constexpr (BLOCK_N == 256) {
- asm volatile("prefetch.tensormap [%0];" :: "l"(&prob->B_tmap_256) : "memory");
- } else {
- asm volatile("prefetch.tensormap [%0];" :: "l"(&prob->B_tmap) : "memory");
- }
- }
}
__syncthreads();
⋯ 39 unchanged lines
const ProblemInfo& prob = global_probs[work.problem_idx];
const int mbar_work_base = mbar_base + ((work_epoch % MBAR_SETS) * MBAR_SET_BYTES);
- if (ENABLE_ROLLING_TMAP_PREFETCH && warp_id == 1 && elect_sync()) {
+ if (warp_id == 1 && elect_sync()) {
int prefetch_idx;
if constexpr (PERSISTENT) {
prefetch_idx = work_idx + gridDim.x;
⋯ 107 unchanged lines
// A is shared across the cluster via multicast. Before reusing a
// ring-buffer stage for the next multicast, rank 0 must wait for
// all CTAs to finish consuming the current stage.
- mbarrier_wait(mbar_work_base + (2*NS + stage) * 8, mma_phase);
- } else {
- mbarrier_wait(mbar_work_base + (NS + stage) * 8, mma_phase);
- }
- } else {
- mbarrier_wait(mbar_work_base + (NS + stage) * 8, mma_phase);
- }
+ mbarrier_wait_reuse(mbar_work_base + (2*NS + stage) * 8, mma_phase);
+ } else {
+ mbarrier_wait_reuse(mbar_work_base + (NS + stage) * 8, mma_phase);
+ }
+ } else {
+ mbarrier_wait_reuse(mbar_work_base + (NS + stage) * 8, mma_phase);
+ }
issue_tma(k_iter, stage);
stage++;
if (stage == NS) {
⋯ 20 unchanged lines
int stage = 0;
int tma_phase = 0;
for (int k_iter = 0; k_iter < num_k_iters; k_iter++) {
- mbarrier_wait(mbar_work_base + stage * 8, tma_phase);
+ mbarrier_wait_tma(mbar_work_base + stage * 8, tma_phase);
const int stage_base = smem_base + stage * STAGE_SIZE;
⋯ 52 unchanged lines
const int last_stage = (num_k_iters - 1) % NS;
const int last_phase = ((num_k_iters - 1) / NS) % 2;
- mbarrier_wait(mbar_work_base + (NS + last_stage) * 8, last_phase);
+ mbarrier_wait_reuse(mbar_work_base + (NS + last_stage) * 8, last_phase);
}
__syncthreads();
⋯ 15 unchanged lines
// In cluster mode we must synchronize across CTAs before re-initializing
// mbarriers; __syncthreads is CTA-local and does not order cluster-wide
// mbarrier arrivals.
- if constexpr (PERSISTENT && CLUSTER_SIZE > 1 && !ENABLE_MBARRIER_DOUBLE_BUFFER) {
- cluster_sync();
- }
-
if constexpr (PERSISTENT) {
if (warp_id == TMA_WARP && elect_sync()) {
// Static grid-stride work distribution avoids global atomics and
// smooths the tail when num_items slightly exceeds one wave.
shared_work_idx = work_idx + gridDim.x;
- if constexpr (!ENABLE_MBARRIER_DOUBLE_BUFFER || CLUSTER_SIZE == 1) {
+ if constexpr (CLUSTER_SIZE == 1) {
if (shared_work_idx < num_items) {
for (int i = 0; i < NS; ++i) {
mbarrier_init(mbar_base + i * 8, 2);
⋯ 14 unchanged lines
__syncthreads();
}
- // Cluster barrier after re-init: ensure all CTAs see re-initialized mbarriers.
- if constexpr (PERSISTENT && CLUSTER_SIZE > 1 && !ENABLE_MBARRIER_DOUBLE_BUFFER) {
- cluster_sync();
- }
-
if constexpr (PERSISTENT) {
work_idx = shared_work_idx;
- if constexpr (CLUSTER_SIZE > 1 && ENABLE_MBARRIER_DOUBLE_BUFFER) {
+ if constexpr (CLUSTER_SIZE > 1) {
work_epoch++;
}
} else {
⋯ 6 unchanged lines
}
}
+ static inline uint64_t max_tmap_rows_u4(const at::Tensor& t, uint64_t global_width) {
+ TORCH_CHECK(global_width >= 256 && (global_width % 256) == 0, "K must be multiple of 256");
+
+ const uint64_t logical_height = (uint64_t)t.size(0);
+ if (t.dim() < 2) {
+ return logical_height;
+ }
+
+ const int64_t elem_size = (int64_t)t.element_size();
+ const int64_t stride0 = t.stride(0);
+ const int64_t row_bytes = (int64_t)(global_width / 2);
+ if (elem_size <= 0 || stride0 <= 0 || (row_bytes % elem_size) != 0) {
+ return logical_height;
+ }
+
+ const int64_t row_elems = row_bytes / elem_size;
+ if (stride0 < row_elems) {
+ return logical_height;
+ }
+
+ const int64_t storage_nbytes = (int64_t)t.storage().nbytes();
+ const int64_t storage_offset_bytes = (int64_t)t.storage_offset() * elem_size;
+ if (storage_offset_bytes > storage_nbytes) {
+ return logical_height;
+ }
+
+ const int64_t available_bytes = storage_nbytes - storage_offset_bytes;
+ if (available_bytes < row_bytes) {
+ return logical_height;
+ }
+
+ const int64_t stride0_bytes = stride0 * elem_size;
+ const uint64_t max_rows = (uint64_t)(1 + (available_bytes - row_bytes) / stride0_bytes);
+ return std::max(logical_height, max_rows);
+ }
+
void init_AB_tmap_u4(
CUtensorMap *tmap,
const void *ptr,
uint64_t global_height, uint64_t global_width,
- uint32_t shared_height, uint32_t shared_width
+ uint32_t shared_height, uint32_t shared_width,
+ uint64_t max_safe_height
) {
TORCH_CHECK(ptr != nullptr, "ptr is null");
TORCH_CHECK(((uintptr_t)ptr % 16) == 0, "ptr must be 16-byte aligned");
TORCH_CHECK(global_width >= 256 && (global_width % 256) == 0, "K must be multiple of 256");
TORCH_CHECK(shared_width == 256, "shared_width must be 256");
+ const uint64_t aligned_height = (global_height + 127ULL) & ~127ULL;
+ if (max_safe_height >= aligned_height) {
+ global_height = aligned_height;
+ }
+
constexpr uint32_t rank = 3;
uint64_t globalDim[rank] = {256, global_height, global_width / 256};
uint64_t globalStrides[rank-1] = {global_width / 2, 128};
⋯ 81 unchanged lines
if (!attrs_set) {
constexpr int SMEM_128_HI_K256 = smem_bytes(128, NS_DEEP_HI, 2);
constexpr int SMEM_256_HI_K256 = smem_bytes(256, NS_WIDE_HI, 2);
- constexpr int SMEM_256_CLUSTER_HI_K256 = smem_bytes(256, NS_WIDE_HI, 3, MBARRIER_CLUSTER_SETS);
+ constexpr int SMEM_256_CLUSTER_HI_K256 = smem_bytes(256, NS_WIDE_HI, 3, 2);
constexpr int SMEM_64_HI_K256 = smem_bytes(64, NS_64_HI, 2);
constexpr int SMEM_64_LO_K256 = smem_bytes(64, NS_64_LO, 2);
⋯ 25 unchanged lines
if (!occ_set) {
constexpr int SMEM_128_HI_K256 = smem_bytes(128, NS_DEEP_HI, 2);
constexpr int SMEM_256_HI_K256 = smem_bytes(256, NS_WIDE_HI, 2);
- constexpr int SMEM_256_CLUSTER_HI_K256 = smem_bytes(256, NS_WIDE_HI, 3, MBARRIER_CLUSTER_SETS);
+ constexpr int SMEM_256_CLUSTER_HI_K256 = smem_bytes(256, NS_WIDE_HI, 3, 2);
constexpr int SMEM_64_HI_K256 = smem_bytes(64, NS_64_HI, 2);
constexpr int SMEM_64_LO_K256 = smem_bytes(64, NS_64_LO, 2);
⋯ 14 unchanged lines
occ_set = true;
}
- std::vector<ProblemInfo> problem_infos(G);
- static thread_local std::vector<WorkItem> cached_work_items_64;
- static thread_local std::vector<WorkItem> cached_work_items_128;
- static thread_local std::vector<WorkItem> cached_work_items_256;
- static thread_local std::vector<WorkItem> cached_work_items_256_cluster;
- static thread_local uint64_t cached_work_hash = 0;
- static thread_local bool cached_work_valid = false;
+ auto options = at::TensorOptions().dtype(at::kByte).device(dev);
+ auto host_pinned_options = at::TensorOptions().dtype(at::kByte).device(at::kCPU).pinned_memory(true);
+ const int64_t probs_bytes = (int64_t)(G * sizeof(ProblemInfo));
+
+ static thread_local at::Tensor h_probs_cache;
+ static thread_local std::vector<ProblemInfo> h_probs_pageable;
+ static thread_local at::Tensor h_work_64_cache;
+ static thread_local at::Tensor h_work_128_cache;
+ static thread_local at::Tensor h_work_256_cache;
+ static thread_local at::Tensor h_work_256_cluster_cache;
+ if ((int64_t)h_probs_pageable.size() < G) {
+ h_probs_pageable.resize((size_t)G);
+ }
+
+ ProblemInfo* problem_infos = h_probs_pageable.data();
+
std::vector<uint8_t> active(G, 0);
std::vector<uint8_t> algo_kind(G, ALGO_128);
std::vector<uint8_t> use_cluster_256(G, 0);
⋯ 38 unchanged lines
work_hash = hash_combine_u64(work_hash, 0);
continue;
}
- const int64_t M = Ms[(size_t)i];
- const int64_t N = Ns[(size_t)i];
+ const int tiles_m = num_tiles_m[(size_t)i];
+ const int tiles_n = num_tiles_n[(size_t)i];
const uint8_t algo = algo_kind[(size_t)i];
const bool is_256_cluster = (use_cluster_256[(size_t)i] != 0);
- work_hash = hash_combine_u64(work_hash, (uint64_t)M);
- work_hash = hash_combine_u64(work_hash, (uint64_t)N);
+ work_hash = hash_combine_u64(work_hash, (uint64_t)tiles_m);
+ work_hash = hash_combine_u64(work_hash, (uint64_t)tiles_n);
work_hash = hash_combine_u64(work_hash, (uint64_t)algo);
work_hash = hash_combine_u64(work_hash, (uint64_t)is_256_cluster);
}
- if (!cached_work_valid || cached_work_hash != work_hash) {
- cached_work_items_64.clear();
- cached_work_items_128.clear();
- cached_work_items_256.clear();
- cached_work_items_256_cluster.clear();
- cached_work_items_64.reserve(G * 32);
- cached_work_items_128.reserve(G * 32);
- cached_work_items_256.reserve(G * 32);
- cached_work_items_256_cluster.reserve(G * 32);
+ int64_t num_items_64_i64 = 0;
+ int64_t num_items_128_i64 = 0;
+ int64_t num_items_256_i64 = 0;
+ int64_t num_items_256_cluster_i64 = 0;
+ for (int64_t i = 0; i < G; i++) {
+ if (!active[(size_t)i]) continue;
+ const int64_t tiles_m = num_tiles_m[(size_t)i];
+ const int64_t tiles_n = num_tiles_n[(size_t)i];
+ const uint8_t algo = algo_kind[(size_t)i];
- int max_tn_64 = 0;
- for (int64_t i = 0; i < G; i++) {
- if (!active[(size_t)i]) continue;
- if (algo_kind[(size_t)i] == ALGO_64) {
- max_tn_64 = std::max(max_tn_64, num_tiles_n[(size_t)i]);
+ if (algo == ALGO_64) {
+ num_items_64_i64 += tiles_m * tiles_n;
+ } else if (algo == ALGO_128) {
+ num_items_128_i64 += tiles_m * tiles_n;
+ } else {
+ if (use_cluster_256[(size_t)i]) {
+ const int64_t tiles_n_padded = ((tiles_n + CLUSTER_SIZE_256 - 1) / CLUSTER_SIZE_256) * CLUSTER_SIZE_256;
+ num_items_256_cluster_i64 += tiles_m * tiles_n_padded;
+ } else {
+ num_items_256_i64 += tiles_m * tiles_n;
}
}
- for (int tn = 0; tn < max_tn_64; tn++) {
- for (int64_t i = 0; i < G; i++) {
- if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_64) continue;
- if (tn >= num_tiles_n[(size_t)i]) continue;
- for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {
- cached_work_items_64.push_back({(int)i, tm, tn});
- }
- }
- }
+ }
- int max_tn_128 = 0;
- for (int64_t i = 0; i < G; i++) {
- if (!active[(size_t)i]) continue;
- if (algo_kind[(size_t)i] == ALGO_128) {
- max_tn_128 = std::max(max_tn_128, num_tiles_n[(size_t)i]);
- }
- }
+ TORCH_CHECK(num_items_64_i64 <= std::numeric_limits<int>::max(), "too many ALGO_64 work items");
+ TORCH_CHECK(num_items_128_i64 <= std::numeric_limits<int>::max(), "too many ALGO_128 work items");
+ TORCH_CHECK(num_items_256_i64 <= std::numeric_limits<int>::max(), "too many ALGO_256 work items");
+ TORCH_CHECK(num_items_256_cluster_i64 <= std::numeric_limits<int>::max(), "too many ALGO_256 cluster work items");
- for (int tn = 0; tn < max_tn_128; tn++) {
- for (int64_t i = 0; i < G; i++) {
- if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_128) continue;
- if (tn >= num_tiles_n[(size_t)i]) continue;
- for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {
- cached_work_items_128.push_back({(int)i, tm, tn});
- }
- }
- }
+ const int num_items_64 = (int)num_items_64_i64;
+ const int num_items_128 = (int)num_items_128_i64;
+ const int num_items_256 = (int)num_items_256_i64;
+ const int num_items_256_cluster = (int)num_items_256_cluster_i64;
- for (int64_t i = 0; i < G; i++) {
- if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_256) continue;
- for (int tm = 0; tm < num_tiles_m[(size_t)i]; tm++) {
- for (int tn = 0; tn < num_tiles_n[(size_t)i]; tn++) {
- if (use_cluster_256[(size_t)i]) {
- cached_work_items_256_cluster.push_back({(int)i, tm, tn});
- } else {
- cached_work_items_256.push_back({(int)i, tm, tn});
- }
- }
- if (use_cluster_256[(size_t)i]) {
- int remainder = num_tiles_n[(size_t)i] % CLUSTER_SIZE_256;
- if (remainder != 0) {
- for (int p = 0; p < CLUSTER_SIZE_256 - remainder; p++) {
- cached_work_items_256_cluster.push_back({(int)i, tm, 0});
- }
- }
- }
- }
- }
-
- cached_work_hash = work_hash;
- cached_work_valid = true;
- }
-
uint64_t probs_hash = 1469598103934665603ULL;
probs_hash = hash_combine_u64(probs_hash, (uint64_t)TMA_BLOCK_K);
for (int64_t i = 0; i < G; i++) {
⋯ 9 unchanged lines
const uint8_t algo = algo_kind[(size_t)i];
- init_AB_tmap_u4(&p.A_tmap, A_list[i].data_ptr(), A_list[i].size(0), K, 128, 256);
+ const uint64_t A_max_safe_height = max_tmap_rows_u4(A_list[i], (uint64_t)K);
+ const uint64_t B_max_safe_height = max_tmap_rows_u4(B_list[i], (uint64_t)K);
+
+ init_AB_tmap_u4(&p.A_tmap, A_list[i].data_ptr(), A_list[i].size(0), K, 128, 256, A_max_safe_height);
const int block_n = block_n_for_algo(algo);
const int tmap_b_height = (block_n == 64) ? 64 : 128;
- init_AB_tmap_u4(&p.B_tmap, B_list[i].data_ptr(), B_list[i].size(0), K, tmap_b_height, 256);
+ init_AB_tmap_u4(&p.B_tmap, B_list[i].data_ptr(), B_list[i].size(0), K, tmap_b_height, 256, B_max_safe_height);
if (algo == ALGO_256) {
- init_AB_tmap_u4(&p.B_tmap_256, B_list[i].data_ptr(), B_list[i].size(0), K, 256, 256);
+ init_AB_tmap_u4(&p.B_tmap_256, B_list[i].data_ptr(), B_list[i].size(0), K, 256, 256, B_max_safe_height);
} else {
p.B_tmap_256 = p.B_tmap;
}
⋯ 11 unchanged lines
probs_hash = hash_combine_u64(probs_hash, (uint64_t)p.Cs1);
}
- if (cached_work_items_64.empty() && cached_work_items_128.empty() && cached_work_items_256.empty() && cached_work_items_256_cluster.empty()) return C_list;
+ if (num_items_64 == 0 && num_items_128 == 0 && num_items_256 == 0 && num_items_256_cluster == 0) return C_list;
- auto options = at::TensorOptions().dtype(at::kByte).device(dev);
+ const int wave_hi_64 = sm_count * occ_64_hi;
+ const int wave_lo_64 = sm_count * occ_64_lo;
+ const bool use_lo_64 = (wave_lo_64 > wave_hi_64) && (num_items_64 > 2 * wave_hi_64);
+ const int wave_cap_64 = use_lo_64 ? wave_lo_64 : wave_hi_64;
+ const int wave_cap_128 = sm_count * occ_128_hi;
+ const int wave_cap_256 = sm_count * occ_256_hi;
+ const int wave_cap_256_cluster = (sm_count * occ_256_cluster_hi / CLUSTER_SIZE_256) * CLUSTER_SIZE_256;
+
+ constexpr int LPT_MIN_WAVES = 3;
+ const bool enable_lpt_64 = (num_items_64 > LPT_MIN_WAVES * wave_cap_64);
+ const bool enable_lpt_128 = (num_items_128 > LPT_MIN_WAVES * wave_cap_128);
+ const bool enable_lpt_256 = (num_items_256 > LPT_MIN_WAVES * wave_cap_256);
+ const bool enable_lpt_256_cluster = (num_items_256_cluster > LPT_MIN_WAVES * wave_cap_256_cluster);
+
static thread_local at::Tensor d_probs_cache;
static thread_local at::Tensor d_work_cache_64;
static thread_local at::Tensor d_work_cache_128;
⋯ 3 unchanged lines
static thread_local uint64_t last_work_hash = 0;
static thread_local bool last_hash_valid = false;
- const int64_t probs_bytes = (int64_t)(G * sizeof(ProblemInfo));
- const int64_t work_bytes_64 = (int64_t)(cached_work_items_64.size() * sizeof(WorkItem));
- const int64_t work_bytes_128 = (int64_t)(cached_work_items_128.size() * sizeof(WorkItem));
- const int64_t work_bytes_256 = (int64_t)(cached_work_items_256.size() * sizeof(WorkItem));
- const int64_t work_bytes_256_cluster = (int64_t)(cached_work_items_256_cluster.size() * sizeof(WorkItem));
+ const int64_t work_bytes_64 = (int64_t)(num_items_64 * sizeof(WorkItem));
+ const int64_t work_bytes_128 = (int64_t)(num_items_128 * sizeof(WorkItem));
+ const int64_t work_bytes_256 = (int64_t)(num_items_256 * sizeof(WorkItem));
+ const int64_t work_bytes_256_cluster = (int64_t)(num_items_256_cluster * sizeof(WorkItem));
bool probs_realloc = false;
+ bool work_realloc = false;
if (!d_probs_cache.defined() || d_probs_cache.device() != dev || d_probs_cache.scalar_type() != at::kByte || d_probs_cache.numel() < probs_bytes) {
d_probs_cache = at::empty({probs_bytes}, options);
probs_realloc = true;
}
if (work_bytes_64 > 0 && (!d_work_cache_64.defined() || d_work_cache_64.device() != dev || d_work_cache_64.scalar_type() != at::kByte || d_work_cache_64.numel() < work_bytes_64)) {
d_work_cache_64 = at::empty({work_bytes_64}, options);
+ work_realloc = true;
}
if (work_bytes_128 > 0 && (!d_work_cache_128.defined() || d_work_cache_128.device() != dev || d_work_cache_128.scalar_type() != at::kByte || d_work_cache_128.numel() < work_bytes_128)) {
d_work_cache_128 = at::empty({work_bytes_128}, options);
+ work_realloc = true;
}
if (work_bytes_256 > 0 && (!d_work_cache_256.defined() || d_work_cache_256.device() != dev || d_work_cache_256.scalar_type() != at::kByte || d_work_cache_256.numel() < work_bytes_256)) {
d_work_cache_256 = at::empty({work_bytes_256}, options);
+ work_realloc = true;
}
if (work_bytes_256_cluster > 0 && (!d_work_cache_256_cluster.defined() || d_work_cache_256_cluster.device() != dev || d_work_cache_256_cluster.scalar_type() != at::kByte || d_work_cache_256_cluster.numel() < work_bytes_256_cluster)) {
d_work_cache_256_cluster = at::empty({work_bytes_256_cluster}, options);
+ work_realloc = true;
}
if (probs_realloc || !last_hash_valid || last_probs_hash != probs_hash) {
- CUDA_CHECK(cudaMemcpyAsync(d_probs_cache.data_ptr(), problem_infos.data(), G * sizeof(ProblemInfo), cudaMemcpyHostToDevice));
+ const bool use_pinned_probs_copy = (probs_bytes >= (64 * 1024));
+ const void* probs_src = problem_infos;
+ if (use_pinned_probs_copy) {
+ if (!h_probs_cache.defined() || h_probs_cache.device().type() != at::kCPU || !h_probs_cache.is_pinned() || h_probs_cache.scalar_type() != at::kByte || h_probs_cache.numel() < probs_bytes) {
+ h_probs_cache = at::empty({probs_bytes}, host_pinned_options);
+ }
+ std::memcpy(h_probs_cache.data_ptr(), problem_infos, (size_t)probs_bytes);
+ probs_src = h_probs_cache.data_ptr();
+ }
+ CUDA_CHECK(cudaMemcpyAsync(d_probs_cache.data_ptr(), probs_src, probs_bytes, cudaMemcpyHostToDevice));
last_probs_hash = probs_hash;
}
- if (!last_hash_valid || last_work_hash != work_hash) {
+ if (work_realloc || !last_hash_valid || last_work_hash != work_hash) {
+ if (work_bytes_64 > 0 && (!h_work_64_cache.defined() || h_work_64_cache.device().type() != at::kCPU || !h_work_64_cache.is_pinned() || h_work_64_cache.scalar_type() != at::kByte || h_work_64_cache.numel() < work_bytes_64)) {
+ h_work_64_cache = at::empty({work_bytes_64}, host_pinned_options);
+ }
+ if (work_bytes_128 > 0 && (!h_work_128_cache.defined() || h_work_128_cache.device().type() != at::kCPU || !h_work_128_cache.is_pinned() || h_work_128_cache.scalar_type() != at::kByte || h_work_128_cache.numel() < work_bytes_128)) {
+ h_work_128_cache = at::empty({work_bytes_128}, host_pinned_options);
+ }
+ if (work_bytes_256 > 0 && (!h_work_256_cache.defined() || h_work_256_cache.device().type() != at::kCPU || !h_work_256_cache.is_pinned() || h_work_256_cache.scalar_type() != at::kByte || h_work_256_cache.numel() < work_bytes_256)) {
+ h_work_256_cache = at::empty({work_bytes_256}, host_pinned_options);
+ }
+ if (work_bytes_256_cluster > 0 && (!h_work_256_cluster_cache.defined() || h_work_256_cluster_cache.device().type() != at::kCPU || !h_work_256_cluster_cache.is_pinned() || h_work_256_cluster_cache.scalar_type() != at::kByte || h_work_256_cluster_cache.numel() < work_bytes_256_cluster)) {
+ h_work_256_cluster_cache = at::empty({work_bytes_256_cluster}, host_pinned_options);
+ }
+
+ WorkItem* h_work_64 = (work_bytes_64 > 0) ? reinterpret_cast<WorkItem*>(h_work_64_cache.data_ptr()) : nullptr;
+ WorkItem* h_work_128 = (work_bytes_128 > 0) ? reinterpret_cast<WorkItem*>(h_work_128_cache.data_ptr()) : nullptr;
+ WorkItem* h_work_256 = (work_bytes_256 > 0) ? reinterpret_cast<WorkItem*>(h_work_256_cache.data_ptr()) : nullptr;
+ WorkItem* h_work_256_cluster = (work_bytes_256_cluster > 0) ? reinterpret_cast<WorkItem*>(h_work_256_cluster_cache.data_ptr()) : nullptr;
+
+ int out_64 = 0;
+ int out_128 = 0;
+ int out_256 = 0;
+ int out_256_cluster = 0;
+
+ int max_tn_64 = 0;
+ for (int64_t i = 0; i < G; ++i) {
+ if (active[(size_t)i] && algo_kind[(size_t)i] == ALGO_64) {
+ max_tn_64 = std::max(max_tn_64, num_tiles_n[(size_t)i]);
+ }
+ }
+ for (int tn = 0; tn < max_tn_64; ++tn) {
+ for (int64_t i = 0; i < G; ++i) {
+ if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_64 || tn >= num_tiles_n[(size_t)i]) continue;
+ for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
+ h_work_64[out_64++] = {(int)i, tm, tn};
+ }
+ }
+ }
+
+ int max_tn_128 = 0;
+ for (int64_t i = 0; i < G; ++i) {
+ if (active[(size_t)i] && algo_kind[(size_t)i] == ALGO_128) {
+ max_tn_128 = std::max(max_tn_128, num_tiles_n[(size_t)i]);
+ }
+ }
+ for (int tn = 0; tn < max_tn_128; ++tn) {
+ for (int64_t i = 0; i < G; ++i) {
+ if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_128 || tn >= num_tiles_n[(size_t)i]) continue;
+ for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
+ h_work_128[out_128++] = {(int)i, tm, tn};
+ }
+ }
+ }
+
+ struct ClusterRow {
+ int problem_idx;
+ int tile_m;
+ int tiles_n;
+ int pad;
+ };
+ std::vector<ClusterRow> cluster_rows;
+ for (int64_t i = 0; i < G; ++i) {
+ if (!active[(size_t)i] || algo_kind[(size_t)i] != ALGO_256) continue;
+ if (!use_cluster_256[(size_t)i]) {
+ for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
+ for (int tn = 0; tn < num_tiles_n[(size_t)i]; ++tn) {
+ h_work_256[out_256++] = {(int)i, tm, tn};
+ }
+ }
+ continue;
+ }
+ const int tiles_n = num_tiles_n[(size_t)i];
+ const int remainder = tiles_n % CLUSTER_SIZE_256;
+ const int pad = (remainder == 0) ? 0 : (CLUSTER_SIZE_256 - remainder);
+ for (int tm = 0; tm < num_tiles_m[(size_t)i]; ++tm) {
+ cluster_rows.push_back({(int)i, tm, tiles_n, pad});
+ }
+ }
+
+ auto sort_by_volume = [&](WorkItem* items, int count) {
+ if (count <= 1) return;
+ std::sort(items, items + count, [&](const WorkItem& a, const WorkItem& b) {
+ const int m_a = std::min(128, (int)Ms[(size_t)a.problem_idx] - a.tile_m * 128);
+ const int m_b = std::min(128, (int)Ms[(size_t)b.problem_idx] - b.tile_m * 128);
+ if (m_a == m_b) return Ks[(size_t)a.problem_idx] > Ks[(size_t)b.problem_idx];
+ return m_a > m_b;
+ });
+ };
+
+ if (enable_lpt_64) sort_by_volume(h_work_64, out_64);
+ if (enable_lpt_128) sort_by_volume(h_work_128, out_128);
+ if (enable_lpt_256) sort_by_volume(h_work_256, out_256);
+
+ if (enable_lpt_256_cluster && !cluster_rows.empty()) {
+ std::sort(cluster_rows.begin(), cluster_rows.end(), [&](const ClusterRow& a, const ClusterRow& b) {
+ const int m_a = std::min(128, (int)Ms[(size_t)a.problem_idx] - a.tile_m * 128);
+ const int m_b = std::min(128, (int)Ms[(size_t)b.problem_idx] - b.tile_m * 128);
+ if (m_a == m_b) return Ks[(size_t)a.problem_idx] > Ks[(size_t)b.problem_idx];
+ return m_a > m_b;
+ });
+ }
+
+ for (const ClusterRow& row : cluster_rows) {
+ for (int tn = 0; tn < row.tiles_n; ++tn) {
+ h_work_256_cluster[out_256_cluster++] = {row.problem_idx, row.tile_m, tn};
+ }
+ for (int p = 0; p < row.pad; ++p) {
+ h_work_256_cluster[out_256_cluster++] = {row.problem_idx, row.tile_m, 0};
+ }
+ }
+
+ TORCH_CHECK(out_64 == num_items_64, "ALGO_64 work-item count mismatch");
+ TORCH_CHECK(out_128 == num_items_128, "ALGO_128 work-item count mismatch");
+ TORCH_CHECK(out_256 == num_items_256, "ALGO_256 work-item count mismatch");
+ TORCH_CHECK(out_256_cluster == num_items_256_cluster, "ALGO_256 cluster work-item count mismatch");
+
if (work_bytes_64 > 0) {
- CUDA_CHECK(cudaMemcpyAsync(d_work_cache_64.data_ptr(), cached_work_items_64.data(), work_bytes_64, cudaMemcpyHostToDevice));
+ CUDA_CHECK(cudaMemcpyAsync(d_work_cache_64.data_ptr(), h_work_64, work_bytes_64, cudaMemcpyHostToDevice));
}
if (work_bytes_128 > 0) {
- CUDA_CHECK(cudaMemcpyAsync(d_work_cache_128.data_ptr(), cached_work_items_128.data(), work_bytes_128, cudaMemcpyHostToDevice));
+ CUDA_CHECK(cudaMemcpyAsync(d_work_cache_128.data_ptr(), h_work_128, work_bytes_128, cudaMemcpyHostToDevice));
}
if (work_bytes_256 > 0) {
- CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256.data_ptr(), cached_work_items_256.data(), work_bytes_256, cudaMemcpyHostToDevice));
+ CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256.data_ptr(), h_work_256, work_bytes_256, cudaMemcpyHostToDevice));
}
if (work_bytes_256_cluster > 0) {
- CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256_cluster.data_ptr(), cached_work_items_256_cluster.data(), work_bytes_256_cluster, cudaMemcpyHostToDevice));
+ CUDA_CHECK(cudaMemcpyAsync(d_work_cache_256_cluster.data_ptr(), h_work_256_cluster, work_bytes_256_cluster, cudaMemcpyHostToDevice));
}
+
last_work_hash = work_hash;
}
last_hash_valid = true;
- if (!cached_work_items_64.empty()) {
- int num_items_64 = (int)cached_work_items_64.size();
-
+ if (num_items_64 > 0) {
const int wave_hi = sm_count * occ_64_hi;
const int wave_lo = sm_count * occ_64_lo;
const bool use_lo = (wave_lo > wave_hi) && (num_items_64 > 2 * wave_hi);
⋯ 26 unchanged lines
}
}
- if (!cached_work_items_128.empty()) {
- int num_items_128 = (int)cached_work_items_128.size();
-
+ if (num_items_128 > 0) {
const int ns = NS_DEEP_HI;
const int occ = occ_128_hi;
const int wave_cap = sm_count * occ;
⋯ 12 unchanged lines
}
}
- if (!cached_work_items_256.empty()) {
- int num_items_256 = (int)cached_work_items_256.size();
-
+ if (num_items_256 > 0) {
const int ns = NS_WIDE_HI;
const int occ = occ_256_hi;
const int wave_cap = sm_count * occ;
⋯ 12 unchanged lines
}
}
- if (!cached_work_items_256_cluster.empty()) {
- int num_items_256_cluster = (int)cached_work_items_256_cluster.size();
+ if (num_items_256_cluster > 0) {
constexpr int CL256 = CLUSTER_SIZE_256;
const int ns = NS_WIDE_HI;
const int occ = occ_256_cluster_hi;
const int wave_cap = (sm_count * occ / CL256) * CL256;
- const int SMEM_SIZE_256 = smem_bytes(256, ns, 3, MBARRIER_CLUSTER_SETS);
+ const int SMEM_SIZE_256 = smem_bytes(256, ns, 3, 2);
const bool persistent_256 = (num_items_256_cluster > wave_cap);
const int launch_ctas_256 = persistent_256 ? wave_cap : num_items_256_cluster;
scrolls · 706 diff lines total

Best evidence level for this revision: reported

JSON